mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
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 `t=−0.90` | +0.0619 `t=+13.8` | | `lilicorr` | 1.2536 | 1.2814 | **1.2882** | +0.0068 `t=+1.42` | +0.0312 `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>
This commit is contained in:
@@ -33,6 +33,7 @@ from modelopt.torch.speculative.eagle.utils import (
|
||||
EagleOfflineDataCollator,
|
||||
OfflineSupervisedDataset,
|
||||
)
|
||||
from modelopt.torch.speculative.plugins.master_weight_adamw import MasterWeightAdamW
|
||||
from modelopt.torch.speculative.utils import get_ttt_msk_func
|
||||
from modelopt.torch.utils import print_rank_0
|
||||
from modelopt.torch.utils.distributed import is_master
|
||||
@@ -188,12 +189,67 @@ class EagleTrainerWithAccLog(Trainer):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.lora_lr_multiplier = lora_lr_multiplier
|
||||
|
||||
def create_optimizer(self):
|
||||
"""Override to give LoRA parameters a higher learning rate."""
|
||||
super().create_optimizer()
|
||||
def create_optimizer(self, model=None):
|
||||
"""Override to give LoRA parameters a higher learning rate.
|
||||
|
||||
``model`` mirrors the base signature. The delayed-creation branch -- FSDP1, FSDP-XLA
|
||||
and SageMaker MP -- calls this with the prepared model rather than ``self.model``,
|
||||
and an override without the parameter is a ``TypeError`` there.
|
||||
"""
|
||||
model = self.model if model is None else model
|
||||
if self.optimizer is None and getattr(model, "dflash_fp32_master_weights", False):
|
||||
# Built here rather than left to HF: the flag asks for an fp32 master copy of the
|
||||
# draft, and that lives in the optimizer. `create_optimizer` below wraps its own
|
||||
# work in `if self.optimizer is None`, so setting it here skips that entirely --
|
||||
# including the decay/no-decay grouping, which is why this reproduces it.
|
||||
# `Trainer.create_optimizer` prefers an explicitly supplied class over
|
||||
# `args.optim`, and so must this: `optimizer_cls_and_kwargs` is the supported way
|
||||
# to pass one without subclassing, and `ModelOptHFTrainer` uses it. Reading
|
||||
# `args.optim` unconditionally would discard it silently.
|
||||
if self.optimizer_cls_and_kwargs is not None:
|
||||
cls, kwargs = self.optimizer_cls_and_kwargs
|
||||
else:
|
||||
cls, kwargs = self.get_optimizer_cls_and_kwargs(self.args, model)
|
||||
if not issubclass(cls, torch.optim.AdamW):
|
||||
raise ValueError(
|
||||
f"dflash_fp32_master_weights needs an AdamW-family optimizer to hold the "
|
||||
f"master weights, but training.optim resolved to {cls.__name__}. Either "
|
||||
f"set training.optim to an adamw_torch variant, or set "
|
||||
f"dflash_fp32_master_weights=false to train the draft without an fp32 "
|
||||
f"master -- which costs acceptance length, but is the only option if the "
|
||||
f"optimizer is the point (adamw_8bit and adafactor both land here)."
|
||||
)
|
||||
# `optim` defaults to adamw_torch_fused, and the fused kernel writes the update
|
||||
# straight into the parameter it was handed -- which for us is the bf16 model
|
||||
# weight, not the fp32 master, so the master would never advance. foreach is the
|
||||
# multi-tensor path and is equivalent here.
|
||||
if kwargs.pop("fused", False):
|
||||
kwargs.setdefault("foreach", True)
|
||||
print_rank_0(
|
||||
"dflash_fp32_master_weights: using the foreach AdamW path instead of "
|
||||
"the fused one, which cannot hold master weights."
|
||||
)
|
||||
decay = self.get_decay_parameter_names(model)
|
||||
named = [(n, p) for n, p in model.named_parameters() if p.requires_grad]
|
||||
self.optimizer = MasterWeightAdamW(
|
||||
[
|
||||
{
|
||||
"params": [p for n, p in named if n in decay],
|
||||
"weight_decay": self.args.weight_decay,
|
||||
},
|
||||
{"params": [p for n, p in named if n not in decay], "weight_decay": 0.0},
|
||||
],
|
||||
**{k: v for k, v in kwargs.items() if k != "weight_decay"},
|
||||
)
|
||||
# Forwarded only when it was given: the parameter is not in every supported
|
||||
# transformers version, and the caller that passes one is that same version.
|
||||
if model is self.model:
|
||||
super().create_optimizer()
|
||||
else:
|
||||
super().create_optimizer(model)
|
||||
if self.lora_lr_multiplier != 1.0:
|
||||
lora_ids = {
|
||||
id(p) for n, p in self.model.named_parameters() if "lora_" in n and p.requires_grad
|
||||
id(p) for n, p in model.named_parameters() if "lora_" in n and p.requires_grad
|
||||
}
|
||||
if lora_ids:
|
||||
new_groups = []
|
||||
|
||||
@@ -61,6 +61,7 @@ from modelopt.torch.speculative.plugins.hf_domino import DominoLambdaCallback
|
||||
from modelopt.torch.speculative.plugins.hf_training_args import (
|
||||
TrainingArguments as SpecTrainingArgs,
|
||||
)
|
||||
from modelopt.torch.speculative.plugins.master_weight_adamw import VerifyMasterWeightsCallback
|
||||
from modelopt.torch.speculative.utils import load_vlm_or_llm, patch_transformers5_params_loading
|
||||
from modelopt.torch.utils import print_rank_0
|
||||
from modelopt.torch.utils.distributed import is_master, local_rank
|
||||
@@ -277,12 +278,11 @@ def train():
|
||||
|
||||
# On the HF-format restore path above, DFlash's modify() ran with the base model still on
|
||||
# meta, so the draft has no device, dtype or rotary buffer yet. Re-apply them here, before
|
||||
# the Trainer is built: create_optimizer freezes the Adam moment dtype off the parameters,
|
||||
# so a draft still sitting at the checkpoint's loaded dtype would silently spend the rest
|
||||
# of the run without fp32 master weights. Passing the checkpoint also restores the
|
||||
# precision `dtype="auto"` dropped on load. A no-op on a fresh convert.
|
||||
# the Trainer is built: DDP's broadcast_buffers hangs on a draft whose rotary buffer is
|
||||
# missing, and the forward only avoids reconciling dtypes because the draft matches the
|
||||
# base. A no-op on a fresh convert.
|
||||
if isinstance(model, HFDFlashModel):
|
||||
model.restore_draft_precision(checkpoint if checkpoint_is_hf else None)
|
||||
model.restore_draft_precision()
|
||||
|
||||
if dry_run:
|
||||
# is_master() is unreliable here: we return before the HF Trainer inits torch.distributed,
|
||||
@@ -322,6 +322,11 @@ def train():
|
||||
and recipe.dflash.dflash_architecture_config.get("projector_type") == "domino"
|
||||
):
|
||||
callbacks.append(DominoLambdaCallback())
|
||||
# fp32 master weights are the optimizer's job, and wiring the optimizer is the training
|
||||
# loop's. This fails the run if that wiring is ever missed, rather than letting the flag
|
||||
# be silently inert for a whole job.
|
||||
if getattr(model, "dflash_fp32_master_weights", False):
|
||||
callbacks.append(VerifyMasterWeightsCallback())
|
||||
# Leave training_args.ignore_data_skip at its default (False). The dataset is
|
||||
# map-style, so HF Trainer's resume skips consumed indices at the batch-sampler
|
||||
# level (accelerate.skip_first_batches) without re-fetching them, landing at the
|
||||
|
||||
Reference in New Issue
Block a user