Skip to content

Disallow use_indexer config combinations that zero the entire training objective - #5280

Open
AusarYao28 wants to merge 1 commit into
mainfrom
yaoyuchen/ds4-pr4-zero-loss-guard
Open

AusarYao28 wants to merge 1 commit into
mainfrom
yaoyuchen/ds4-pr4-zero-loss-guard

Conversation

@AusarYao28

@AusarYao28 AusarYao28 commented Sep 18, 2026 •

Copy link
Copy Markdown
Collaborator

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.

  • validate_train_config rejects use_indexer=True with indexer_sparse_training=False and
    indexer_loss_scaling_factor <= 0.0.
  • deepseek4-284b.yml and deepseek4-tiny.yml default indexer_sparse_training: true, so a plain model_name
    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.
  • Tests added to tests/unit/train_utils_test.py::TestValidateTrainConfig.

FIXES: b/563046810

Tests

  • tests/unit/train_utils_test.py::TestValidateTrainConfig: 13 passed.
  • train_utils_test + attention_compressed_test + csa_streamindex_test + deepseek_v4_rope_test +
    deepseek_v4_indexer_loss_test: 57 passed, 2 skipped.
  • pyconfig + validate_train_config with plain model_name=deepseek4-284b / deepseek4-tiny: pass with this
    change; without the yml default, deepseek4-284b raises "zeroes the entire training objective".
  • E2E (v5p-128, pre-review code): https://cloudlogging.app.goo.gl/btd35GVEi93uhmUq5

Checklist

  • 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 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.

Comment thread src/maxtext/trainers/pre_train/train.py Outdated
Comment on lines 218 to 224
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

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

medium

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.0

@AusarYao28
AusarYao28 added this pull request to stack #5284 September 18, 2026 21:26
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr3-checkpoint-tid2eid branch from 9276829 to fcb37de Compare September 18, 2026 21:38
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr4-zero-loss-guard branch from e538896 to 1dc3141 Compare September 18, 2026 21:38
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr3-checkpoint-tid2eid branch from fcb37de to 7dd5f86 Compare September 22, 2026 23:57
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr4-zero-loss-guard branch from 1dc3141 to 4386164 Compare September 22, 2026 23:57
@codecov

codecov Bot commented Sep 23, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.

📢 Thoughts on this report? Let us know!

@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr4-zero-loss-guard branch from 4386164 to c6b3628 Compare September 23, 2026 01:48
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr3-checkpoint-tid2eid branch 2 times, most recently from 1081036 to d0fc29f Compare September 23, 2026 18:25
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr4-zero-loss-guard branch from c6b3628 to 3761a46 Compare September 23, 2026 18:25
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr3-checkpoint-tid2eid branch from d0fc29f to 7ee4589 Compare September 24, 2026 18:17
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr4-zero-loss-guard branch from 3761a46 to 45df4dc Compare September 24, 2026 18:17
Base automatically changed from yaoyuchen/ds4-pr3-checkpoint-tid2eid to yaoyuchen/ds4-pr2-yarn-rope September 24, 2026 18:17
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr2-yarn-rope branch 2 times, most recently from 9df1c92 to 13ea3e7 Compare September 28, 2026 22:58
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr4-zero-loss-guard branch from 18da2cd to 3cf5bdd Compare September 28, 2026 22:58
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr2-yarn-rope branch from 13ea3e7 to b84c309 Compare September 29, 2026 17:46
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr4-zero-loss-guard branch from 3cf5bdd to 538f42a Compare September 29, 2026 17:46
@AusarYao28
AusarYao28 removed this pull request from stack #5284 September 29, 2026 18:13
@AusarYao28
AusarYao28 added this pull request to stack #5444 September 29, 2026 18:13
@@ -0,0 +1,73 @@
"""Unit tests for the zero-objective indexer configuration validator."""

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.

Add license header here

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

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):

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.

These tests can be moved into tests/unit/train_utils_test.py::TestValidateTrainConfig

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

The tests moved into train_utils_test.py::TestValidateTrainConfig, and the standalone file was deleted, so no new file needs a header.

Comment thread src/maxtext/utils/train_utils.py Outdated

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(

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.

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)

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

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.

@AusarYao28
AusarYao28 removed this pull request from stack #5444 September 30, 2026 16:46
@AusarYao28
AusarYao28 added this pull request to stack #5468 September 30, 2026 16:47
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr2-yarn-rope branch from b84c309 to a0b0217 Compare September 30, 2026 17:48
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr4-zero-loss-guard branch 3 times, most recently from 044af50 to ebc54b6 Compare October 2, 2026 18:27
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr2-yarn-rope branch 2 times, most recently from 8b791c7 to fd628d6 Compare October 2, 2026 19:24
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr4-zero-loss-guard branch from ebc54b6 to d484ae4 Compare October 2, 2026 19:24
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr4-zero-loss-guard branch from d484ae4 to 7ff4aa7 Compare October 8, 2026 00:24
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr2-yarn-rope branch from fd628d6 to cd5bcc8 Compare October 8, 2026 00:24
Base automatically changed from yaoyuchen/ds4-pr2-yarn-rope to main October 8, 2026 05:30
@AusarYao28
AusarYao28 removed this pull request from stack #5468 October 8, 2026 05:35
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr4-zero-loss-guard branch from 7ff4aa7 to ceba669 Compare October 8, 2026 05:46

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

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.

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.

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.py should be reverted;
  • deepseek3.2-671b / deepseek-custom need no YAML change;
  • the DSv4 warm-up command in Run_DeepSeek.md works unchanged.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

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.

@AusarYao28 AusarYao28 Oct 8, 2026 •

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

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
@AusarYao28
AusarYao28 force-pushed the yaoyuchen/ds4-pr4-zero-loss-guard branch from ceba669 to d07cd58 Compare October 8, 2026 17:33
@AusarYao28 AusarYao28 changed the title Disallow use_indexer config combinations that zero the entire training objective Keep the LM loss when the indexer loss is disabled Oct 8, 2026
@AusarYao28 AusarYao28 changed the title Keep the LM loss when the indexer loss is disabled Disallow use_indexer config combinations that zero the entire training objective Oct 8, 2026
@AusarYao28

Copy link
Copy Markdown
Collaborator Author

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.

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.

This branch has not been deployed

No deployments
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.

4 participants