Skip to content

Enabled auxilillary loss free load balancing and sequence wise load balancing for DSV4 [Replacement for PR-https://github.com/AI-Hypercomputer/maxtext/pull/4233] - #4753

Open
dipakg-lang wants to merge 1 commit into
AI-Hypercomputer:mainfrom
dipakg-lang:dsv4_load_balancing

Conversation

@dipakg-lang

@dipakg-lang dipakg-lang commented Aug 6, 2026

Copy link
Copy Markdown
Collaborator

Description

Implement DSV4 auxillary loss free and sequence-wise load balancing via pure NNX MoEBiasVar

Replacement for PR-https://github.com/AI-Hypercomputer/maxtext/pull/4233

The rest of the description includes relevant details and context, examples:

This commit migrates the DeepSeek V4 auxiliary-loss-free routing bias to a pure nnx.Variable (MoEBiasVar), automatically isolating it from the optimizer and standard sequence-wise gradients without requiring the jax.lax.stop_gradient hacks or Optax global masking.

DeepSeek V3 backward compatibility remains completely untouched and functional via the legacy nnx.Param paths.

FIXES: b/509933890
FIXES: b/521990776

Tests

Tested by running training loop with new tiny Deeepseek V4 model added as part of the commit,
here are the logs for testing and commands used for this :

export JAX_PLATFORM_NAME=cpu & export XLA_FLAGS=--xla_force_host_platform_device_count=8 & python3 -m maxtext.trainers.pre_train.train src/maxtext/configs/base.yml override_model_config=True model_name=deepseek4-tiny enable_checkpointing=False base_output_directory=/tmp/maxtext_output/ dataset_type=synthetic hardware=cpu skip_jax_distributed_system=True attention=dot_product per_device_batch_size=1 steps=100 max_target_length=256 async_checkpointing=false dtype=bfloat16 weight_dtype=bfloat16 megablox=False sparse_matmul=False ici_expert_parallelism=-1 sharding_tolerance=1.0 ici_fsdp_parallelism=1 indexer_topk=16 routed_bias=True load_balance_loss_weight=0.0001 routed_bias_update_rate=0.001
To prove the variance reduction and model convergence, the DeepSeek-V4-Flash (284B) model was trained for 300 steps on a Ironwood cluster. The following analysis was performed:

* Run A (Baseline): Variance at Step 300 with NO load balancing active (routed_bias_update_rate=0.0, load_balance_loss_weight=0.0). Shows natural degradation and expert collapse.
* Run B (With Load Balancing): Variance at Step 300 with FULL load balancing active (routed_bias_update_rate=0.001, load_balance_loss_weight=0.0001). Shows healthy token distribution and stable loss curve comparable or better than the PR-4497 base architecture.

=== DeepSeek V4 Load Balancing Variance Analysis (Step 300) ===

| Layer Index | Routing Type | Baseline (No LB) | With Load Balancing | Improvement |
|-------------|--------------|------------------|---------------------|-------------|
|           0 | Hash Routed  |        587520.00 |           587520.00 |       0.00% |
|           1 | Hash Routed  |        587520.00 |           587520.00 |       0.00% |
|           2 | Hash Routed  |        587520.00 |           587520.00 |       0.00% |
|           3 | Top-K Routed |         14623.82 |             2447.45 |      83.26% |
|           4 | Top-K Routed |         15871.75 |             3436.97 |      78.35% |
|           5 | Top-K Routed |         15261.11 |             3110.88 |      79.62% |
|           6 | Top-K Routed |         15717.74 |             1812.94 |      88.47% |
|           7 | Top-K Routed |         15728.96 |             2615.92 |      83.37% |
|           8 | Top-K Routed |         15999.51 |             4967.23 |      68.95% |
|           9 | Top-K Routed |         15240.05 |             3919.24 |      74.28% |
|          10 | Top-K Routed |         16086.67 |             2348.59 |      85.40% |
|          11 | Top-K Routed |         16092.98 |             2607.32 |      83.80% |
|          12 | Top-K Routed |         14255.88 |             2896.68 |      79.68% |
|          13 | Top-K Routed |         14632.31 |             4420.16 |      69.79% |
|          14 | Top-K Routed |         15010.89 |             4668.00 |      68.90% |
|          15 | Top-K Routed |         16004.66 |             2190.59 |      86.31% |
|          16 | Top-K Routed |         14625.14 |             2506.02 |      82.86% |
|          17 | Top-K Routed |         13549.16 |             2880.44 |      78.74% |
|          18 | Top-K Routed |         15119.30 |             3142.03 |      79.22% |
|          19 | Top-K Routed |         14413.22 |             4134.32 |      71.32% |
|          20 | Top-K Routed |         14239.44 |             2910.45 |      79.56% |
|          21 | Top-K Routed |         13616.19 |             2687.55 |      80.26% |
|          22 | Top-K Routed |         14360.10 |             1947.84 |      86.44% |
|          23 | Top-K Routed |         14887.30 |             3880.00 |      73.94% |
|          24 | Top-K Routed |         14978.84 |             4074.12 |      72.80% |
|          25 | Top-K Routed |         14996.25 |             1401.96 |      90.65% |
|          26 | Top-K Routed |         14232.96 |             3889.90 |      72.67% |
|          27 | Top-K Routed |         14977.17 |             1858.18 |      87.59% |
|          28 | Top-K Routed |         14378.70 |             2400.12 |      83.31% |
|          29 | Top-K Routed |         13691.07 |             2181.08 |      84.07% |
|          30 | Top-K Routed |         15055.45 |             3521.05 |      76.61% |
|          31 | Top-K Routed |         14677.20 |             5095.23 |      65.28% |
|          32 | Top-K Routed |         16188.89 |             3978.03 |      75.43% |
|          33 | Top-K Routed |         14369.13 |             2921.98 |      79.66% |
|          34 | Top-K Routed |         14483.14 |             5714.06 |      60.55% |
|          35 | Top-K Routed |         15321.41 |             2911.53 |      81.00% |
|          36 | Top-K Routed |         13667.09 |             2495.93 |      81.74% |
|          37 | Top-K Routed |         14704.00 |             3494.09 |      76.24% |
|          38 | Top-K Routed |         14339.20 |             2385.22 |      83.37% |
|          39 | Top-K Routed |         14203.30 |             2454.55 |      82.72% |
|          40 | Top-K Routed |         13582.94 |             2371.51 |      82.54% |
|          41 | Top-K Routed |         16522.23 |             3767.62 |      77.20% |
|          42 | Top-K Routed |         16277.68 |             1494.26 |      90.82% |
|-------------|--------------|------------------|---------------------|-------------|
| TOTAL/AVG   | Top-K Only   |        595982.83 |           123941.04 |      79.20% |

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 implements DeepSeek V4's auxiliary-loss-free routing bias and sequence-wise load balancing mechanisms in MaxText. It introduces a custom MoEBiasVar to isolate routing bias updates from the global optimizer, adds sequence-wise load balancing loss, and updates configurations, decoder layers, and the training loop to support these features. The code review feedback identifies three key improvement opportunities: refactoring the bias initialization in GateLogit to avoid immediately overwriting an nnx.Param with MoEBiasVar while maintaining RNG consistency, removing an unused self.layers list initialization in nnx_decoders.py, and replacing a brittle string-based type check with isinstance in train.py.

Comment thread src/maxtext/layers/moe.py
Comment thread src/maxtext/layers/nnx_decoders.py Outdated
Comment thread src/maxtext/trainers/pre_train/train.py
…ia pure NNX MoEBiasVar

This commit migrates the DeepSeek V4 auxiliary-loss-free routing bias to a pure nnx.Variable (MoEBiasVar), automatically isolating it from the optimizer and standard sequence-wise gradients without requiring the jax.lax.stop_gradient hacks or Optax global masking.

DeepSeek V3 backward compatibility remains completely untouched and functional via the legacy nnx.Param paths.

To prove the variance reduction and model convergence, the DeepSeek-V4-Flash (284B) model was trained for 300 steps on a Ironwood cluster. The following analysis was performed:

* Run A (Baseline): Variance at Step 300 with NO load balancing active (routed_bias_update_rate=0.0, load_balance_loss_weight=0.0). Shows natural degradation and expert collapse.
* Run B (With Load Balancing): Variance at Step 300 with FULL load balancing active (routed_bias_update_rate=0.001, load_balance_loss_weight=0.0001). Shows healthy token distribution and stable loss curve comparable or better than the PR-4497 base architecture.

=== DeepSeek V4 Load Balancing Variance Analysis (Step 300) ===

| Layer Index | Routing Type | Baseline (No LB) | With Load Balancing | Improvement |
|-------------|--------------|------------------|---------------------|-------------|
|           0 | Hash Routed  |        587520.00 |           587520.00 |       0.00% |
|           1 | Hash Routed  |        587520.00 |           587520.00 |       0.00% |
|           2 | Hash Routed  |        587520.00 |           587520.00 |       0.00% |
|           3 | Top-K Routed |         14623.82 |             2447.45 |      83.26% |
|           4 | Top-K Routed |         15871.75 |             3436.97 |      78.35% |
|           5 | Top-K Routed |         15261.11 |             3110.88 |      79.62% |
|           6 | Top-K Routed |         15717.74 |             1812.94 |      88.47% |
|           7 | Top-K Routed |         15728.96 |             2615.92 |      83.37% |
|           8 | Top-K Routed |         15999.51 |             4967.23 |      68.95% |
|           9 | Top-K Routed |         15240.05 |             3919.24 |      74.28% |
|          10 | Top-K Routed |         16086.67 |             2348.59 |      85.40% |
|          11 | Top-K Routed |         16092.98 |             2607.32 |      83.80% |
|          12 | Top-K Routed |         14255.88 |             2896.68 |      79.68% |
|          13 | Top-K Routed |         14632.31 |             4420.16 |      69.79% |
|          14 | Top-K Routed |         15010.89 |             4668.00 |      68.90% |
|          15 | Top-K Routed |         16004.66 |             2190.59 |      86.31% |
|          16 | Top-K Routed |         14625.14 |             2506.02 |      82.86% |
|          17 | Top-K Routed |         13549.16 |             2880.44 |      78.74% |
|          18 | Top-K Routed |         15119.30 |             3142.03 |      79.22% |
|          19 | Top-K Routed |         14413.22 |             4134.32 |      71.32% |
|          20 | Top-K Routed |         14239.44 |             2910.45 |      79.56% |
|          21 | Top-K Routed |         13616.19 |             2687.55 |      80.26% |
|          22 | Top-K Routed |         14360.10 |             1947.84 |      86.44% |
|          23 | Top-K Routed |         14887.30 |             3880.00 |      73.94% |
|          24 | Top-K Routed |         14978.84 |             4074.12 |      72.80% |
|          25 | Top-K Routed |         14996.25 |             1401.96 |      90.65% |
|          26 | Top-K Routed |         14232.96 |             3889.90 |      72.67% |
|          27 | Top-K Routed |         14977.17 |             1858.18 |      87.59% |
|          28 | Top-K Routed |         14378.70 |             2400.12 |      83.31% |
|          29 | Top-K Routed |         13691.07 |             2181.08 |      84.07% |
|          30 | Top-K Routed |         15055.45 |             3521.05 |      76.61% |
|          31 | Top-K Routed |         14677.20 |             5095.23 |      65.28% |
|          32 | Top-K Routed |         16188.89 |             3978.03 |      75.43% |
|          33 | Top-K Routed |         14369.13 |             2921.98 |      79.66% |
|          34 | Top-K Routed |         14483.14 |             5714.06 |      60.55% |
|          35 | Top-K Routed |         15321.41 |             2911.53 |      81.00% |
|          36 | Top-K Routed |         13667.09 |             2495.93 |      81.74% |
|          37 | Top-K Routed |         14704.00 |             3494.09 |      76.24% |
|          38 | Top-K Routed |         14339.20 |             2385.22 |      83.37% |
|          39 | Top-K Routed |         14203.30 |             2454.55 |      82.72% |
|          40 | Top-K Routed |         13582.94 |             2371.51 |      82.54% |
|          41 | Top-K Routed |         16522.23 |             3767.62 |      77.20% |
|          42 | Top-K Routed |         16277.68 |             1494.26 |      90.82% |
|-------------|--------------|------------------|---------------------|-------------|
| TOTAL/AVG   | Top-K Only   |        595982.83 |           123941.04 |      79.20% |

Raw data logs and loss curve trajectories collected for this analysis:

[Run : With Load Balancing Logs]: https://paste.googleplex.com/5739467624284160
[Run : PR-4497 Legacy Baseline Logs]: https://paste.googleplex.com/6184775940440064

@RissyRan RissyRan 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.

Could you help add original PR link to the description as well?

@parambole parambole 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.

LGTM

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