Skip to content

[AutoTP] Support tied embeddings for row-parallel output-head training - #8680

Open
jinyouzhi wants to merge 4 commits into
deepspeedai:masterfrom
jinyouzhi:fix/autotp-row-head-tie-snapshot
Open

jinyouzhi wants to merge 4 commits into
deepspeedai:masterfrom
jinyouzhi:fix/autotp-row-head-tie-snapshot

Conversation

@jinyouzhi

@jinyouzhi jinyouzhi commented Sep 27, 2026 •

Copy link
Copy Markdown
Contributor

Follow-up to #8538 (part of #8173).

Summary

Row-parallel output-head training (#8538) rejected an lm_head/embed_out whose weight is tied to
the input embedding, and that rejection could be bypassed silently (see Bug below). This PR fixes
the bug and supports the tied case by sharding the tied nn.Embedding along the same hidden
dimension as the row-parallel head, so both share one physical per-rank shard.

Bug: tie silently broken when the embedding is visited first

The tie was detected by scanning self.module.modules() during the replacement walk. The walk
visits embed_tokens before lm_head, so when partition_config also matches the embedding,
_slice_embedding swaps it for a new module with a new Parameter first. By the time lm_head was
checked, no other module aliased its weight, the check passed, and the head was sharded while the
embedding kept a separate copy — with no error.

AutoTPConfig(layer_specs=[
    TPLayerSpec(patterns=[r".*embed_tokens\.weight$"], partition_type=PartitionType.ROW),
    TPLayerSpec(patterns=[r".*lm_head\.weight$"], partition_type=PartitionType.ROW),
])
# before: no error, model.lm_head.weight is model.embed_tokens.weight -> False
# after:  both modules share one hidden-dimension shard

Fixed by snapshotting weight ties from the original, pre-replacement model in AutoTP.__init__,
following the approach the vocab-parallel path already takes for the same ordering issue
(introduced in #8309, now _tied_vocab_parallel_embedding_ids).

Changes

  • HiddenParallelEmbedding (layers.py): looks up each token's hidden-dimension slice and
    all-gathers the slices into the replicated activation (GatherFromTensorParallelRegion, so
    uneven hidden shards work). When tied, it reuses the row-parallel head's already-sharded weight
    Parameter and partition sizes; its universal-checkpoint metadata matches LinearAllreduce's.
    padding_idx, scale_grad_by_freq, and Gemma3's scaled embedding are preserved; max_norm,
    sparse, and other custom forwards are rejected.
  • AutoTP (auto_tp.py):
    • Snapshots each module's weight-tie partners from the original model, and records embeddings
      tied to a row-parallel training head so _slice_embedding leaves them alone regardless of
      traversal order.
    • _create_row_parallel_layer validates all tie partners before partitioning the shared
      weight, then replaces every tied embedding alias with a HiddenParallelEmbedding. A tie to a
      non-nn.Embedding module still raises; a conflicting explicit embedding spec is superseded
      with a warning.
  • Docs: autotp-training.md describes the tied path and notes that vocab_parallel_lm_head is
    usually preferable for tied models, since a row-parallel head all-reduces full vocabulary logits.

Not changed: untied embeddings that match a spec still go through _slice_embedding, whose
un-gathered hidden slice is intentional for per-head embeddings such as T5's
relative_attention_bias.

Tests

  • Unit (test_tp_partition_config_path.py): tie is shared with/without an embedding spec
    visited first (the embedding-spec case reproduces the bug above on master); rejection for
    non-embedding ties and max_norm leaves the model untouched; Gemma3 scaled lookup matches the
    original module.
  • Integration (TestTiedRowParallelOutputHeadTraining, 2 GPUs): three optimizer steps against an
    unsharded reference comparing logits, the full tied gradient (lookup + projection
    contributions), and updated weights, covering even/uneven hidden (32/35), padding_idx,
    embedding-spec ordering, and ZeRO-0/2; universal-checkpoint affine map rebuilds the tied weight
    from uneven shards. Mutation check: detaching the embedding's weight in the lookup fails all
    four training cases.
  • tests/unit/module_inject/ and tests/unit/v1/autotp/: 225 passed, 12 skipped on
    3x NVIDIA GeForce RTX 5090 D.

…ut-head training

The row-parallel output-head tie check scanned the module tree during the
replacement walk. When a tied embedding matched a partition_config spec, it
was sliced into a new Parameter before lm_head was visited, so the check no
longer saw the tie and the head was sharded while the embedding kept a
separate copy. Snapshot weight ties from the original model in __init__ and
consult that snapshot instead.

Signed-off-by: Jin, Youzhi <youzhi.jin@intel.com>
@jinyouzhi
jinyouzhi force-pushed the fix/autotp-row-head-tie-snapshot branch from bf0a8b5 to f96d8bb Compare September 27, 2026 10:39
@jinyouzhi jinyouzhi changed the title [AutoTP] Detect tied weights before replacement for row-parallel output-head training [AutoTP] Support tied embeddings for row-parallel output-head training Sep 27, 2026
Row-parallel output-head training rejected a weight tied to the input
embedding. Shard every tied nn.Embedding along the same hidden dimension
with a new HiddenParallelEmbedding that reuses the head's weight shard and
all-gathers each token's slices into the replicated activation. Tied
embeddings are skipped by _slice_embedding so traversal order cannot break
the tie, and all tie partners are validated before the shared weight is
partitioned.

Signed-off-by: Jin, Youzhi <youzhi.jin@intel.com>
@delock
delock self-requested a review September 28, 2026 14:40
name=embed_full_name,
tp_meta=self.tp_meta)
# Guards re-entry when the replacement walk later reaches this module.
setattr(hidden_parallel_embedding, "replaced", True)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Every test that exercises the tied row-parallel output head registers embed_tokens before lm_head, and in those fixtures the guard setattr(hidden_parallel_embedding, "replaced", True) is never consulted — the embedding is protected by _tied_row_parallel_embedding_ids instead. (A tie alias reached later in the walk does consult it, but no such fixture uses a row-parallel head.)

When the head is visited first the guard becomes load-bearing: _replace_with_config then dispatches the freshly installed HiddenParallelEmbedding (2-dim weight, not an nn.Embedding) into _create_row_parallel_layer and wraps it a second time, and the replaced flag is the only check that stops it — without it embed_tokens ends up as a LinearAllreduce.

A head-first OutputModel case would cover it.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Good catch — added a head-first parametrization in 12ae314. It fails without the replaced guard (embed_tokens becomes LinearAllreduce) and passes with it.

@delock
delock enabled auto-merge October 4, 2026 06:51
@delock
delock added this pull request to the merge queue Oct 4, 2026
@github-merge-queue
github-merge-queue Bot removed this pull request from the merge queue due to failed status checks Oct 4, 2026
@delock
delock added this pull request to the merge queue Oct 4, 2026
@github-merge-queue
github-merge-queue Bot removed this pull request from the merge queue due to failed status checks Oct 4, 2026
@jinyouzhi

Copy link
Copy Markdown
Contributor Author

Thanks for your help with the merge, @delock. I checked the latest run and found two failing cases, which I’ve reproduced. I’m investigating whether they’re due to a bug or rounding differences.

@jinyouzhi

Copy link
Copy Markdown
Contributor Author

Thanks for your help with the merge, @delock. I checked the latest run and found two failing cases, which I’ve reproduced. I’m investigating whether they’re due to a bug or rounding differences.

My suspicion is that this is a precision-related issue. I was able to reproduce the failure with TF32 enabled, but it did not occur when I forced FP32. I’m not sure what the exact CI environment is, though.

I suggest we first address the timeout in modal-torch-latest, then use the resulting failure report to decide on the next steps.

@delock Thanks

@delock

delock commented Oct 4, 2026

Copy link
Copy Markdown
Collaborator

Thanks for your help with the merge, @delock. I checked the latest run and found two failing cases, which I’ve reproduced. I’m investigating whether they’re due to a bug or rounding differences.

My suspicion is that this is a precision-related issue. I was able to reproduce the failure with TF32 enabled, but it did not occur when I forced FP32. I’m not sure what the exact CI environment is, though.

I suggest we first address the timeout in modal-torch-latest, then use the resulting failure report to decide on the next steps.

@delock Thanks

Looks like modal is hitting 5400 sec (1.5 hour limit). Let's watch another nightly run to see what is the actual finish time. It is 1h08m on Oct-4.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants