Files
Model-Optimizer/tests
h-guo18andClaude Opus 5 87f7d1432f fix(speculative): hold the DFlash draft's fp32 master weights in the optimizer (#2483)
### What does this PR do?

Type of change: Bug fix

**Follow-up to #2342**, which split this out on review (commit
`c67784d9`), and a rethink of how
the flag is implemented.

`dflash_fp32_master_weights` exists because the DFlash draft is cast to
the frozen bf16 target's
dtype, so AdamW allocates its moments in bf16 — and bf16 is too coarse
to hold them. At
`beta2=0.999` a single step changes `v` by at most **0.100%**, while the
smallest change bf16 can
represent near `v` is **0.164% mean / 0.388% max** (measured): every
decrease rounds away, `v` only
grows, and the effective step size decays on its own from step 1.

#2342 fixed that by **promoting the draft model to fp32**. Everything
else followed from giving the
model a dtype the rest of it does not have — a bf16 autocast at every
entry point, two transformers
loader hints so `from_pretrained(dtype="auto")` would not round the
draft away, a post-condition
check because those hints fail silently, and a doubled DDP gradient
all-reduce.

**This PR puts the fp32 in the optimizer instead**, where Megatron-LM,
DeepSpeed and apex put it.
`MasterWeightAdamW` holds an fp32 master copy of each non-fp32 parameter
plus fp32 moments in
`self.state[p]`, steps on the master, and copies back at the parameter's
dtype. The model is never
anything but the base dtype, so every one of those follow-on pieces is
deleted, gradients stay
bf16, and the exported drafter is unchanged. What the placement costs is
that wiring the optimizer
becomes the training loop's job:
`EagleTrainerWithAccLog.create_optimizer` builds it, and
`VerifyMasterWeightsCallback` raises at the end of step 1 if the moments
are not fp32.

**The default flips to `True`** — the flag now changes optimizer memory
and optimizer arithmetic
and nothing else. Flipping it on the model-promoted implementation turns
**25 of 259** unit tests
red; flipping it here is **259 passed**. Set it to `False` to reclaim
the memory, about 12 bytes per
draft parameter instead of 4.

<details>
<summary>Three drive-by fixes, independent of the above</summary>

- `_place_draft` is folded back into `modify()` — it fused the draft's
dtype, its device and an
  eager rotary buffer behind one meta guard.
- The module docstring's claim that `DFlashModule` has an `_apply`
meta-buffer fix is removed
  (`grep "def _apply"` matches nothing, and never did).
- #2342's field description no longer lists `evaluation` as a broken
path — `forward`
short-circuits to the base model when `not self.training`, so the draft
never runs there.

</details>

### Usage

No API change. `dflash_fp32_master_weights` now means the *optimizer*
holds fp32 master weights
rather than the draft model being fp32.

### Testing

**1 · The refactor is arithmetically a no-op.** Both implementations run
AdamW on an fp32 tensor,
so given the same starting values and the same gradients the
trajectories are identical — 1000
steps, `weight_decay=0.01`:

```
old fp32 parameter  vs  new fp32 master : bitwise equal = True  (max |diff| 0.0e+00)
exp_avg / exp_avg_sq                    : bitwise equal = True
optimizer state dtypes                  : ['torch.float32']
model parameter dtype                   : torch.bfloat16
```

Initial values have to be matched at bf16 first, or the bf16 arm's
one-time rounding of the draw
shows up as a 2e-4 "difference" that is not arithmetic. With that
controlled, the two
implementations differ only in their *inputs*: gradient precision (fp32
vs bf16 — torch 2.10
requires `grad.dtype == param.dtype`) and that one-time rounding.

**2 · End to end on GPU: the effect survives the refactor.** Qwen3-1.7B
base, real corpus, one GPU
per arm, three arms — pure bf16 (flag off), the #2342 implementation,
and this one — on two
algorithms trained independently, sharing seed, data order and
initialisation within an algorithm.

<img width="2925" height="960" alt="image"
src="https://github.com/user-attachments/assets/40d37059-ab8a-4924-b049-85f76b70b156"
/>


