[tunix] Accumulate weighted loss gradients globally - #1744
Open
ethannnnnn wants to merge 2 commits into
Open
Conversation
Define a target-aligned, batch-major diffusion batch contract and typed adapter/scorer protocols without depending on MaxText or a specific training algorithm. Validate shapes and dtypes at construction and scoring boundaries, while preserving JAX pytree, JIT, and sharding compatibility. Tests: 10 diffusion contract tests; pyink/isort; pylint; pyrefly; py_compile.
Accumulate LossOutput gradients as unreduced sums and normalize once by the total denominator across microbatches. Preserve denominator-one behavior for scalar losses and return zero gradients when every weight is zero. Select auxiliary-metric reducers by value type in training and evaluation: globally combine weighted metrics while averaging ordinary scalar metrics. Reject per-key type changes across microbatches and preserve consistent epsilon and minimum-denominator bounds during global reduction. Preserve the dtype selected by each Optax optimizer-state initializer across conditional update and skip branches. This keeps explicit bf16 moments in bf16, retains explicit fp32 moments, and prevents Flax NNX branch-type mismatches without special-casing a particular accumulation count. Tests cover weighted and fractional denominators, zero-weight batches, mixed weighted/plain train and eval metrics, reducer invariants, and a real PeftTrainer + nnx.jit matrix over direct/injected AdamW and gradient accumulation counts 1 and 2. The complete PeftTrainer suite passes 56 tests; the cumulative focused validation passes 119 tests with six optional engine tests deselected. Ruff and git diff checks pass.
ethannnnnn
requested review from
abheesht17,
hgao327,
jiangyangmu,
lc5211,
s-noghabi,
sizhit2,
tianshub and
wang2yn84
as code owners
July 23, 2026 22:43
|
Thanks for your pull request! It looks like this may be your first contribution to a Google open source project. Before we can look at your pull request, you'll need to sign a Contributor License Agreement (CLA). View this failed invocation of the CLA check for more information. For the most up to date status, view the checks section at the bottom of the pull request. |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Motivation
Normalizing each microbatch independently biases gradients when microbatches contain different numbers or fractional weights of valid tokens. Diffusion objectives require one global numerator and denominator.
Scope
LossOutputgradients as unreduced weighted sums.diagnostics alongside weighted objectives.
Design
For each microbatch, the trainer differentiates the loss numerator and accumulates its denominator. The final gradient is the sum of numerator gradients divided by the sum of denominators. A zero total denominator returns zero gradients without division by zero. Scalar/legacy loss functions use denominator one and retain their existing mean-of-microbatches behavior.
Before an optimizer update, the trainer records the optimizer-state dtype tree and restores floating leaves to that contract after Optax returns. This keeps
nnx.condupdate and skip branches type-identical for bfloat16 models without globally promoting optimizer state at construction time.Buffered
WeightedMetricvalues are combined as one numerator and denominatorin both training and evaluation. Plain array metrics continue to use their
existing arithmetic mean. The reducer is selected from the metric value
itself, so a weighted primary loss does not incorrectly force unrelated RL
diagnostics through the weighted metric API. A metric key cannot change kind
within one buffer, and combined weighted metrics must agree on
epsandmin_denom; inconsistent inputs fail before producing misleading telemetry.Compatibility
Ordinary scalar losses and plain auxiliary diagnostics are unchanged.
LossOutputcallers receive mathematically correct weighted reduction forweighted training and evaluation metrics. Direct and injected Optax
transformations retain their configured state dtype.
Extensibility
The reduction is objective agnostic and supports sparse masks, fractional weights, variable sequence lengths, and future weighted objectives beyond diffusion.
Tests
PeftTrainerregressions cover unequal and fractional denominators, all-zero weights, gradients, and metric aggregation.tests/sft/peft_trainer_test.py: 56 passed, covering mixed weighted/plaintrain and evaluation metrics, per-key kind consistency,
denominator-bound consistency, unequal and fractional denominators,
all-zero weights, and optimizer-state dtype regressions.
PeftTrainertests and 63 focused RL/role-sharding tests. Six optionalvLLM/SGL-JAX integration tests were deselected.
updates after the fix, with finite losses and rewards, nonzero gradient
norms, a finalized step checkpoint, and a verified model export. This is
correctness smoke evidence, not a throughput or model-quality claim.
git diff --checkpassed.
Known limitations
This PR does not add a diffusion loss by itself; it provides the denominator-aware trainer primitive used by later adapters.
Stack
Depends on the preceding upstream PR: #1743
Tunix block-diffusion design document