Repository navigation
Disallow use_indexer config combinations that zero the entire training objective - #5280
AusarYao28 wants to merge 1 commit into
Conversation
There was a problem hiding this comment.
Code Review
This pull request introduces validation logic to prevent training configurations where 'use_indexer' is enabled without 'indexer_sparse_training' and with a non-positive 'indexer_loss_scaling_factor', which would otherwise result in a zeroed training objective. The validation is implemented in 'validate_train_config' and supported by a new unit test. The reviewer correctly identified that the redundant check added to 'loss_fn' in 'train.py' should be removed to maintain clean code, avoid duplication, and prevent potential issues with raising exceptions inside JIT-compiled functions.
| if config.indexer_loss_scaling_factor <= 0.0: | ||
| raise ValueError( | ||
| "use_indexer=True with indexer_sparse_training=False and indexer_loss_scaling_factor<=0.0 " | ||
| "zeroes the entire training objective. Enable indexer_sparse_training or set a positive indexer_loss_scaling_factor." | ||
| ) | ||
| xent_sum = 0.0 | ||
| total_z_loss = 0.0 |
There was a problem hiding this comment.
This configuration check is redundant because the exact same validation is already performed at startup in validate_train_config (inside src/maxtext/utils/train_utils.py). Centralizing configuration validation in validate_train_config is the standard pattern in MaxText, keeping the loss function clean and avoiding duplication. Additionally, raising Python exceptions inside JIT-compiled functions (like loss_fn during tracing) is best avoided for static configuration parameters. We should remove this check from loss_fn.
xent_sum = 0.0
total_z_loss = 0.09276829 to
fcb37de
Compare
e538896 to
1dc3141
Compare
fcb37de to
7dd5f86
Compare
1dc3141 to
4386164
Compare
Codecov Report✅ All modified and coverable lines are covered by tests. 📢 Thoughts on this report? Let us know! |
4386164 to
c6b3628
Compare
1081036 to
d0fc29f
Compare
c6b3628 to
3761a46
Compare
d0fc29f to
7ee4589
Compare
3761a46 to
45df4dc
Compare
9df1c92 to
13ea3e7
Compare
18da2cd to
3cf5bdd
Compare
13ea3e7 to
b84c309
Compare
3cf5bdd to
538f42a
Compare
| @@ -0,0 +1,73 @@ | |||
| """Unit tests for the zero-objective indexer configuration validator.""" | |||
There was a problem hiding this comment.
Add license header here
There was a problem hiding this comment.
File deleted; the tests now live in tests/unit/train_utils_test.py, which already has the header.
| from maxtext.utils.train_utils import validate_train_config | ||
|
|
||
|
|
||
| class TestIndexerZeroObjectiveValidator(unittest.TestCase): |
There was a problem hiding this comment.
These tests can be moved into tests/unit/train_utils_test.py::TestValidateTrainConfig
There was a problem hiding this comment.
The tests moved into train_utils_test.py::TestValidateTrainConfig, and the standalone file was deleted, so no new file needs a header.
|
|
||
| if getattr(config, "use_indexer", False) and not getattr(config, "indexer_sparse_training", False): | ||
| if getattr(config, "indexer_loss_scaling_factor", 0.0) <= 0.0: | ||
| raise ValueError( |
There was a problem hiding this comment.
deepseek4-284b.yml sets use_indexer: true but inherits indexer_sparse_training=false / indexer_loss_scaling_factor=0.0 from base.yml, so a plain model_name=deepseek4-284b run now fails validation.
Note that main's 2_test_deepseek.sh pre-train stage uses exactly those defaults, so this needs to land with/after #5283 or add indexer_sparse_training=true to that script here. Optionally, consider defaulting one of these in deepseek4-284b.yml so the model config is self-consistent.
Suggest defaulting indexer_sparse_training: true / indexer_loss_scaling_factor: 1.0 in deepseek4-284b.yml (and -tiny)
There was a problem hiding this comment.
Confirmed and fixed in this PR, so it no longer depends on #5283 landing first. Both DS4 ymls now set
indexer_sparse_training: true; main's 2_test_deepseek.sh pretrain stage (no indexer flags,
max_target_length=2048) passes validation with it.
I kept indexer_loss_scaling_factor at 0.0 rather than 1.0: the compressed-attention check in types.py
(indexer_topk >= max_target_length // 4 with scaling > 0 raises) runs for every mode, so a 1.0 default
rejects plain model_name=deepseek4-284b at base.yml's max_target_length 2048 (512 >= 512) and the decode
stage at 512 (512 >= 128). The e2e pretrain stage passes indexer_loss_scaling_factor=1.0 explicitly at
max_target_length=4096. If you prefer 1.0 as the default, the options are adding
indexer_loss_scaling_factor=0.0 to every inference invocation or scoping that check to training.
b84c309 to
a0b0217
Compare
044af50 to
ebc54b6
Compare
8b791c7 to
fd628d6
Compare
ebc54b6 to
d484ae4
Compare
d484ae4 to
7ff4aa7
Compare
fd628d6 to
cd5bcc8
Compare
7ff4aa7 to
ceba669
Compare
There was a problem hiding this comment.
Thanks for enhancing the robustness! Approve to unblock. Please address the comments from other reviewers.
if keeping the current approach: setting indexer_sparse_training: true in deepseek4-*.yml breaks the documented DSv4 warm-up command. The CLI passes indexer_sparse_training=false without override_model_config=true, so pyconfig rejects it as a model-config/CLI conflict. Please add override_model_config=true to the warm-up command in Run_DeepSeek.md#L310-L328. Also note that deepseek3.2-671b and deepseek-custom set use_indexer: true with the defaults, so plain train.py runs on them will now fail the new validator.
alternative: fix the root cause in train.py instead (see inline comment). That makes the YAML edits, the validator, and the doc change unnecessary.
There was a problem hiding this comment.
Fix the root cause instead of adding a validator.
The dense warm-up gate in train.py ignores indexer_loss_scaling_factor, so it disagrees with the logits-skip gate and with the config docs:
| Location | Condition |
|---|---|
| train.py#L222-L226 (zeroes LM loss) | use_indexer and not sparse |
| nnx_decoders.py#L2308-L2315 (skips logits) | use_indexer and scale > 0 and not sparse |
| base.yml#L556-L562 (docs) | indexer_sparse_training "only active when indexer_loss_scaling_factor > 0" |
With scale=0 and sparse=false (the defaults), the model computes full logits and then train.py discards them: wasted compute and zero loss.
Suggested change at train.py#L222:
if (config.use_indexer and config.indexer_loss_scaling_factor > 0.0
and not config.indexer_sparse_training) and is_train:This aligns all three conditions. scale=0 then means normal LM training with the indexer used for selection only, which covers continued training of converted V3.2/V4 checkpoints. With this change:
- the YAML edits and the new validator in
train_utils.pyshould be reverted; deepseek3.2-671b/deepseek-customneed no YAML change;- the DSv4 warm-up command in Run_DeepSeek.md works unchanged.
There was a problem hiding this comment.
Done, I went with the root-cause fix. train.py now gates the warm-up branch on indexer_loss_scaling_factor > 0.0, the same way nnx_decoders.py and attention_compressed.py already do. I reverted the validator and both deepseek4 yml edits. deepseek3.2-671b and deepseek-custom work unchanged, and so does the documented DSv4 warm-up command. I updated the two warm-up tests in train_nnx_test.py to set a positive scale, and added a test that scale 0 keeps the LM loss.
There was a problem hiding this comment.
Done, I went with the root-cause fix. train.py now gates the warm-up branch on indexer_loss_scaling_factor > 0.0, the same way nnx_decoders.py and attention_compressed.py already do. I reverted the validator and both deepseek4 yml edits. deepseek3.2-671b and deepseek-custom work unchanged, and so does the documented DSv4 warm-up command. I updated the two warm-up tests in train_nnx_test.py to set a positive scale, and added a test that scale 0 keeps the LM loss.
loss_fn treated use_indexer=True with indexer_sparse_training=False as the indexer dense warm-up and zeroed the LM loss without checking indexer_loss_scaling_factor. The decoder logits skip (nnx_decoders.py), the compressed-attention mask (attention_compressed.py) and the base.yml docs all gate dense warm-up on indexer_loss_scaling_factor > 0. With the base.yml defaults (indexer_sparse_training=False, indexer_loss_scaling_factor=0.0), every use_indexer model (deepseek3.2-671b, deepseek4-284b, deepseek4-tiny, deepseek-custom) computed full logits, discarded them, and trained on total_loss == 0.0 with zero gradients. - Gate the dense warm-up branch in train.py loss_fn on indexer_loss_scaling_factor > 0, matching nnx_decoders.py and attention_compressed.py. With scale 0 the model trains on the LM loss and uses the indexer for top-k selection only. - No model config or validator changes, so use_indexer model configs and the documented warm-up command work unchanged. - train_nnx_test: set a positive indexer_loss_scaling_factor in the two warm-up tests and add a test that scale 0 keeps the LM loss. Bug: b/563046810
ceba669 to
d07cd58
Compare
Both points are moot now: the yml defaults are reverted, so the warm-up command in Run_DeepSeek.md no longer conflicts, and deepseek3.2-671b / deepseek-custom need no change. |
Description
With use_indexer=True and indexer_sparse_training=False, train.py runs the indexer dense warm-up:
loss_fn sets xent_sum = 0.0 and the only signal is the indexer KL loss scaled by
indexer_loss_scaling_factor. base.yml defaults those to false / 0.0, so a model config that enabled
use_indexer without overriding either flag (deepseek4-284b, deepseek4-tiny) trained on total_loss == 0.0
without any error.
indexer_loss_scaling_factor <= 0.0.
run trains on the LM loss. indexer_loss_scaling_factor stays 0.0: with a positive default, the
types.py check indexer_topk < max_target_length // 4 (config-level, not train-only) rejects
deepseek4-284b (indexer_topk 512) at base.yml's max_target_length 2048 and at the decode stage's 512.
FIXES: b/563046810
Tests
deepseek_v4_indexer_loss_test: 57 passed, 2 skipped.
change; without the yml default, deepseek4-284b raises "zeroes the entire training objective".
Checklist
gemini-reviewlabel.