mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
### What does this PR do?
Type of change: Bug fix
**EAGLE-3 context-parallel training (`--cp_size > 1`) is fixed, and its
tests run again.** CP has been broken since `accelerate` 1.13, and the
tests never caught it: the guard compared `Version("2.10.0a0")` against
`Version("2.10.0")`, which is False on every NGC alpha torch build, so
`test_llama_eagle3[cp_size=2]` has never actually run in CI.
Five fixes:
- **`main.py`** — rebuild the FSDP2 plugin accelerate requires for
`cp_size > 1`. The `--fsdp full_shard --fsdp_config` launcher flags that
used to supply it were dropped from `launch_train.sh`, so CP could not
start at all. Also pass the CP degree to the draft model.
- **`modeling_eagle.py`** — apply the draft model's first input norm
inside `layers[0]`'s own forward, where FSDP2 has actually unsharded its
weights, and only stash the input embeds on the path whose pre-hook
consumes them.
- **`hf_eagle.py`** — skip the dense eagle attention mask under CP
(causal masking comes from `is_causal`, TTT masking from the
ring-attention patch), and warn that padded positions are therefore
unmasked. Also stop `(eagle_loss or 0)` replacing a `0.0` loss tensor
with a plain `int`, which detached the graph.
- **`eagle_utils.py`** — key TTT-mask injection off the backward call's
`grad_out` kwarg, since newer torch omits `attn_bias` on the forward
call, silently disabling TTT masking.
- **`utils.py`** — CUDNN-only SDPA under CP; the `MATH` backend
decomposes SDPA and breaks on DTensors. Scoped to `cp_size > 1`, since
this context manager wraps every training forward and CPU has no cudnn
backend.
**Drops the `speculative_decoding` 26.01 container override.** It was
added when the lane ran 25.06 and spec-dec needed something *newer* — a
floor. Later bumps moved the default past it, so it had silently become
a ceiling holding spec-dec on a 6-month-old image.
### Testing
Ran `tests/examples/speculative_decoding` in
`nvcr.io/nvidia/pytorch:26.07-py3` on 2 GPUs, reproducing the CI install
steps (`pip uninstall -y nvidia-modelopt`, `pip install -e
".[hf,dev-test]"`, example requirements): **16 passed, 2 skipped** — the
2 skipped being pre-existing `--run-manual` tests. All four
`test_llama_eagle3` cases pass, including both `cp_size=2` ones.
### Before your PR is "*Ready for review*"
- Is this change backward compatible?: ✅
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: N/A
- Did you write any new necessary tests?: N/A — the existing `cp_size=2`
tests are re-enabled
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅
- Did you get Claude approval on this PR?: ❌ — not yet run
Signed-off-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com>
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
364 lines
16 KiB
Python
364 lines
16 KiB
Python
# Adapted from https://github.com/tatsu-lab/stanford_alpaca/blob/3783d18/train.py
|
|
|
|
# Copyright 2023 Rohan Taori, Ishaan Gulrajani, Tianyi Zhang, Yann Dubois, Xuechen Li
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# http://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
|
|
# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
|
# SPDX-License-Identifier: Apache-2.0
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# http://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
|
|
import argparse
|
|
import dataclasses
|
|
import os
|
|
|
|
import fsdp2_buffer_patch
|
|
import torch
|
|
import transformers
|
|
from eagle_utils import (
|
|
DFlashFSDP2ShardedSDExportCallback,
|
|
EagleTrainerWithAccLog,
|
|
EagleTrainingPlot,
|
|
LoRAWarmupCallback,
|
|
make_speculative_data_module,
|
|
patch_ring_attention_for_ttt,
|
|
)
|
|
from rich.pretty import pprint
|
|
from transformers.trainer_utils import get_last_checkpoint
|
|
|
|
import modelopt.torch.opt as mto
|
|
import modelopt.torch.speculative as mtsp
|
|
from modelopt.recipe import load_recipe
|
|
from modelopt.recipe.config import (
|
|
ModelOptDFlashRecipe,
|
|
ModelOptEagleRecipe,
|
|
ModelOptMedusaRecipe,
|
|
ModelOptSpeculativeRecipeBase,
|
|
)
|
|
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.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
|
|
|
|
torch.manual_seed(0)
|
|
mto.enable_huggingface_checkpointing()
|
|
|
|
if os.environ.get("PATCH_FSDP2_BUFFERS_TF457") == "1":
|
|
fsdp2_buffer_patch.apply()
|
|
|
|
|
|
# HF-compatible TrainingArguments with our speculative-decoding extensions, auto-derived
|
|
# from :class:`SpecTrainingArgs` so its field set can't drift from the Pydantic recipe schema.
|
|
# Used at runtime as ``HfTrainingArguments(**recipe.training.model_dump())`` to obtain a
|
|
# ``transformers.Trainer``-compatible dataclass.
|
|
HfTrainingArguments = dataclasses.make_dataclass(
|
|
"HfTrainingArguments",
|
|
[
|
|
(name, fi.annotation, dataclasses.field(default=fi.default))
|
|
for name, fi in SpecTrainingArgs.model_fields.items()
|
|
],
|
|
bases=(transformers.TrainingArguments,),
|
|
)
|
|
|
|
|
|
def _parse_cli() -> tuple[str, bool, list[str]]:
|
|
"""Parse --config (required) and --dry_run from argv; return remaining args as dotlist overrides.
|
|
|
|
Extra positional args use dotlist syntax, e.g.
|
|
``model.model_name_or_path=meta-llama/Llama-3.2-1B training.output_dir=ckpts/test``.
|
|
"""
|
|
p = argparse.ArgumentParser(add_help=False)
|
|
p.add_argument(
|
|
"--config",
|
|
required=True,
|
|
help=(
|
|
"Path to a modelopt speculative-decoding recipe YAML "
|
|
"(speculative_eagle / speculative_dflash / speculative_medusa)."
|
|
),
|
|
)
|
|
p.add_argument(
|
|
"--dry_run",
|
|
action="store_true",
|
|
help="Skip training: load base + mtsp.convert + save_pretrained, then exit. "
|
|
"Produces a ModelOpt HF checkpoint with untrained draft-head weights, suitable "
|
|
"for end-to-end plumbing tests (e.g. running scripts/export_hf_checkpoint.py).",
|
|
)
|
|
args, overrides = p.parse_known_args()
|
|
return args.config, args.dry_run, overrides
|
|
|
|
|
|
def init_distributed_env(training_args: transformers.TrainingArguments) -> None:
|
|
"""Resolve dp_shard_size from the live env and attach a ParallelismConfig in-place.
|
|
|
|
Reads ``WORLD_SIZE`` / ``torch.cuda.device_count()`` and (when actually distributed)
|
|
builds an ``accelerate.ParallelismConfig`` on ``training_args``. Kept out of the
|
|
Pydantic schema so the recipe stays a pure declarative spec.
|
|
"""
|
|
if training_args.cp_size < 1:
|
|
raise ValueError(f"cp_size must be >= 1, got {training_args.cp_size}.")
|
|
world_size = int(os.environ.get("WORLD_SIZE", torch.cuda.device_count() or 1))
|
|
if training_args.dp_shard_size is None:
|
|
training_args.dp_shard_size = world_size // training_args.cp_size
|
|
if training_args.dp_shard_size < 1:
|
|
raise ValueError(
|
|
f"dp_shard_size resolved to {training_args.dp_shard_size}; "
|
|
f"WORLD_SIZE ({world_size}) must be >= cp_size ({training_args.cp_size})."
|
|
)
|
|
|
|
if training_args.cp_size > 1 or training_args.dp_shard_size > 1:
|
|
parallel_size = training_args.dp_shard_size * training_args.cp_size
|
|
if world_size % parallel_size != 0:
|
|
raise ValueError(
|
|
f"world_size ({world_size}) must be divisible by "
|
|
f"dp_shard_size ({training_args.dp_shard_size}) * "
|
|
f"cp_size ({training_args.cp_size}) = {parallel_size}"
|
|
)
|
|
try:
|
|
from accelerate import ParallelismConfig
|
|
except ImportError as e:
|
|
raise ImportError(
|
|
"cp_size>1 or dp_shard_size>1 requires `accelerate` for ParallelismConfig. "
|
|
"Install it via `pip install accelerate`."
|
|
) from e
|
|
training_args.parallelism_config = ParallelismConfig(
|
|
cp_size=training_args.cp_size,
|
|
dp_shard_size=training_args.dp_shard_size,
|
|
dp_replicate_size=world_size // parallel_size,
|
|
)
|
|
|
|
|
|
def _is_hf_format_checkpoint(checkpoint: str | None) -> bool:
|
|
"""True if the checkpoint dir holds consolidated HF weights (from_pretrained-loadable).
|
|
|
|
FSDP2 SHARDED_STATE_DICT checkpoints contain only distributed shards
|
|
(``pytorch_model_fsdp_*/``), no ``model.safetensors`` — those return False, signalling
|
|
the caller to load the base model and resume via the Trainer instead. This inspects the
|
|
on-disk format of the *resume* checkpoint, which is a property of the existing bytes and
|
|
is independent of the current run's save mode (the two can differ across runs), so it's
|
|
intentionally separate from the save-time FSDP state-dict-type gate used for the export
|
|
callback.
|
|
"""
|
|
if not checkpoint:
|
|
return False
|
|
hf_files = ("model.safetensors", "pytorch_model.bin", "model.safetensors.index.json")
|
|
return any(os.path.isfile(os.path.join(checkpoint, f)) for f in hf_files)
|
|
|
|
|
|
def train():
|
|
config_path, dry_run, overrides = _parse_cli()
|
|
recipe = load_recipe(config_path, overrides=overrides)
|
|
if not isinstance(recipe, ModelOptSpeculativeRecipeBase):
|
|
raise ValueError(
|
|
f"main.py expects a speculative-decoding recipe (eagle / dflash / medusa); "
|
|
f"got {type(recipe).__name__} from {config_path!r}."
|
|
)
|
|
|
|
# Pydantic-typed sections flow straight through as *_args; only TrainingArguments is
|
|
# reconstructed as an HF dataclass so it can be handed to transformers.Trainer.
|
|
training_args = HfTrainingArguments(**recipe.training.model_dump())
|
|
init_distributed_env(training_args)
|
|
|
|
if not dry_run and recipe.data.mode in ("online", "streaming") and not recipe.data.data_path:
|
|
raise ValueError(f"data.mode={recipe.data.mode!r} requires data.data_path.")
|
|
if training_args.cp_size > 1:
|
|
patch_ring_attention_for_ttt()
|
|
# accelerate requires an fsdp_plugin when cp_size > 1; the --fsdp launcher flags that
|
|
# used to provide one were dropped from launch_train.sh.
|
|
if not training_args.fsdp_plugin_args:
|
|
training_args.fsdp = "full_shard"
|
|
training_args.fsdp_config = {"fsdp_version": 2}
|
|
training_args.fsdp_plugin_args = training_args._process_fsdp_args()
|
|
if is_master():
|
|
pprint(recipe)
|
|
|
|
# Detect checkpoint to resume from
|
|
last_checkpoint = (
|
|
get_last_checkpoint(training_args.output_dir)
|
|
if os.path.isdir(training_args.output_dir)
|
|
else None
|
|
)
|
|
if last_checkpoint:
|
|
print_rank_0(f"Last checkpoint detected: {last_checkpoint}")
|
|
|
|
checkpoint = training_args.resume_from_checkpoint or last_checkpoint
|
|
|
|
use_offline_training = recipe.data.mode != "online"
|
|
|
|
# Resume path depends on the existing checkpoint's on-disk format: consolidated HF
|
|
# weights load via from_pretrained; FSDP sharded checkpoints load the base model and
|
|
# resume through the Trainer.
|
|
checkpoint_is_hf = _is_hf_format_checkpoint(checkpoint)
|
|
|
|
if checkpoint_is_hf:
|
|
assert checkpoint is not None # guaranteed by checkpoint_is_hf
|
|
with patch_transformers5_params_loading():
|
|
model = load_vlm_or_llm(
|
|
checkpoint, dtype="auto", trust_remote_code=recipe.model.trust_remote_code
|
|
)
|
|
tokenizer = transformers.AutoTokenizer.from_pretrained(
|
|
checkpoint, trust_remote_code=recipe.model.trust_remote_code
|
|
)
|
|
else:
|
|
if checkpoint:
|
|
print_rank_0(
|
|
f"Checkpoint {checkpoint} is not in HF format (FSDP distributed checkpoint). "
|
|
f"Loading base model and resuming via Trainer."
|
|
)
|
|
model_name_or_path = recipe.model.model_name_or_path
|
|
if model_name_or_path is None:
|
|
raise ValueError(
|
|
"model.model_name_or_path must be set in the recipe YAML or via a dotlist override."
|
|
)
|
|
# To avoid OOM for large models, we load and convert model on CPU first.
|
|
# Model will be moved to GPU during HF trainer.init().
|
|
model = load_vlm_or_llm(
|
|
model_name_or_path,
|
|
use_fake_base=recipe.model.use_fake_base_for_offline,
|
|
use_offline_training=use_offline_training,
|
|
dtype="auto",
|
|
device_map="cpu",
|
|
trust_remote_code=recipe.model.trust_remote_code,
|
|
)
|
|
tokenizer = transformers.AutoTokenizer.from_pretrained(
|
|
model_name_or_path,
|
|
model_max_length=training_args.training_seq_len,
|
|
trust_remote_code=recipe.model.trust_remote_code,
|
|
)
|
|
if isinstance(recipe, ModelOptMedusaRecipe):
|
|
medusa_cfg: dict = recipe.medusa.model_dump()
|
|
mtsp.convert(model, [("medusa", medusa_cfg)])
|
|
elif isinstance(recipe, ModelOptEagleRecipe):
|
|
eagle_cfg: dict = recipe.eagle.model_dump()
|
|
mtsp.convert(model, [("eagle", eagle_cfg)])
|
|
# Load draft vocab cache
|
|
mtsp.plugins.HFEagleModel.load_draft_vocab_cache(model, recipe.data.draft_vocab_cache)
|
|
elif isinstance(recipe, ModelOptDFlashRecipe):
|
|
# Fall back to tokenizer.mask_token_id when not set in the recipe; require one of the two.
|
|
if recipe.dflash.dflash_mask_token_id is None:
|
|
recipe.dflash.dflash_mask_token_id = getattr(tokenizer, "mask_token_id", None)
|
|
if recipe.dflash.dflash_mask_token_id is None:
|
|
raise ValueError(
|
|
"dflash.dflash_mask_token_id is required: set it in the recipe YAML "
|
|
"or use a tokenizer that defines mask_token_id."
|
|
)
|
|
dflash_cfg: dict = recipe.dflash.model_dump()
|
|
mtsp.convert(model, [("dflash", dflash_cfg)])
|
|
else:
|
|
raise ValueError(f"Unsupported speculative recipe type: {type(recipe).__name__}")
|
|
|
|
if dry_run:
|
|
# is_master() is unreliable here: we return before the HF Trainer inits torch.distributed,
|
|
# so use local_rank() (env-based) to keep a single writer to output_dir.
|
|
if local_rank() == 0:
|
|
os.makedirs(training_args.output_dir, exist_ok=True)
|
|
model.save_pretrained(training_args.output_dir)
|
|
tokenizer.save_pretrained(training_args.output_dir)
|
|
print_rank_0(
|
|
f"[dry-run] saved ModelOpt HF checkpoint (untrained draft head) to "
|
|
f"{training_args.output_dir}"
|
|
)
|
|
return
|
|
|
|
print_rank_0("Loading dataset...")
|
|
is_dflash = isinstance(recipe, ModelOptDFlashRecipe)
|
|
data_module = make_speculative_data_module(
|
|
tokenizer,
|
|
recipe.data,
|
|
train_len=training_args.training_seq_len,
|
|
answer_only_loss=training_args.answer_only_loss,
|
|
shift_labels=not is_dflash,
|
|
)
|
|
|
|
callbacks = [EagleTrainingPlot(training_args.ar_validate_steps, training_args.estimate_ar)]
|
|
if (
|
|
isinstance(recipe, ModelOptEagleRecipe)
|
|
and recipe.eagle.eagle_base_lora
|
|
and recipe.eagle.eagle_base_lora_warmup_steps > 0
|
|
):
|
|
callbacks.append(LoRAWarmupCallback(recipe.eagle.eagle_base_lora_warmup_steps))
|
|
# Domino (dflash recipe with projector_type=domino) needs the lambda_base
|
|
# curriculum schedule driven by the trainer's global step.
|
|
if (
|
|
isinstance(recipe, ModelOptDFlashRecipe)
|
|
and recipe.dflash.dflash_architecture_config.get("projector_type") == "domino"
|
|
):
|
|
callbacks.append(DominoLambdaCallback())
|
|
# 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
|
|
# exact data position. Setting it True would restart the data order from the top.
|
|
|
|
# Tell the draft model the CP degree so it skips the dense eagle mask under CP.
|
|
model.eagle_cp_size = training_args.cp_size
|
|
|
|
trainer = EagleTrainerWithAccLog(
|
|
model=model,
|
|
processing_class=tokenizer,
|
|
args=training_args,
|
|
callbacks=callbacks,
|
|
**data_module,
|
|
)
|
|
|
|
if os.environ.get("PATCH_FSDP2_BUFFERS_TF457") == "1":
|
|
fsdp2_buffer_patch.patch_accelerator(trainer.accelerator)
|
|
|
|
# DFlash: export the draft submodule after each checkpoint save — but only under FSDP2
|
|
# SHARDED_STATE_DICT, where checkpoints are distributed shards the post-training
|
|
# export_hf_checkpoint.py pass can't read. Gate by reading the live FSDP state dict
|
|
# type off the accelerator; full-state-dict runs (DDP, single-device, FSDP2
|
|
# FULL_STATE_DICT) use the launcher's post-run export instead.
|
|
if isinstance(recipe, ModelOptDFlashRecipe):
|
|
fsdp_plugin = getattr(trainer.accelerator.state, "fsdp_plugin", None)
|
|
sd_type = str(getattr(fsdp_plugin, "state_dict_type", "") or "")
|
|
if "SHARDED_STATE_DICT" in sd_type:
|
|
trainer.add_callback(DFlashFSDP2ShardedSDExportCallback())
|
|
print_rank_0("DFlash: FSDP2 SHARDED_STATE_DICT — enabling per-save draft export.")
|
|
else:
|
|
print_rank_0(
|
|
f"DFlash: checkpoints use {sd_type or 'a full state dict'}; relying on the "
|
|
"launcher's post-run export (no per-save export callback added)."
|
|
)
|
|
|
|
# Manually enable this to return loss in eval
|
|
trainer.can_return_loss = True
|
|
# Make sure label_smoother is None
|
|
assert trainer.label_smoother is None, (
|
|
"label_smoother is not supported in speculative decoding!"
|
|
)
|
|
|
|
# Diagnostic (no-op unless DFLASH_LOG_PARAM_DTYPES=1): verifies FSDP2 dtype sync.
|
|
fsdp2_buffer_patch.log_param_dtypes(trainer.model)
|
|
|
|
print_rank_0("Start training...")
|
|
trainer.train(resume_from_checkpoint=checkpoint)
|
|
trainer.save_state()
|
|
trainer.save_model(training_args.output_dir)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
train()
|