The two fp32 arms sit on top of each other for the whole run while bf16
stays above both, and the
old-vs-new gap is 10–23× smaller than the fp32-vs-bf16 effect it has to
be compared against.

**Acceptance length says the same thing, and settles what the loss could
not.** All six drafters at
the end of those curves were exported and served under vLLM against the
same base, and measured on
MT-Bench (80 prompts, 8 categories, greedy, one request at a time,
`num_speculative_tokens` =
trained `block_size` − 1, every knob but the drafter held fixed):

| | pure bf16 | fp32 in model (#2342) | fp32 in optimizer (this PR) |
new − old | fp32 − bf16 |
|---|---|---|---|---|---|
| `dflash` | 1.3068 | 1.3708 | **1.3666** | −0.0042 &nbsp;`t=−0.90` |
+0.0619 &nbsp;`t=+13.8` |
| `lilicorr` | 1.2536 | 1.2814 | **1.2882** | +0.0068 &nbsp;`t=+1.42` |
+0.0312 &nbsp;`t=+9.2` |

Paired by prompt, n=80. On both algorithms the new-vs-old 95% CI
straddles zero (`dflash`
[−0.0134, +0.0051], `lilicorr` [−0.0028, +0.0164]) while fp32-vs-bf16
does not come close to it,
and the sign of new-vs-old **flips between the two algorithms** — what a
rounding difference looks
like, not a bias. This is also the comparison the training loss could
not give: all three arms are
**exported and served in bf16**, so the old implementation's fp32 draft
weights are rounded at
export exactly as they would be for deployment, and the "its loss was
computed on a more precise
forward" caveat below does not apply. `lilicorr` needs

[vllm-project/vllm#57934](https://github.com/vllm-project/vllm/pull/57934),
applied as an overlay so
that both algorithms are measured on one engine build.

The right panel is the mechanism, and the one signal that depends on
neither the seed nor the
choice of loss statistic: Adam's updates to the draft's RMSNorm gains
are smaller than the bf16 ULP
at 1.0 (0.0078), so in the bf16 arm every one of them rounds away and
the gains never move — not
one of `dflash`'s 14 in 30000 steps, and two of `lilicorr`'s 20 by
3e-06. Both fp32 arms move
all of them, by the same amount.

Two results behind the figure rather than in it. **fp32-vs-bf16 grows
with the horizon** while
old-vs-new does not — on `dflash` −0.129 at 1500 steps → −0.262 at 15000
→ −0.341 at 30000, and on
`lilicorr` −0.191 → −0.220 → −0.285, against an old-vs-new difference
that stays near 0.02 at every
horizon and changes sign between them (−0.026 → +0.028 on `lilicorr`).
That is what a compounding
bias and a rounding difference respectively should look like, and it is
the reason the longer runs
were worth doing. And
**across seeds**, the paired old-vs-new difference at 1500 steps is
+0.0003 (n=6) on `dflash` and
+0.0643 (n=10) on `lilicorr`, both with a 95% CI straddling zero.

<details>
<summary>Limits of the above, stated rather than smoothed over</summary>

At 5 seeds the `lilicorr` paired difference read +0.1610 ± 0.0557
(t=+2.89, 4/5 seeds in the same
direction) — nominally significant, suggesting the new implementation
was genuinely worse there.
Four further `lilicorr` seeds were run against that pre-declared
question; two came back strongly
negative and the estimate settled at +0.0643 (95% CI [−0.086, +0.214]).
The earlier reading was
small-sample noise.

At 1500 steps on `lilicorr` that CI is *not* narrower than the
fp32-vs-bf16 effect it is being
compared against, so the 1500-step sweep alone cannot certify
equivalence there — `lilicorr` is
still at loss 9.3 and deep in its early transient, and it is the long
runs that resolve it. On
`dflash` the 1500-step CI (±0.031) is already 4× tighter than the effect
(−0.129).

One asymmetry the loss comparison cannot separate: the old
implementation held the draft weights in
fp32 *at forward time*, so its training loss was computed on a more
precise forward, while both
implementations export bf16. Any residual advantage it appears to have
is therefore an upper bound.

</details>

**3 · Unit tests.** `tests/unit/torch/speculative/` — **259 passed** on
**transformers 5.0.0** and
**5.3.0**, both ends of the supported `>=5.0,<5.13` (CPU, torch 2.10).
`TestDFlashFp32MasterWeights`
is rewritten for the new mechanism; the two that would have caught the
traps in this design are
`test_resume_does_not_round_the_master_back_down`
(`Optimizer.load_state_dict` casts float state to
its parameter's dtype, so a naive subclass rounds the master and both
moments to bf16 on *every*
resume, silently, with the loss still falling) and
`test_the_callback_refuses_a_loop_that_forgot_the_optimizer`. The rest
cover the draft's dtype with
the flag either way, that no forward path needs an autocast any more,
that plain AdamW really does
leave the moments in bf16, and that an fp32 model allocates no redundant
master. A sharded FSDP2
`DTensor` keeps an fp32 master and fp32 moments through a step.

### Before your PR is "*Ready for review*"

- Is this change backward compatible?: ✅ for artifacts, with one
intentional default change.
The draft's stored dtype goes back to matching the base, as it was
before #2342; existing
checkpoints load unchanged and the exported drafter is unaffected. The
flag now defaults to
**`True`** — the measurements above are the reason, and the cost is fp32
master + fp32 moments for
the draft only. A training loop that builds its own optimizer instead of
using the shipped
`create_optimizer` gets plain AdamW and none of this;
`VerifyMasterWeightsCallback` makes that
  fail loudly at step 1 rather than skip the feature quietly.
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ — no new
dependencies.
- Did you write any new necessary tests?: ✅
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅
- Did you get Claude approval on this PR?: ❌ — draft; will run `/claude
review` before marking ready.

### Additional Information

**On the 7–14% acceptance-length gain quoted in #2342:** that was
measured on the fp32-model
arithmetic and is not re-derived here. What is measured above is the
like-for-like comparison this
PR has to answer — same corpus, same horizon, same serving path, one
implementation swapped.

**History:** commits 1–3 restore the autocast design as it was split
out; commits 4–6 replace it.
Happy to squash before review.

<details>
<summary>Alternatives measured and rejected, so they do not get
re-proposed</summary>

- **Swapping `p.data` to the master and calling `super().step()`**
(reuses all of AdamW, ~20 lines
instead of ~50): bit-identical on ordinary parameters over 25 steps, but
silently wrong under
FSDP2 — assigning `.data` on a `DTensor` parameter updates the wrapper's
reported dtype while the
local shard keeps the model's, so `p.dtype` reads fp32, `p.data.dtype`
reads bf16, and
`zeros_like(p)` allocates the moments in bf16 anyway. CPU tests pass
either way.
- **Narrowing the autocast from `__call__` to `forward`** (while it
still existed): turns 10
Domino/DSpark tests red — the variants apply their heads in their own
`forward` overrides,
  outside `DFlashModule.forward`.
- **Building the rotary buffer on meta and letting the loader
materialise it**: makes RoPE
correctness depend on transformers selecting a branch by class-name
substring
(`"RotaryEmbedding" in module.__class__.__name__`), and the `if not
hasattr` guard is then
permanently satisfied, so a later `to_empty()` leaves garbage forever —
measured `4.56e-41`, i.e.
  cos=1 / sin=0, no positional encoding at all.
- **Building it eagerly in `DFlashModule.__init__`**: lands before the
dtype cast, so `Module.to`
rounds the RoPE frequencies to bf16 on the default path — measured
`0.8659643530845642` →
  `0.8671875`, loss `3.47230935097` → `3.47114777565`.

</details>

🤖 Generated with [Claude Code](https://claude.com/claude-code)




<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

- **New Features**
- DFlash now uses FP32 optimizer master weights and Adam moments by
default while keeping draft parameters in the base model’s dtype.
- Master-weight training preserves optimizer precision when restoring
checkpoints.
  - The feature can be disabled to reduce optimizer memory usage.
  - Draft models consistently follow the base model’s dtype and device.

- **Bug Fixes**
  - DFlash workflows now support operation without autocast.
- Added validation for compatible AdamW-family optimizers and
master-weight precision, including resumed training runs.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-09-23 15:39:33 +08:00
..