Add fsdp-as-dp-for-attn-rf mesh rule and eval warmup - #5528
Conversation
There was a problem hiding this comment.
Code Review
This pull request introduces performance optimizations for large-scale training and evaluation, including a new custom mesh rule fsdp-as-dp-for-attn-rf that shards the MoE router weight on FSDP, and two new options: warm_eval_input_reshard_before_run_start to warm the eval input resharding path and eval_cache_prefill_in_background to prefill the eval cache in a background thread. Feedback suggests guarding the new unit test against missing tensorflow dependencies to prevent test failures, and notes that shaped_eval_batch is conditionally initialized, which may silently skip the eval reshard warming optimization when loading pre-compiled trainsteps.
Codecov Report❌ Patch coverage is
📢 Thoughts on this report? Let us know! |
…rt-up flags Two independent, default-off additions used by the DeepSeek-V3 671B MLPerf recipe on TPU7x: 1. custom_mesh_and_rule=fsdp-as-dp-for-attn-rf: identical to fsdp-as-dp-for-attn except that the MoE router weight (embed_router) is sharded on fsdp instead of replicated, so its gradient is reduced with a reduce-scatter that overlaps the GMM backward instead of a synchronous TensorCore all-reduce inside the backward scan (-93 ms/step at 1k chips, -96 ms/step at 2k). 2. warm_eval_input_reshard_before_run_start / eval_cache_prefill_in_background: compile the eval input reshard before run_start with an all-zero synthetic batch (no dataset access, same rule as warm_input_reshard_before_run_start) and fill the eval tf.data cache on a daemon thread right after run_start, overlapped with the first training steps. Removes the first-eval excess (0.22 s at 1k chips ... 1.52 s at 8k) from the timed region; per-step losses are bit-identical.
6b2f5fc to
451b3bc
Compare
NuojCheng
left a comment
There was a problem hiding this comment.
First custom mesh not submitted by me, Niiice!
…ter-gradient reduction
Description
Adds the
fsdp-as-dp-for-attn-rfcustom mesh and logical-axis rule (embed_routersharded onfsdp) and two pre-run_start / background eval startup flags (warm_eval_input_reshard_before_run_start,eval_cache_prefill_in_background) for MLPerf v6.1 DeepSeek-V3 671B training on TPU7x.custom_mesh_and_rule=fsdp-as-dp-for-attn-rf(-96 ms/stepat 2k,-98 ms/stepat 1k,-85 ms/stepat 4k):fsdp-as-dp-for-attnthat shards the MoE router weight (embed_router: ['fsdp']) instead of replicating it acrossfsdp.all-reduceof the router weight gradient over all devices in the backward scan loop with areduce-scatteroverfsdp(offloaded to SparseCore alongside the forward/rematall-gather).warm_eval_input_reshard_before_run_start&eval_cache_prefill_in_background(-1.52 son first eval at step 42):warm_eval_input_reshard_before_run_start: warms the eval inputdevice_putreshard program (P(data, fsdp, expert)->P(data, fsdp, None)) beforerun_startusing a synthetic zero array (no eval dataset access).eval_cache_prefill_in_background: iterates thec4_mlperfevaltf.datadataset once on a daemon thread immediately afterrun_startso its.cache()is populated in memory during the first ~3 training steps.Tests
tests/unit/train_utils_test.py(TestEvalCachePrefill,TestPrepareBeforeRunStart).1k,2k,4kchips):1k(V7X_8X8X16,ranran-1k-r233-1435): 162.8 s E2E (55steps,2.839 s/step,eval_loss=3.5980 <= 3.6).2k(V7X_8X16X16,ranran-2k-r231-1435): 138.8 s E2E (51steps,2.608 s/step,eval_loss=3.5973 <= 3.6).4k(V7X_8X16X32,ranran-4k-r195-0402): 171.8 s E2E (50steps,eval_loss=3.5975 <= 3.6).Checklist
Before submitting this PR, please make sure (put X in square brackets):
gemini-reviewlabel.