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
Conversation
There was a problem hiding this comment.
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.
…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
078a044 to
da3533f
Compare
RissyRan
left a comment
There was a problem hiding this comment.
Could you help add original PR link to the description as well?
Description
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 :
Checklist
Before submitting this PR, please make sure (put X in square brackets):