Skip to content

Add fsdp-as-dp-for-attn-rf mesh rule and eval warmup - #5528

Merged
copybara-service[bot] merged 3 commits into
mainfrom
ranran/dev/dsv3-rf-rule-and-eval-startup
Oct 4, 2026
Merged

copybara-service[bot] merged 3 commits into
mainfrom
ranran/dev/dsv3-rf-rule-and-eval-startup

Conversation

@RissyRan

@RissyRan RissyRan commented Oct 4, 2026 •

Copy link
Copy Markdown
Collaborator

Description

Adds the fsdp-as-dp-for-attn-rf custom mesh and logical-axis rule (embed_router sharded on fsdp) 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.

  1. custom_mesh_and_rule=fsdp-as-dp-for-attn-rf (-96 ms/step at 2k, -98 ms/step at 1k, -85 ms/step at 4k):
    • Variant of fsdp-as-dp-for-attn that shards the MoE router weight (embed_router: ['fsdp']) instead of replicating it across fsdp.
    • Replaces the synchronous 58-layer TensorCore all-reduce of the router weight gradient over all devices in the backward scan loop with a reduce-scatter over fsdp (offloaded to SparseCore alongside the forward/remat all-gather).
  2. warm_eval_input_reshard_before_run_start & eval_cache_prefill_in_background (-1.52 s on first eval at step 42):
    • warm_eval_input_reshard_before_run_start: warms the eval input device_put reshard program (P(data, fsdp, expert) -> P(data, fsdp, None)) before run_start using a synthetic zero array (no eval dataset access).
    • eval_cache_prefill_in_background: iterates the c4_mlperf eval tf.data dataset once on a daemon thread immediately after run_start so its .cache() is populated in memory during the first ~3 training steps.

Tests

  • Unit tests in tests/unit/train_utils_test.py (TestEvalCachePrefill, TestPrepareBeforeRunStart).
  • End-to-end MLPerf v6.1 DeepSeek-V3 671B validation on TPU7x (1k, 2k, 4k chips):
    • 1k (V7X_8X8X16, ranran-1k-r233-1435): 162.8 s E2E (55 steps, 2.839 s/step, eval_loss=3.5980 <= 3.6).
    • 2k (V7X_8X16X16, ranran-2k-r231-1435): 138.8 s E2E (51 steps, 2.608 s/step, eval_loss=3.5973 <= 3.6).
    • 4k (V7X_8X16X32, ranran-4k-r195-0402): 171.8 s E2E (50 steps, eval_loss=3.5975 <= 3.6).

Checklist

Before submitting this PR, please make sure (put X in square brackets):

  • I have performed a self-review of my code. For an optional AI review, add the gemini-review label.
  • I have necessary comments in my code, particularly in hard-to-understand areas.
  • I have run end-to-end tests tests and provided workload links above if applicable.
  • I have made or will make corresponding changes to the doc if needed, including adding new documentation pages to the relevant Table of Contents (toctree directive) as explained in our documentation.

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

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.

Comment thread tests/unit/train_utils_test.py
Comment thread src/maxtext/trainers/pre_train/train.py Outdated
@codecov

codecov Bot commented Oct 4, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 31.25000% with 33 lines in your changes missing coverage. Please review.

Files with missing lines Patch % Lines
src/maxtext/utils/train_utils.py 22.85% 23 Missing and 4 partials ⚠️
src/maxtext/trainers/pre_train/train.py 50.00% 4 Missing and 2 partials ⚠️

📢 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.
@RissyRan RissyRan changed the title MLPerf DeepSeek-V3: add fsdp-as-dp-for-attn-rf mesh rule and eval sta… Add fsdp-as-dp-for-attn-rf mesh rule and eval warmup Oct 4, 2026

@NuojCheng NuojCheng left a comment

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.

First custom mesh not submitted by me, Niiice!

@copybara-service
copybara-service Bot merged commit c3d7f89 into main Oct 4, 2026
88 of 90 checks passed
@copybara-service
copybara-service Bot deleted the ranran/dev/dsv3-rf-rule-and-eval-startup branch October 4, 2026 20:39
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