Skip to content

- Lazily promote PeftTrainer.grad_accumulator to persistent mode (allocate_grads=True) and clear the JIT cache on first fwd_bwd() invocation so split fwd_bwd() + update() calls under nnx.cached_partial (cache_nnx_graph=True) and dynamic sequence-packing microsteps (gradient_accumulation_steps == 1) have a pre-allocated, sharded gradient buffer across the JIT boundary, while preserving the zero-allocation fused train_step() fast path for train(). - #2729

Merged
copybara-service[bot] merged 1 commit into
mainfrom
test_994701268
Oct 7, 2026

Conversation

@copybara-service

Copy link
Copy Markdown
  • Lazily promote PeftTrainer.grad_accumulator to persistent mode (allocate_grads=True) and clear the JIT cache on first fwd_bwd() invocation so split fwd_bwd() + update() calls under nnx.cached_partial (cache_nnx_graph=True) and dynamic sequence-packing microsteps (gradient_accumulation_steps == 1) have a pre-allocated, sharded gradient buffer across the JIT boundary, while preserving the zero-allocation fused train_step() fast path for train().
  • Plumb MAX_SEQ_TOKEN_PER_TPU in frozenlake_dist/launcher.sh to the trainer and orchestrator commands and skip the static MINI_BATCH_SIZE * NUM_GENERATIONS % TRAIN_MICRO_BATCH_SIZE divisibility check when sequence packing is enabled.
  • Add unit tests in peft_trainer_v2_test.py covering split fwd_bwd() + update() at depth 1 and dynamic multi-microstep accumulation under cache_nnx_graph=True.

…allocate_grads=True`) and clear the JIT cache on first `fwd_bwd()` invocation so split `fwd_bwd()` + `update()` calls under `nnx.cached_partial` (`cache_nnx_graph=True`) and dynamic sequence-packing microsteps (`gradient_accumulation_steps == 1`) have a pre-allocated, sharded gradient buffer across the JIT boundary, while preserving the zero-allocation fused `train_step()` fast path for `train()`.

- Plumb `MAX_SEQ_TOKEN_PER_TPU` in `frozenlake_dist/launcher.sh` to the trainer and orchestrator commands and skip the static `MINI_BATCH_SIZE * NUM_GENERATIONS % TRAIN_MICRO_BATCH_SIZE` divisibility check when sequence packing is enabled.
- Add unit tests in `peft_trainer_v2_test.py` covering split `fwd_bwd()` + `update()` at depth 1 and dynamic multi-microstep accumulation under `cache_nnx_graph=True`.

PiperOrigin-RevId: 994716064
@copybara-service
copybara-service Bot merged commit 1ba318a into main Oct 7, 2026
2 checks passed
@copybara-service
copybara-service Bot deleted the test_994701268 branch October 7, 2026 00:02

This branch was successfully deployed

1 active deployment
testing — 1ba318a9 Deployed Oct 7, 2026 by copybara-service[bot] via tunix_tpu_unit_tests / run_dev #11076
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant