Conversation
…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>
bf0a8b5 to
f96d8bb
Compare
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>
| 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) |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
Good catch — added a head-first parametrization in 12ae314. It fails without the replaced guard (embed_tokens becomes LinearAllreduce) and passes with it.
Signed-off-by: Jin, Youzhi <youzhi.jin@intel.com>
|
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 @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. |
Follow-up to #8538 (part of #8173).
Summary
Row-parallel output-head training (#8538) rejected an
lm_head/embed_outwhose weight is tied tothe 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.Embeddingalong the same hiddendimension 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 walkvisits
embed_tokensbeforelm_head, so whenpartition_configalso matches the embedding,_slice_embeddingswaps it for a new module with a new Parameter first. By the timelm_headwaschecked, no other module aliased its weight, the check passed, and the head was sharded while the
embedding kept a separate copy — with no error.
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 andall-gathers the slices into the replicated activation (
GatherFromTensorParallelRegion, souneven hidden shards work). When tied, it reuses the row-parallel head's already-sharded weight
Parameterand partition sizes; its universal-checkpoint metadata matchesLinearAllreduce'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):tied to a row-parallel training head so
_slice_embeddingleaves them alone regardless oftraversal order.
_create_row_parallel_layervalidates all tie partners before partitioning the sharedweight, then replaces every tied embedding alias with a
HiddenParallelEmbedding. A tie to anon-
nn.Embeddingmodule still raises; a conflicting explicit embedding spec is supersededwith a warning.
autotp-training.mddescribes the tied path and notes thatvocab_parallel_lm_headisusually 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, whoseun-gathered hidden slice is intentional for per-head embeddings such as T5's
relative_attention_bias.Tests
test_tp_partition_config_path.py): tie is shared with/without an embedding specvisited first (the embedding-spec case reproduces the bug above on master); rejection for
non-embedding ties and
max_normleaves the model untouched; Gemma3 scaled lookup matches theoriginal module.
TestTiedRowParallelOutputHeadTraining, 2 GPUs): three optimizer steps against anunsharded 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/andtests/unit/v1/autotp/: 225 passed, 12 skipped on3x NVIDIA GeForce RTX 5090 D.