[1/3][Refactor]: File reorg; deprecate ParallelDraft (#1296)

### What does this PR do?

Type of change: refactoring

Part 1 of a 3-PR series splitting #1271:
- **[1/3] this PR**: File reorg + deprecate `ParallelDraft`
- **[2/3] #1295**: Offline DFlash training
- **[3/3] #1297**: Extract `HFSpecDecMixin`

Changes:
- **File reorg**: `transformers.py` → `hf_eagle.py`; extract
`HFMedusaModel` → `hf_medusa.py`; extract `EagleModule` /
`EagleBaseModelOutput` → `modeling_eagle.py`; extract `DFlashModule` /
`DFlashAttention` / `DFlashDecoderLayer` / `build_target_layer_ids` /
`apply_rotary_pos_emb` → `modeling_dflash.py`.
- **Deprecate `ParallelDraft`**: remove `parallel_draft_step`,
`parallel_draft_heads_num_layers`, and the `ParallelDraft` module from
HF Eagle; remove the `EagleMedusaExporter` branch from
`HFEagleModel.get_exporter()` (the `EagleMedusaExporter` class itself
still lives in `hf_spec_export.py` for Megatron parity).
- **Rename**: `_draft_model_config` → `eagle_config` in export plugin.
- Update imports in `examples/speculative_decoding/` and
`modelopt/torch/speculative/utils.py` to follow the module rename.

### Testing

Validated with existing Eagle and DFlash training scripts (re-run after
`9ae5302729 revert behavior change`).

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

Make sure you read and follow [Contributor
guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)
and your commits are signed (`git commit -s -S`).

Make sure you read and follow the [Security Best
Practices](https://github.com/NVIDIA/Model-Optimizer/blob/main/SECURITY.md#security-coding-practices-for-contributors)
(e.g. avoiding hardcoded `trust_remote_code=True`, `torch.load(...,
weights_only=False)`, `pickle`, etc.).

- Is this change backward compatible?: ❌ — renames
`modelopt.torch.speculative.plugins.transformers` → `.hf_eagle`; removes
`parallel_draft_step` / `parallel_draft_heads_num_layers` from Eagle
config; renames `_draft_model_config` → `eagle_config` in export plugin.
- 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 — pure refactor; existing
tests updated for the rename. `test_hf_spec_rope_export.py` assertions
were also corrected to reflect the actual production path (the old
assertions were masked by `MagicMock` not invoking the
`_draft_model_config` `@property`).
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
❌

### Additional Information

Breaking changes:
- `modelopt.torch.speculative.plugins.transformers` → `.hf_eagle`
- `parallel_draft_step` / `parallel_draft_heads_num_layers` removed from
Eagle config
- `_draft_model_config` → `eagle_config` in export plugin

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

* **Refactoring**
* Reorganized speculative-decoding plugins into focused modules,
converting the legacy "transformers" entry into a deprecated shim that
re-exports the new plugin surface.
* Consolidated DFlash implementation into a shared modeling component
and introduced a dedicated EAGLE decoder module.

* **New Features**
* Added a Medusa speculative-decoding plugin with configurable heads and
combined-loss training behavior.

* **Chores**
  * Updated pre-commit license-hook exclusion and feature-flag wiring.

* **Tests**
  * Updated export tests to expect rope-scaling fallback semantics.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com>
This commit is contained in:
h-guo18
2026-04-24 14:46:10 -07:00
committed by GitHub
parent 946639aa19
commit 7c80d85751
14 changed files with 1511 additions and 1464 deletions
+1 -1
View File
@@ -99,7 +99,7 @@ repos:
modelopt/torch/quantization/plugins/attention.py|
modelopt/torch/sparsity/attention_sparsity/methods/vsa_utils.py|
modelopt/torch/speculative/eagle/utils.py|
modelopt/torch/speculative/plugins/transformers.py|
modelopt/torch/speculative/plugins/hf_medusa.py|
modelopt/torch/utils/plugins/megatron_mmlu.py|
examples/chained_optimizations/bert_prune_distill_quantize.py|
examples/deepseek/quantize_to_nvfp4.py|
+1 -1
View File
@@ -358,7 +358,7 @@ def get_patched_templated_ring_attn(orig_templated_attn: Callable):
original_op = args[2]
# This patch is only enabled for eagle model by context manager, not base model.
patch_enbabled = modelopt.torch.speculative.plugins.transformers.ENABLE_CP_TTT_PATCH
patch_enbabled = modelopt.torch.speculative.plugins.hf_eagle.ENABLE_CP_TTT_PATCH
if patch_enbabled and original_op != torch.ops.aten._scaled_dot_product_cudnn_attention:
raise ValueError(f"CP TTT only supports cudnn attention now. Got: {original_op}")
@@ -27,7 +27,7 @@ from tqdm import tqdm
from transformers import AutoTokenizer
import modelopt.torch.opt as mto
from modelopt.torch.speculative.plugins.transformers import HFARValidation
from modelopt.torch.speculative.plugins.hf_eagle import HFARValidation
from modelopt.torch.speculative.utils import load_vlm_or_llm
mto.enable_huggingface_checkpointing()
@@ -171,8 +171,8 @@ class EagleExporter(SpeculativeDecodingExporter):
template_config = deepcopy(template_config)
def _get_config_from_draft_or_base(key: str, model: nn.Module):
if getattr(model._draft_model_config, key, None) is not None:
return getattr(model._draft_model_config, key)
if getattr(model.eagle_config, key, None) is not None:
return getattr(model.eagle_config, key)
elif getattr(model.config, key, None) is not None:
return getattr(model.config, key)
else:
@@ -37,6 +37,7 @@ default_eagle_config = {
"use_aux_hidden_state": False,
"eagle_aux_hidden_state_layer_ids": [],
"use_mtp_layernorm": False,
# Deprecated on the HF flow; TODO: remove once the Megatron flow stops reading these.
"parallel_draft_step": 1,
"parallel_draft_heads_num_layers": 1,
"has_lm_head": False,
@@ -107,6 +108,7 @@ default_kimik2_eagle_config = {
"use_aux_hidden_state": True,
"eagle_aux_hidden_state_layer_ids": [],
"use_mtp_layernorm": False,
# Deprecated on the HF flow; TODO: remove once the Megatron flow stops reading these.
"parallel_draft_step": 1,
"parallel_draft_heads_num_layers": 1,
"has_lm_head": False,
@@ -18,7 +18,7 @@
Please check out the source code of this module for examples of how plugins work and how you can
write your own one. Currently, we support plugins for
- :meth:`transformers<modelopt.torch.speculative.plugins.transformers>`
- :meth:`hf_eagle<modelopt.torch.speculative.plugins.hf_eagle>`
"""
from modelopt.torch.utils import import_plugin
@@ -31,4 +31,5 @@ with import_plugin("megatron_medusa"):
with import_plugin("transformers"):
from .hf_dflash import *
from .transformers import *
from .hf_eagle import *
from .hf_medusa import *
+1 -214
View File
@@ -54,21 +54,14 @@ import logging
import torch
import torch.nn.functional as F
from torch import nn
from transformers import PreTrainedModel
from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS
from transformers.models.qwen3.configuration_qwen3 import Qwen3Config as _Qwen3Config
from transformers.models.qwen3.modeling_qwen3 import Qwen3MLP as _MLP_CLS # noqa: N814
from transformers.models.qwen3.modeling_qwen3 import Qwen3RMSNorm as _NORM_CLS # noqa: N814
from transformers.models.qwen3.modeling_qwen3 import (
Qwen3RotaryEmbedding as _ROTARY_CLS, # noqa: N814
)
from transformers.models.qwen3.modeling_qwen3 import rotate_half as _rotate_half
from transformers.trainer_pt_utils import LabelSmoother
from transformers.utils import ModelOutput
from ..dflash.conversion import DFlashDMRegistry
from ..dflash.dflash_model import DFlashModel
from .modeling_dflash import DFlashAttention, DFlashModule, build_target_layer_ids # noqa: F401
from .modeling_fakebase import _BASE_MODEL_PATHS, _EMBED_TOKENS_PATHS, _LM_HEAD_PATHS
logger = logging.getLogger(__name__)
@@ -76,212 +69,6 @@ logger = logging.getLogger(__name__)
__all__ = ["HFDFlashModel"]
def build_target_layer_ids(num_target_layers, num_draft_layers):
"""Select layers uniformly from the target model for feature extraction."""
if num_target_layers < num_draft_layers:
raise ValueError(
f"num_target_layers ({num_target_layers}) must be >= num_draft_layers ({num_draft_layers})"
)
if num_draft_layers == 1:
return [num_target_layers // 2]
start = min(1, num_target_layers - 1)
end = max(start, num_target_layers - 3)
span = end - start
return [round(start + (i * span) / (num_draft_layers - 1)) for i in range(num_draft_layers)]
def apply_rotary_pos_emb(q, k, cos, sin):
"""Apply RoPE. Q uses last q_len positions, K uses all positions."""
cos = cos.unsqueeze(1) # [B, 1, seq, dim]
sin = sin.unsqueeze(1)
q_len = q.size(2)
q_embed = (q * cos[:, :, -q_len:, :]) + (_rotate_half(q) * sin[:, :, -q_len:, :])
k_embed = (k * cos) + (_rotate_half(k) * sin)
return q_embed, k_embed
class DFlashAttention(nn.Module):
"""Attention with KV injection, using HF's attention dispatch."""
def __init__(self, config, layer_idx):
"""Initialize DFlash attention with KV injection projections and QK-norm."""
super().__init__()
self.config = config
self.layer_idx = layer_idx
self.head_dim = getattr(
config, "head_dim", config.hidden_size // config.num_attention_heads
)
self.num_heads = config.num_attention_heads
self.num_kv_heads = config.num_key_value_heads
self.num_key_value_groups = self.num_heads // self.num_kv_heads
self.scaling = self.head_dim**-0.5
self.attention_dropout = getattr(config, "attention_dropout", 0.0)
self.is_causal = False
attn_bias = getattr(config, "attention_bias", False)
self.q_proj = nn.Linear(config.hidden_size, self.num_heads * self.head_dim, bias=attn_bias)
self.k_proj = nn.Linear(
config.hidden_size, self.num_kv_heads * self.head_dim, bias=attn_bias
)
self.v_proj = nn.Linear(
config.hidden_size, self.num_kv_heads * self.head_dim, bias=attn_bias
)
self.o_proj = nn.Linear(self.num_heads * self.head_dim, config.hidden_size, bias=attn_bias)
self.q_norm = _NORM_CLS(self.head_dim, eps=config.rms_norm_eps)
self.k_norm = _NORM_CLS(self.head_dim, eps=config.rms_norm_eps)
# Resolve HF attention function
self._attn_fn = None
# Qwen3 uses sliding window attention on some layers (config.layer_types)
if hasattr(config, "layer_types") and hasattr(config, "sliding_window"):
is_sliding = config.layer_types[layer_idx] == "sliding_attention"
self.sliding_window = config.sliding_window if is_sliding else None
else:
self.sliding_window = None
def _get_attn_fn(self):
"""Lazily resolve the HF attention function (default: sdpa)."""
if self._attn_fn is not None:
return self._attn_fn
impl = self.config._attn_implementation # default set in dflash/default_config.py
self._attn_fn = ALL_ATTENTION_FUNCTIONS.get(impl, ALL_ATTENTION_FUNCTIONS["sdpa"])
return self._attn_fn
def forward(self, hidden_states, target_hidden, position_embeddings, attention_mask=None):
"""Forward with KV injection.
Q is projected from the noise block (draft token embeddings: [anchor, mask, mask, ...]).
K and V are projected from the concatenation of target hidden states (context from the
base model) and noise block, so the draft can attend to both context and its own block.
"""
bsz, q_len, _ = hidden_states.shape
ctx_len = target_hidden.shape[1]
# Q from noise block only (the draft tokens being predicted), with QK-norm
q = self.q_proj(hidden_states).view(bsz, q_len, -1, self.head_dim)
q = self.q_norm(q).transpose(1, 2)
# K from context + noise, with QK-norm
k_ctx = self.k_proj(target_hidden)
k_noise = self.k_proj(hidden_states)
k = torch.cat([k_ctx, k_noise], dim=1).view(bsz, ctx_len + q_len, -1, self.head_dim)
k = self.k_norm(k).transpose(1, 2)
# V from context + noise (no norm)
v_ctx = self.v_proj(target_hidden)
v_noise = self.v_proj(hidden_states)
v = (
torch.cat([v_ctx, v_noise], dim=1)
.view(bsz, ctx_len + q_len, -1, self.head_dim)
.transpose(1, 2)
)
# RoPE
cos, sin = position_embeddings
q, k = apply_rotary_pos_emb(q, k, cos, sin)
# Use HF's attention dispatch (handles GQA internally)
attn_fn = self._get_attn_fn()
attn_output, _ = attn_fn(
self,
q,
k,
v,
attention_mask,
dropout=0.0 if not self.training else self.attention_dropout,
scaling=self.scaling,
sliding_window=self.sliding_window,
)
attn_output = attn_output.reshape(bsz, q_len, -1)
return self.o_proj(attn_output)
class DFlashDecoderLayer(nn.Module):
"""Draft decoder layer with KV injection."""
def __init__(self, config, layer_idx):
"""Initialize decoder layer with attention, MLP, and layer norms."""
super().__init__()
self.self_attn = DFlashAttention(config, layer_idx)
self.mlp = _MLP_CLS(config)
self.input_layernorm = _NORM_CLS(config.hidden_size, eps=config.rms_norm_eps)
self.post_attention_layernorm = _NORM_CLS(config.hidden_size, eps=config.rms_norm_eps)
def forward(self, hidden_states, target_hidden, position_embeddings, attention_mask=None):
"""Forward pass with residual connections."""
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
hidden_states = self.self_attn(
hidden_states, target_hidden, position_embeddings, attention_mask
)
hidden_states = residual + hidden_states
residual = hidden_states
hidden_states = self.post_attention_layernorm(hidden_states)
hidden_states = self.mlp(hidden_states)
hidden_states = residual + hidden_states
return hidden_states
class DFlashModule(nn.Module):
"""DFlash draft module using Qwen3 components (MLP, RMSNorm, RotaryEmbedding)."""
def __init__(self, config):
"""Initialize DFlash module with feature fusion, decoder layers, and rotary embeddings."""
super().__init__()
self.config = config
self.block_size = config.block_size
# Feature fusion
num_fused_layers = len(config.target_layer_ids)
self.fc = nn.Linear(num_fused_layers * config.hidden_size, config.hidden_size, bias=False)
self.hidden_norm = _NORM_CLS(config.hidden_size, eps=config.rms_norm_eps)
# Decoder layers
self.layers = nn.ModuleList(
[DFlashDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
)
self.norm = _NORM_CLS(config.hidden_size, eps=config.rms_norm_eps)
self._rotary_config = config # Used by _maybe_init_rotary_emb
# Explicit weight init is needed because DFlashModule is instantiated via
# mtsp.convert() AFTER the base model's post_init() has already run, so HF's
# automatic _init_weights walk doesn't reach these new layers.
self._init_weights(config)
def _maybe_init_rotary_emb(self, device=None):
"""Lazily initialize rotary embeddings on first forward call.
Same pattern as EAGLE3's _maybe_init_rope. Avoids creating rotary_emb
during __init__ (which runs on meta device during from_pretrained),
preventing the meta-tensor inv_freq issue on checkpoint resume.
"""
if not hasattr(self, "rotary_emb"):
self.rotary_emb = _ROTARY_CLS(config=self._rotary_config, device=device)
def _init_weights(self, config):
"""Initialize weights matching HF PreTrainedModel._init_weights."""
std = getattr(config, "initializer_range", 0.02)
for module in self.modules():
if isinstance(module, nn.Linear):
nn.init.normal_(module.weight, mean=0.0, std=std)
if module.bias is not None:
nn.init.zeros_(module.bias)
def forward(self, noise_embedding, target_hidden, position_ids, attention_mask=None):
"""Forward with feature fusion, KV injection, and position embeddings."""
hidden_states = noise_embedding
target_hidden = self.hidden_norm(self.fc(target_hidden))
self._maybe_init_rotary_emb(device=hidden_states.device)
position_embeddings = self.rotary_emb(hidden_states, position_ids)
for layer in self.layers:
hidden_states = layer(hidden_states, target_hidden, position_embeddings, attention_mask)
return self.norm(hidden_states)
@DFlashDMRegistry.register({PreTrainedModel: "hf.PreTrainedModel"})
class HFDFlashModel(DFlashModel):
"""DFlash Model for HuggingFace transformers."""
@@ -0,0 +1,855 @@
# 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.
"""Support speculative decoding for huggingface models."""
import contextlib
import copy
from typing import Any
import torch
from torch.nn.attention.flex_attention import BlockMask, create_block_mask
from transformers import Cache, DynamicCache, PreTrainedModel
from transformers.models.llama.modeling_llama import LlamaDecoderLayer
from transformers.utils import ModelOutput
from ...export.plugins.hf_spec_export import EagleExporter, SpeculativeDecodingExporter
from ..eagle.conversion import EagleDMRegistry
from ..eagle.eagle_model import EagleModel
from ..eagle.utils import expand_mask, make_causal_mask
from ..utils import (
AcceptanceRateValidation,
_setup_kimi_k2_decoder,
enable_cp_ttt_patch,
get_ttt_msk_func,
temporary_set_config_value,
)
from .modeling_eagle import EagleBaseModelOutput, EagleModule
from .modeling_fakebase import _BASE_MODEL_PATHS, _EMBED_TOKENS_PATHS, _LM_HEAD_PATHS
__all__ = ["HFARValidation", "HFEagleModel"]
ENABLE_CP_TTT_PATCH = False
# module variable to cache attention mask for cp ttt
CACHED_SHARD_TTT_MASKS = {}
@EagleDMRegistry.register({PreTrainedModel: "hf.PreTrainedModel"})
class HFEagleModel(EagleModel):
"""Eagle Model Class for huggingface models."""
@property
def _base_model(self):
return self.get_submodule(self.base_model_path)
@property
def _base_model_embeddings(self):
return self.get_submodule(self.base_model_embeddings_path)
@property
def _base_model_lm_head(self):
return self.get_submodule(self.base_model_lm_head_path)
@property
def _base_llm_config(self):
"""Return the llm config for the base model, from LLM or VLM."""
return (
getattr(self.config, "text_config", None)
or getattr(self.config, "llm_config", None)
or self.config
)
def _nvtx_range(self, name):
"""Optionally create an NVTX range for the given name when config.eagle_enable_nvtx is set."""
if not self.eagle_enable_nvtx:
return contextlib.nullcontext()
try:
import torch.cuda.nvtx as nvtx
return nvtx.range(name)
except Exception as e:
print(f"Failed to create NVTX range {name}: {e}")
return contextlib.nullcontext()
def _find_base_model_parts(self):
"""Find model parts from different models and set base_{part}_path attributes."""
base_model_parts_mapping = {
"base_model_path": _BASE_MODEL_PATHS,
"base_model_embeddings_path": _EMBED_TOKENS_PATHS,
"base_model_lm_head_path": _LM_HEAD_PATHS,
}
for name, paths in base_model_parts_mapping.items():
found_submodule = False
for path in paths:
try:
submodule = self.get_submodule(path)
assert isinstance(submodule, torch.nn.Module)
print(f"Found {name} at {path}")
found_submodule = True
setattr(self, name, path)
break
except Exception:
continue
if not found_submodule:
raise ValueError(f"Part {name} not found in model")
def _activate_torch_compile(self):
import torch._dynamo
torch._dynamo.config.suppress_errors = True # Allow fallback to eager mode
compile_targets = [
("_prepare_eagle_inputs", {}),
("_eagle_forward", {"mode": "max-autotune"}),
("_eagle_loss", {"fullgraph": True}),
]
for name, kwargs in compile_targets:
try:
setattr(self, name, torch.compile(getattr(self, name), dynamic=False, **kwargs))
except Exception: # noqa: PERF203
print(f"Disabling torch.compile for {name} due to compilation error.")
def get_dummy_inputs(self) -> dict:
"""Construct dummy inputs for export forward pass."""
device = self.device
dummy_inputs = {
"input_ids": torch.ones(1, 2, dtype=torch.long, device=device),
}
if self.eagle_offline:
device = self.device
dtype = next(self.parameters()).dtype
hidden_size = self._base_llm_config.hidden_size
base_model_outputs = {
"base_model_hidden_states": torch.zeros(
1, 2, hidden_size, dtype=dtype, device=device
),
"base_model_input_embeds": torch.zeros(
1, 2, hidden_size, dtype=dtype, device=device
),
}
if self.eagle_config.use_aux_hidden_state:
num_aux = len(self.eagle_config.eagle_aux_hidden_state_layer_ids)
base_model_outputs["aux_hidden_states"] = torch.zeros(
1, 2, hidden_size * num_aux, dtype=dtype, device=device
)
dummy_inputs["base_model_outputs"] = base_model_outputs
return dummy_inputs
def get_exporter(self) -> SpeculativeDecodingExporter:
"""Get the exporter for the draft model."""
return EagleExporter(self)
def _enable_cp_ttt(self):
if self.training and not self.eagle_mix_hidden_states:
return enable_cp_ttt_patch()
return contextlib.nullcontext()
def _set_default_aux_hidden_state_layers(self):
# Read a custom config attribute since we override num_hidden_layers for offline training
num_layers = self._base_llm_config.num_hidden_layers
if self.eagle_offline and (num_layers is None or num_layers <= 0):
num_layers = getattr(self.config, "num_orig_hidden_layers", 0)
self.eagle_config.eagle_aux_hidden_state_layer_ids = [
1,
max(0, num_layers // 2 - 1),
max(0, num_layers - 4),
]
self.eagle_config.eagle_aux_hidden_state_layer_ids = list(
set(self.eagle_config.eagle_aux_hidden_state_layer_ids)
)
def _collect_aux_hidden_states_forward_hook(self, module, input, output) -> None:
"""Collect auxiliary hidden states from base model intermediate layers, save them in attribute."""
raw = output if isinstance(output, torch.Tensor) else output[0]
# With LoRA co-training (after warmup), keep grad so EAGLE loss
# backpropagates through hidden states to LoRA.
if self.training and getattr(self, "_lora_cotraining_active", False):
self._aux_hidden_states.append(raw.clone())
else:
self._aux_hidden_states.append(raw.clone().detach())
def pop_and_gather_aux_hiddens(self):
"""Pop auxiliary hidden states from base model and gather them on the draft model device."""
if not self.eagle_config.use_aux_hidden_state:
return None
# In PTQ, forward method will be called with try and except to find max batch size.
# This leads to uncleared aux hidden states in the front of the list.
# To fix it, we only return the last num_aux_h items in the list.
num_aux_h = len(self.eagle_config.eagle_aux_hidden_state_layer_ids)
aux_h_list = self._aux_hidden_states[-num_aux_h:]
self._aux_hidden_states.clear()
# Gather aux hidden states on the draft model device
aux_hiddens = torch.cat(
[h.to(self.eagle_module.fc.weight.device) for h in aux_h_list], dim=-1
)
return aux_hiddens
def _get_eagle_device(self):
"""Return the device where we should place eagle module."""
if self.eagle_offline:
# For offline training, the base model has no layers.
# Read the device from the base model lm_head instead.
return self._base_model_lm_head.weight.device
else:
# When there is a base model, put eagle on the last layer's device.
base_model_last_layer = self._base_model.layers[-1]
return next(base_model_last_layer.parameters()).device
def _inject_base_lora(self):
"""Inject HF PEFT LoRA adapters into the base model in-place and unfreeze them."""
from peft import LoraConfig
from peft.mapping import inject_adapter_in_model
target_modules = self.eagle_base_lora_target_modules or None
lora_config = LoraConfig(
r=self.eagle_base_lora_rank,
lora_alpha=self.eagle_base_lora_alpha,
target_modules=target_modules,
bias="none",
)
inject_adapter_in_model(lora_config, self._base_model, adapter_name="default")
# Unfreeze LoRA parameters unless we have a warmup phase
freeze_lora = self.eagle_base_lora_warmup_steps > 0
for name, param in self._base_model.named_parameters():
if "lora_" in name:
param.requires_grad = not freeze_lora
def _set_base_lora_enabled(self, enabled: bool) -> None:
"""Enable or disable LoRA adapters in the base model."""
from peft.tuners.lora import LoraLayer
for module in self._base_model.modules():
if isinstance(module, LoraLayer):
module.enable_adapters(enabled)
def _preservation_loss(
self, ref_logits: torch.Tensor, lora_logits: torch.Tensor
) -> torch.Tensor:
"""KL divergence encouraging LoRA output to stay close to the original base model.
KL(softmax(ref) || log_softmax(lora)) weighted by eagle_base_lora_preservation_loss_weight.
"""
loss = torch.nn.Softmax(dim=-1)(ref_logits.detach()) * torch.nn.LogSoftmax(dim=-1)(
lora_logits
)
return -loss.sum(dim=-1).mean() * self.eagle_base_lora_preservation_loss_weight
def modify(
self,
config,
):
"""Constructor.
Args:
config: The config for eagle decoder layers.
"""
super().modify(config)
if self.eagle_decoder_type == "llama":
# Use default eagle config
decoder_cls = LlamaDecoderLayer
elif self.eagle_decoder_type == "kimik2":
decoder_cls = _setup_kimi_k2_decoder()
arch_config = config.eagle_architecture_config
# Populate base-model-dependent fields before constructing PretrainedConfig,
# since transformers >=5.4 validates rope_scaling during __init__.
arch_config["hidden_size"] = self._base_llm_config.hidden_size
arch_config["vocab_size"] = self._base_llm_config.vocab_size
arch_config["max_position_embeddings"] = self._base_llm_config.max_position_embeddings
rope_scaling = arch_config.get("rope_scaling")
if rope_scaling and "rope_theta" not in rope_scaling and "rope_theta" in arch_config:
rope_scaling["rope_theta"] = arch_config["rope_theta"]
# Use the base model's config class so fields like max_position_embeddings are declared
# before transformers>=5.5 rope standardization runs in __post_init__.
base_config_cls = type(self._base_llm_config)
self.eagle_config = base_config_cls.from_dict(arch_config)
self.eagle_config.eagle_decoder_type = self.eagle_decoder_type
self.eagle_config.draft_vocab_size = getattr(
self.eagle_config, "draft_vocab_size", self.eagle_config.vocab_size
)
if self.eagle_config._attn_implementation is None:
self.eagle_config._attn_implementation = "sdpa"
# Set default aux_hidden_state layers
if (
self.eagle_config.use_aux_hidden_state
and len(self.eagle_config.eagle_aux_hidden_state_layer_ids) == 0
):
self._set_default_aux_hidden_state_layers()
# Freeze all parameters
if self.eagle_freeze_base_model:
for _, param in self.named_parameters():
param.requires_grad = False
self.eagle_module = EagleModule(
self.eagle_config,
decoder_cls,
)
# find base model, lm head, and embeddings paths
self._find_base_model_parts()
self.eagle_module.to(self._base_model.dtype).to(self._get_eagle_device())
# EAGLE-3 auxiliary hidden_states
if (not self.eagle_offline) and self.eagle_config.use_aux_hidden_state:
self._aux_hidden_states = []
for layer_idx, layer in enumerate(self._base_model.layers):
if layer_idx in self.eagle_config.eagle_aux_hidden_state_layer_ids:
layer.register_forward_hook(self._collect_aux_hidden_states_forward_hook)
# Inject HF PEFT LoRA adapters into the base model for co-training
if self.eagle_base_lora:
if self.eagle_offline:
raise ValueError("eagle_base_lora is incompatible with eagle_offline=True")
self._inject_base_lora()
# Whether LoRA co-training is active this step. Controlled by the
# trainer based on warmup schedule. When False, LoRA params are
# frozen and logits are always detached (eagle-only training).
self._lora_cotraining_active = self.eagle_base_lora_warmup_steps == 0
# delete base model layers for offline training
if self.eagle_offline:
self._base_model._modules.pop("layers")
# NOTE: this is a temporary hack to bypass hf trainer check:
# https://github.com/huggingface/transformers/blob/v4.56-release/src/transformers/trainer.py#L566
self.is_quantized = False
if self.eagle_use_torch_compile:
self._activate_torch_compile()
self._cached_attn_blk_masks = {}
def _get_ttt_attention_mask(self, batch_size, seq_length, ttt_step):
# compile and cached flex attention masks in first call
if ttt_step not in self._cached_attn_blk_masks:
self._cached_attn_blk_masks.update(
{ttt_step: self._compute_ttt_attention_mask(batch_size, seq_length, ttt_step)}
)
return self._cached_attn_blk_masks[ttt_step]
def _prepare_decoder_attention_mask(
self, attention_mask, input_shape, past_key_values_length, device, dtype
):
"""Expand the 2-D attention mask to 4-D and apply causal mask."""
# create causal mask
# [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
combined_attention_mask = None
# construct causal mask
if input_shape[-1] > 1:
combined_attention_mask = make_causal_mask(
input_shape,
dtype,
device=device,
past_key_values_length=past_key_values_length,
)
# merge causal mask with padding mask
if attention_mask is not None:
# [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
expanded_attn_mask = expand_mask(attention_mask, dtype, tgt_len=input_shape[-1]).to(
device
)
combined_attention_mask = (
expanded_attn_mask
if combined_attention_mask is None
else expanded_attn_mask + combined_attention_mask
)
return combined_attention_mask
def _prepare_eagle_inputs(
self,
input_ids,
attention_mask,
position_ids,
eagle_cache,
base_outputs,
):
"""Helper function to prepare eagle inputs for the 0th eagle forward pass."""
b, seq_length = input_ids.shape
past_kv_len = eagle_cache.get_seq_length() if eagle_cache is not None else 0
seq_len_with_past = seq_length + past_kv_len
# Prepare eagle_input_embeds: Shift left 1 token
with torch.no_grad():
if base_outputs.input_embeds is None:
eagle_input_embeds = self._base_model_embeddings(input_ids.roll(-1, 1))
else:
eagle_input_embeds = base_outputs.input_embeds.roll(-1, 1)
# Prepare eagle_input_hiddens
if self.eagle_config.use_aux_hidden_state:
# concat base model intermediate (pre-norm) hiddens
eagle_input_hiddens = self.eagle_module.fc(base_outputs.aux_hiddens)
else:
# use base model output (post-norm)hiddens
eagle_input_hiddens = base_outputs.out_hiddens
# Prepare attention_mask
if attention_mask is None:
eagle_attention_mask = torch.ones( # default: all tokens are valid
(b, seq_len_with_past), dtype=torch.bool, device=eagle_input_hiddens.device
)
else:
eagle_attention_mask = attention_mask.roll(-1, 1) # Shift left 1 token
# Expand the 2-D attention mask to 4-D and apply causal mask.
eagle_attention_mask = self._prepare_decoder_attention_mask(
eagle_attention_mask,
(b, seq_length),
past_kv_len,
eagle_input_hiddens.device,
eagle_input_hiddens.dtype,
)
# Prepare position_ids
if position_ids is None:
eagle_position_ids = (
torch.arange(
past_kv_len,
seq_len_with_past,
dtype=torch.long,
device=eagle_input_hiddens.device,
)
.unsqueeze(0)
.view(-1, seq_length)
)
else:
eagle_position_ids = position_ids.view(-1, seq_length).long()
base_model_logits = base_outputs.logits
if self.eagle_config.draft_vocab_size != self.eagle_config.vocab_size:
base_model_logits = self._map_logits_to_draft_vocab(base_model_logits)
base_output_predict_tok = base_model_logits.argmax(dim=-1).detach()
base_output_softmax_logits = torch.softmax(base_model_logits, dim=2)
# After LoRA warmup, stochastically detach logits — acts as dropout
# regularization on the eagle-loss-to-LoRA gradient path, preventing
# LoRA from degenerating to maximize EAGLE acc at cost of base quality.
# During warmup or when LoRA is off, always detach.
lora_active = getattr(self, "_lora_cotraining_active", False) and self.training
if lora_active and torch.rand(1).item() >= self.eagle_base_lora_logits_detach_prob:
pass # keep gradients flowing through logits to LoRA
else:
base_output_softmax_logits = base_output_softmax_logits.detach()
return (
eagle_input_embeds,
eagle_input_hiddens,
eagle_attention_mask,
eagle_position_ids,
base_output_predict_tok,
base_output_softmax_logits,
)
def _compute_ttt_attention_mask(
self, batch_size, seq_length, ttt_step
) -> BlockMask | torch.Tensor:
"""Return TTT attention_mask tensor of type BlockMask or Tensor depends on eagle attn impl."""
msk_func = get_ttt_msk_func(seq_length, ttt_step)
dtype = (
getattr(self._base_llm_config, "dtype", None)
or self.eagle_module.layers[0].input_layernorm.weight.dtype
)
dtypemin = torch.finfo(dtype).min
q_len = seq_length
kv_len = seq_length * (1 + ttt_step)
if self.eagle_config._attn_implementation == "flex_attention":
# Return block mask for flex attention
block_mask = create_block_mask(msk_func, B=None, H=None, Q_LEN=q_len, KV_LEN=kv_len)
return block_mask
else:
# Return tensor mask for non-flex attention
tensor_mask = msk_func(
None,
None,
torch.arange(q_len).view(1, 1, q_len, 1),
torch.arange(kv_len).view(1, 1, 1, kv_len),
).to(self.device)
tensor_mask = torch.full_like(
tensor_mask, 0, dtype=dtype, device=self.device
).masked_fill(~tensor_mask, dtypemin)
return tensor_mask
def _eagle_base_model_forward(
self,
input_ids,
attention_mask,
position_ids,
past_key_values,
freeze_base_model,
labels,
**kwargs,
):
def _run_forward(no_grad):
with torch.no_grad() if no_grad else contextlib.nullcontext():
return super(HFEagleModel, self).forward(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_values,
output_hidden_states=True,
**kwargs,
)
# With LoRA co-training, run a reference forward (LoRA disabled, no grad)
# to get the original base model logits for preservation loss, then run
# the main forward with LoRA enabled and gradients flowing.
# During warmup (_lora_cotraining_active=False), skip entirely.
lora_active = getattr(self, "_lora_cotraining_active", False) and self.training
ref_logits = None
if lora_active:
self._set_base_lora_enabled(False)
try:
ref_logits = _run_forward(no_grad=True).logits
finally:
if hasattr(self, "_aux_hidden_states"):
self._aux_hidden_states.clear()
self._set_base_lora_enabled(True)
outputs = _run_forward(no_grad=freeze_base_model and not lora_active)
past_key_values = getattr(outputs, "past_key_values", None)
base_model_logits = outputs.logits
if ref_logits is not None:
base_model_loss = self._preservation_loss(ref_logits, base_model_logits)
elif not freeze_base_model and labels is not None:
loss_fct = torch.nn.CrossEntropyLoss()
base_model_loss = loss_fct(
base_model_logits.view(-1, base_model_logits.shape[-1]), labels.view(-1)
)
else:
base_model_loss = None
return EagleBaseModelOutput(
input_embeds=outputs.hidden_states[0],
aux_hiddens=self.pop_and_gather_aux_hiddens(),
out_hiddens=outputs.hidden_states[-1],
logits=base_model_logits,
loss=base_model_loss,
), past_key_values
def _map_logits_to_draft_vocab(self, full_logits):
assert hasattr(self.eagle_module, "d2t"), "d2t buffer not initialized"
reverse_mapping = (
torch.arange(len(self.eagle_module.d2t)).to(self.eagle_module.d2t.device)
+ self.eagle_module.d2t
)
return full_logits[:, :, reverse_mapping]
def _eagle_forward(
self,
eagle_input_hidden_states,
inputs_embeds,
attention_mask,
position_ids,
eagle_cache=None,
):
eagle_postnorm_h, eagle_prenorm_h, eagle_cache = self.eagle_module(
eagle_input_hidden_states,
inputs_embeds,
attention_mask=attention_mask,
position_ids=position_ids,
use_cache=True,
past_key_values=eagle_cache,
)
eagle_lm_head = (
self.eagle_module.lm_head
if hasattr(self.eagle_module, "lm_head")
else self._base_model_lm_head
)
eagle_logits = eagle_lm_head(eagle_postnorm_h)
return eagle_postnorm_h, eagle_prenorm_h, eagle_logits, eagle_cache
def forward(
self,
input_ids: torch.LongTensor,
attention_mask: torch.Tensor | None = None,
position_ids: torch.LongTensor | None = None,
past_key_values: Cache | None = None,
inputs_embeds: torch.FloatTensor | None = None,
labels: torch.LongTensor | None = None,
use_cache: bool | None = None,
output_attentions: bool | None = None,
output_hidden_states: bool | None = None,
cache_position: torch.LongTensor | None = None,
logits_to_keep: int = 0,
loss_mask: torch.Tensor | None = None,
**kwargs,
) -> Any:
"""Forward pass of the EagleModel.
Returns:
loss: Loss of base model or eagle model.
logits: Base model logits.
past_key_values: Base model past key values with eagle cache attached.
hidden_states: Base model hidden states.
train_acc: Drafter training accuracies.
"""
eagle_cache = getattr(past_key_values, "eagle_cache", None)
if self.training:
assert past_key_values is None, "past_key_values should be None in training"
if loss_mask is None:
# By default, mask out padding tokens in loss computation
loss_mask = (
attention_mask.clone().detach()
if attention_mask is not None
else torch.ones_like(input_ids, dtype=torch.bool)
)
# ====First, run base model forward====
if self.eagle_offline:
# Parse base model outputs forwarded from teacher
assert "base_model_outputs" in kwargs
base_outputs = EagleBaseModelOutput.from_offline_dict(kwargs["base_model_outputs"])
if base_outputs.logits is None:
base_outputs.logits = self._base_model_lm_head(base_outputs.out_hiddens)
past_key_values = None
else:
with self._nvtx_range("base_model_forward"):
base_outputs, past_key_values = self._eagle_base_model_forward(
input_ids,
attention_mask,
position_ids,
past_key_values,
self.eagle_freeze_base_model,
labels,
**kwargs,
)
if not isinstance(past_key_values, Cache):
past_key_values = DynamicCache(config=self._base_llm_config)
if not isinstance(eagle_cache, Cache):
eagle_cache = DynamicCache(config=self.eagle_module.config)
past_key_values.eagle_cache = eagle_cache
# ====Prepare inputs for the first eagle forward pass====
eagle_loss = None
num_ttt = self.eagle_ttt_steps
train_accs = torch.zeros(1, num_ttt, device=input_ids.device)
b, seq_length, _ = base_outputs.out_hiddens.shape
with self._nvtx_range("prepare_eagle_inputs"):
(
eagle_input_embeds,
eagle_input_hiddens,
eagle_attn_mask_0,
eagle_position_ids,
base_output_predict_tok,
base_output_softmax_logits,
) = self._prepare_eagle_inputs(
input_ids,
attention_mask,
position_ids,
eagle_cache,
base_outputs,
)
# Init rotary_emb outside _eagle_forward so inv_freq is not captured by CUDAGraph.
self.eagle_module._maybe_init_rope(device=eagle_input_hiddens.device)
# ====Run eagle forward with extra training-time-test steps====
num_ttt_steps = self.eagle_ttt_steps if self.training else 1
for ttt_step in range(num_ttt_steps):
# TODO: (hg) during cp training, this mask is not used. Maybe turn it off then.
eagle_attention_mask = (
eagle_attn_mask_0
if self.eagle_mix_hidden_states or ttt_step == 0
else self._get_ttt_attention_mask(b, seq_length, ttt_step)
)
with self._enable_cp_ttt(), self._nvtx_range("eagle_forward"):
_, eagle_output_hiddens, eagle_logits, eagle_cache = self._eagle_forward(
eagle_input_hiddens,
eagle_input_embeds,
eagle_attention_mask,
eagle_position_ids,
None if self.eagle_mix_hidden_states else eagle_cache,
)
eagle_output_hiddens = eagle_output_hiddens.roll(1, 1)
if self.eagle_mix_hidden_states:
batch_size, seq_len_s, _ = eagle_input_hiddens.shape
num_to_replace = max(1, seq_len_s // (2**ttt_step + 1))
# Randomly select positions for each batch to replace
rand_indices = torch.rand(
batch_size, seq_len_s, device=eagle_input_hiddens.device
).argsort(dim=1)[:, :num_to_replace]
# Clone to avoid inplace modification that breaks autograd
eagle_input_hiddens = eagle_input_hiddens.clone()
batch_indices = torch.arange(batch_size)[:, None]
eagle_input_hiddens[batch_indices, rand_indices] = eagle_output_hiddens[
batch_indices, rand_indices
]
else:
eagle_input_hiddens = eagle_output_hiddens
with self._nvtx_range("eagle_loss"):
classification_loss, acc = self._eagle_loss(
# base model predict +1 tok, while eagle predict +2
# so we shift base model outputs compared to eagle outputs
# additionally, we mask the first n tok of eagle outputs at nth TTT step
base_output_softmax_logits[:, 1 + ttt_step :],
base_output_predict_tok[:, 1 + ttt_step :],
eagle_logits[:, ttt_step:-1],
loss_mask[:, 1 + ttt_step :],
)
# Apply loss decay factor to focus on early steps
classification_loss *= self.eagle_loss_decay_factor**ttt_step
eagle_loss = (
classification_loss if eagle_loss is None else eagle_loss + classification_loss
)
train_accs[0, ttt_step] = acc
train_accs = train_accs[:, :num_ttt_steps].tolist()
# Merge eagle loss and preservation loss (if LoRA co-training)
if base_outputs.loss is None and eagle_loss is None:
loss = None
assert not self.training, "At least one loss must be computed for training."
else:
loss = (base_outputs.loss or 0) + (eagle_loss or 0)
return ModelOutput(
loss=loss,
logits=base_outputs.logits,
past_key_values=past_key_values,
hidden_states=base_outputs.out_hiddens,
train_acc=train_accs,
eagle_loss=eagle_loss,
preservation_loss=base_outputs.loss if self.eagle_base_lora else None,
)
def _eagle_loss(
self,
base_output_softmax_logits,
base_output_predict_tok,
eagle_logits,
loss_mask,
):
"""Function for EAGLE loss computing."""
loss_mask = loss_mask[:, : eagle_logits.shape[1], None]
eagle_logsoft = torch.log_softmax(eagle_logits, dim=2)
classification_loss = -torch.sum(
torch.sum(loss_mask * base_output_softmax_logits * eagle_logsoft, 2)
) / (loss_mask.sum() + 1e-5)
# Compute accuracy (returned as tensor to avoid sync; .item() called after TTT loop)
eagle_predict_tok = eagle_logits.detach().argmax(dim=-1)
valid = loss_mask[:, :, 0].bool()
correct = (base_output_predict_tok == eagle_predict_tok) & valid
denom = valid.sum().clamp_min(1).float()
accuracy = correct.sum().float() / denom
return classification_loss, accuracy
@torch.no_grad()
def pseudo_speculative_generate(
self,
input_ids: torch.Tensor,
steps: int = 1,
):
"""Pseudo generate of the EAGLE GPTModel.
Returns:
base_token (torch.Tensor): token from base model
draft_tokens (torch.Tensor): draft tokens from eagle module
"""
base_model_outputs = super().forward(
input_ids=input_ids,
output_hidden_states=True,
)
base_model_hidden_states = base_model_outputs.hidden_states[-1]
base_model_logits = base_model_outputs.logits
base_token = base_model_logits[:, -1:, :].argmax(dim=-1).to(input_ids.device)
# Early return
if steps < 1:
if hasattr(self, "_aux_hidden_states"):
_ = self.pop_and_gather_aux_hiddens()
return base_token, None
eagle_ids = torch.cat((input_ids[:, 1:], base_token), dim=-1)
if self.eagle_config.use_aux_hidden_state:
# Only the first iteration input_hidden_states are from aux_hidden_state layers
# Gather _aux_hidden_states from all devices before concatenation
eagle_input_hidden_states = self.eagle_module.fc(self.pop_and_gather_aux_hiddens())
else:
eagle_input_hidden_states = base_model_hidden_states
# Init rotary_emb outside _eagle_forward so inv_freq is not captured by CUDAGraph.
self.eagle_module._maybe_init_rope(device=eagle_input_hidden_states.device)
draft_tokens = []
for step in range(steps):
b, seq_length = eagle_ids.shape
eagle_attention_mask = self._prepare_decoder_attention_mask(
None,
(b, seq_length),
0,
eagle_input_hidden_states.device,
eagle_input_hidden_states.dtype,
)
# Use SDPA attention during generation for both stability and performance
with (
temporary_set_config_value(self.eagle_config, "_attn_implementation", "sdpa"),
self._nvtx_range("eagle_forward"),
):
_, eagle_prenorm_h, eagle_logits, _ = self._eagle_forward(
eagle_input_hidden_states,
self._base_model_embeddings(eagle_ids),
eagle_attention_mask,
None,
)
draft_token = eagle_logits[:, -1:, :].argmax(dim=-1)
if self.eagle_config.draft_vocab_size != self.eagle_config.vocab_size:
draft_token += self.eagle_module.d2t[draft_token]
draft_tokens.append(draft_token)
eagle_ids = torch.cat((eagle_ids, draft_token.to(eagle_ids.device)), dim=-1)
eagle_input_hidden_states = torch.cat(
(eagle_input_hidden_states, eagle_prenorm_h[:, -1:, :]), dim=1
)
draft_tokens = torch.cat(draft_tokens, dim=-1).to(base_token.device)
return base_token, draft_tokens
class HFARValidation(AcceptanceRateValidation):
"""This is the subclass for HF model AR validation."""
def get_ground_truth(self, input_ids, osl):
"""This function returns ground truth output tokens from the base model."""
input_ids = copy.deepcopy(input_ids).to(torch.cuda.current_device())
for _ in range(osl):
input_id, _ = self.model.pseudo_speculative_generate(input_ids, steps=0)
input_ids = torch.cat((input_ids, input_id.to(input_ids.device)), dim=-1)
if input_id[0, 0] == self.end_token:
break
return input_ids
@@ -0,0 +1,167 @@
# Adapted from: https://github.com/ctlllll/axolotl/blob/f86767e/src/axolotl/monkeypatch/medusa_utils.py
#
# 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.
"""Medusa speculative decoding plugin for HuggingFace models."""
import contextlib
from typing import Any
import torch
from torch import nn
from torch.nn import CrossEntropyLoss
from transformers import Cache, PreTrainedModel
from transformers.trainer_pt_utils import LabelSmoother
from transformers.utils import ModelOutput
from ..medusa.conversion import MedusaDMRegistry
from ..medusa.medusa_model import MedusaModel
from ..utils import ResBlock
__all__ = ["HFMedusaModel"]
IGNORE_TOKEN_ID = LabelSmoother.ignore_index
@MedusaDMRegistry.register({PreTrainedModel: "hf.PreTrainedModel"})
class HFMedusaModel(MedusaModel):
"""Medusa Model Class for huggingface models."""
def modify(self, medusa_num_heads=0, medusa_num_layers=0):
"""Constructor.
Args:
medusa_num_heads: number of medusa heads.
medusa_num_layers: number of ResBlock layers in each head.
"""
super().modify(medusa_num_heads=medusa_num_heads, medusa_num_layers=medusa_num_layers)
self.config.medusa = {
"num_medusa_heads": medusa_num_heads,
"num_medusa_layers": medusa_num_layers,
}
hidden_size = self.lm_head.weight.shape[-1]
vocab_size = self.lm_head.weight.shape[0]
# Create a list of Medusa heads
self.medusa_heads = nn.ModuleList(
[
nn.Sequential(
*([ResBlock(hidden_size) for _ in range(self.medusa_num_layers)]),
nn.Linear(hidden_size, vocab_size, bias=False),
)
for _ in range(self.medusa_num_heads)
]
)
# Ensure medusa_head's dtype and device align with the base_model
self.medusa_heads.to(self.lm_head.weight.dtype).to(self.lm_head.weight.device)
self.medusa_heads.device = self.lm_head.weight.device
if hasattr(self, "hf_device_map") and "lm_head" in self.hf_device_map:
self.hf_device_map["medusa_heads"] = self.hf_device_map["lm_head"]
def forward(
self,
input_ids: torch.LongTensor | None = None,
attention_mask: torch.Tensor | None = None,
position_ids: torch.LongTensor | None = None,
past_key_values: Cache | None = None,
inputs_embeds: torch.FloatTensor | None = None,
labels: torch.LongTensor | None = None,
use_cache: bool | None = None,
output_attentions: bool | None = None,
output_hidden_states: bool | None = None,
cache_position: torch.LongTensor | None = None,
logits_to_keep: int | torch.Tensor = 0,
freeze_base_model: bool = True,
medusa_heads_coefficient: float | None = 0.2,
medusa_decay_coefficient: float | None = 0.8,
**kwargs,
) -> Any:
"""Forward pass of the MedusaModel.
Returns:
torch.Tensor: A tensor containing predictions from all Medusa heads.
"""
# Pass input through the base model
with torch.no_grad() if freeze_base_model else contextlib.nullcontext():
outputs = self.model(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_values,
inputs_embeds=inputs_embeds,
use_cache=use_cache,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
rcache_position=cache_position,
**kwargs,
)
hidden_states = outputs.last_hidden_state
# Only compute necessary logits, and do not upcast them to float if we are not computing the loss
slice_indices = (
slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep
)
logits = self.lm_head(hidden_states[:, slice_indices, :])
medusa_logits = [
self.medusa_heads[i](hidden_states[:, slice_indices, :])
for i in range(self.medusa_num_heads)
]
if labels is not None:
loss = 0
loss_fct = CrossEntropyLoss()
# Base model loss
if not freeze_base_model:
loss_logits = logits.view(-1, logits.shape[-1])
loss_labels = labels.view(-1)
base_model_loss = loss_fct(loss_logits, loss_labels)
loss += base_model_loss
# Medusa loss
for i in range(self.medusa_num_heads):
labels = labels[..., 1:].contiguous()
loss_logits = medusa_logits[i][:, : -(1 + i)].contiguous()
loss_logits = loss_logits.view(-1, loss_logits.shape[-1])
loss_labels = labels.view(-1)
loss += (
loss_fct(loss_logits, loss_labels)
* medusa_decay_coefficient**i
* medusa_heads_coefficient
)
else:
loss = None
return ModelOutput(
loss=loss,
logits=logits,
medusa_logits=medusa_logits,
past_key_values=outputs.past_key_values,
hidden_states=outputs.hidden_states,
attentions=outputs.attentions,
)
@@ -0,0 +1,239 @@
# 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.
"""DFlash draft model architecture (DFlashModule) and related components.
Draft model components use Qwen3 (MLP, RMSNorm, RotaryEmbedding) from
``transformers.models.qwen3``, matching z-lab's reference checkpoint format.
The draft architecture is independent of the target model.
"""
import torch
from torch import nn
from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS
from transformers.models.qwen3.modeling_qwen3 import Qwen3MLP as _MLP_CLS # noqa: N814
from transformers.models.qwen3.modeling_qwen3 import Qwen3RMSNorm as _NORM_CLS # noqa: N814
from transformers.models.qwen3.modeling_qwen3 import (
Qwen3RotaryEmbedding as _ROTARY_CLS, # noqa: N814
)
from transformers.models.qwen3.modeling_qwen3 import rotate_half as _rotate_half
__all__ = ["DFlashModule", "build_target_layer_ids"]
def build_target_layer_ids(num_target_layers, num_draft_layers):
"""Select layers uniformly from the target model for feature extraction."""
if num_target_layers < num_draft_layers:
raise ValueError(
f"num_target_layers ({num_target_layers}) must be >= num_draft_layers ({num_draft_layers})"
)
if num_draft_layers == 1:
return [num_target_layers // 2]
start = min(1, num_target_layers - 1)
end = max(start, num_target_layers - 3)
span = end - start
return [round(start + (i * span) / (num_draft_layers - 1)) for i in range(num_draft_layers)]
def apply_rotary_pos_emb(q, k, cos, sin):
"""Apply RoPE. Q uses last q_len positions, K uses all positions."""
cos = cos.unsqueeze(1) # [B, 1, seq, dim]
sin = sin.unsqueeze(1)
q_len = q.size(2)
q_embed = (q * cos[:, :, -q_len:, :]) + (_rotate_half(q) * sin[:, :, -q_len:, :])
k_embed = (k * cos) + (_rotate_half(k) * sin)
return q_embed, k_embed
class DFlashAttention(nn.Module):
"""Attention with KV injection, using HF's attention dispatch."""
def __init__(self, config, layer_idx):
"""Initialize DFlash attention with KV injection projections and QK-norm."""
super().__init__()
self.config = config
self.layer_idx = layer_idx
self.head_dim = getattr(
config, "head_dim", config.hidden_size // config.num_attention_heads
)
self.num_heads = config.num_attention_heads
self.num_kv_heads = config.num_key_value_heads
self.num_key_value_groups = self.num_heads // self.num_kv_heads
self.scaling = self.head_dim**-0.5
self.attention_dropout = getattr(config, "attention_dropout", 0.0)
self.is_causal = False
attn_bias = getattr(config, "attention_bias", False)
self.q_proj = nn.Linear(config.hidden_size, self.num_heads * self.head_dim, bias=attn_bias)
self.k_proj = nn.Linear(
config.hidden_size, self.num_kv_heads * self.head_dim, bias=attn_bias
)
self.v_proj = nn.Linear(
config.hidden_size, self.num_kv_heads * self.head_dim, bias=attn_bias
)
self.o_proj = nn.Linear(self.num_heads * self.head_dim, config.hidden_size, bias=attn_bias)
self.q_norm = _NORM_CLS(self.head_dim, eps=config.rms_norm_eps)
self.k_norm = _NORM_CLS(self.head_dim, eps=config.rms_norm_eps)
# Resolve HF attention function
self._attn_fn = None
# Qwen3 uses sliding window attention on some layers (config.layer_types)
if hasattr(config, "layer_types") and hasattr(config, "sliding_window"):
is_sliding = config.layer_types[layer_idx] == "sliding_attention"
self.sliding_window = config.sliding_window if is_sliding else None
else:
self.sliding_window = None
def _get_attn_fn(self):
"""Lazily resolve the HF attention function (default: sdpa)."""
if self._attn_fn is not None:
return self._attn_fn
impl = self.config._attn_implementation # default set in dflash/default_config.py
self._attn_fn = ALL_ATTENTION_FUNCTIONS.get(impl, ALL_ATTENTION_FUNCTIONS["sdpa"])
return self._attn_fn
def forward(self, hidden_states, target_hidden, position_embeddings, attention_mask=None):
"""Forward with KV injection.
Q is projected from the noise block (draft token embeddings: [anchor, mask, mask, ...]).
K and V are projected from the concatenation of target hidden states (context from the
base model) and noise block, so the draft can attend to both context and its own block.
"""
bsz, q_len, _ = hidden_states.shape
ctx_len = target_hidden.shape[1]
# Q from noise block only (the draft tokens being predicted), with QK-norm
q = self.q_proj(hidden_states).view(bsz, q_len, -1, self.head_dim)
q = self.q_norm(q).transpose(1, 2)
# K from context + noise, with QK-norm
k_ctx = self.k_proj(target_hidden)
k_noise = self.k_proj(hidden_states)
k = torch.cat([k_ctx, k_noise], dim=1).view(bsz, ctx_len + q_len, -1, self.head_dim)
k = self.k_norm(k).transpose(1, 2)
# V from context + noise (no norm)
v_ctx = self.v_proj(target_hidden)
v_noise = self.v_proj(hidden_states)
v = (
torch.cat([v_ctx, v_noise], dim=1)
.view(bsz, ctx_len + q_len, -1, self.head_dim)
.transpose(1, 2)
)
# RoPE
cos, sin = position_embeddings
q, k = apply_rotary_pos_emb(q, k, cos, sin)
# Use HF's attention dispatch (handles GQA internally)
attn_fn = self._get_attn_fn()
attn_output, _ = attn_fn(
self,
q,
k,
v,
attention_mask,
dropout=0.0 if not self.training else self.attention_dropout,
scaling=self.scaling,
sliding_window=self.sliding_window,
)
attn_output = attn_output.reshape(bsz, q_len, -1)
return self.o_proj(attn_output)
class DFlashDecoderLayer(nn.Module):
"""Draft decoder layer with KV injection."""
def __init__(self, config, layer_idx):
"""Initialize decoder layer with attention, MLP, and layer norms."""
super().__init__()
self.self_attn = DFlashAttention(config, layer_idx)
self.mlp = _MLP_CLS(config)
self.input_layernorm = _NORM_CLS(config.hidden_size, eps=config.rms_norm_eps)
self.post_attention_layernorm = _NORM_CLS(config.hidden_size, eps=config.rms_norm_eps)
def forward(self, hidden_states, target_hidden, position_embeddings, attention_mask=None):
"""Forward pass with residual connections."""
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
hidden_states = self.self_attn(
hidden_states, target_hidden, position_embeddings, attention_mask
)
hidden_states = residual + hidden_states
residual = hidden_states
hidden_states = self.post_attention_layernorm(hidden_states)
hidden_states = self.mlp(hidden_states)
hidden_states = residual + hidden_states
return hidden_states
class DFlashModule(nn.Module):
"""DFlash draft module using Qwen3 components (MLP, RMSNorm, RotaryEmbedding)."""
def __init__(self, config):
"""Initialize DFlash module with feature fusion, decoder layers, and rotary embeddings."""
super().__init__()
self.config = config
self.block_size = config.block_size
# Feature fusion
num_fused_layers = len(config.target_layer_ids)
self.fc = nn.Linear(num_fused_layers * config.hidden_size, config.hidden_size, bias=False)
self.hidden_norm = _NORM_CLS(config.hidden_size, eps=config.rms_norm_eps)
# Decoder layers
self.layers = nn.ModuleList(
[DFlashDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
)
self.norm = _NORM_CLS(config.hidden_size, eps=config.rms_norm_eps)
self._rotary_config = config # Used by _maybe_init_rotary_emb
# Explicit weight init is needed because DFlashModule is instantiated via
# mtsp.convert() AFTER the base model's post_init() has already run, so HF's
# automatic _init_weights walk doesn't reach these new layers.
self._init_weights(config)
def _maybe_init_rotary_emb(self, device=None):
"""Lazily initialize rotary embeddings on first forward call.
Same pattern as EAGLE3's _maybe_init_rope. Avoids creating rotary_emb
during __init__ (which runs on meta device during from_pretrained),
preventing the meta-tensor inv_freq issue on checkpoint resume.
"""
if not hasattr(self, "rotary_emb"):
self.rotary_emb = _ROTARY_CLS(config=self._rotary_config, device=device)
def _init_weights(self, config):
"""Initialize weights matching HF PreTrainedModel._init_weights."""
std = getattr(config, "initializer_range", 0.02)
for module in self.modules():
if isinstance(module, nn.Linear):
nn.init.normal_(module.weight, mean=0.0, std=std)
if module.bias is not None:
nn.init.zeros_(module.bias)
def forward(self, noise_embedding, target_hidden, position_ids, attention_mask=None):
"""Forward with feature fusion, KV injection, and position embeddings."""
hidden_states = noise_embedding
target_hidden = self.hidden_norm(self.fc(target_hidden))
self._maybe_init_rotary_emb(device=hidden_states.device)
position_embeddings = self.rotary_emb(hidden_states, position_ids)
for layer in self.layers:
hidden_states = layer(hidden_states, target_hidden, position_embeddings, attention_mask)
return self.norm(hidden_states)
@@ -0,0 +1,215 @@
# 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.
"""EAGLE draft model architecture (EagleModule) and related data structures."""
from dataclasses import dataclass
import torch
from torch import nn
from transformers import Cache
from transformers.models.llama.modeling_llama import LlamaRMSNorm, LlamaRotaryEmbedding
__all__ = ["EagleBaseModelOutput", "EagleModule"]
class EagleModule(nn.Module):
"""Eagle module used in EAGLE model."""
def __init__(self, config, decoder_layer_cls, bias=False):
"""Init function for EagleModule."""
super().__init__()
self.config = config
self.layers = nn.ModuleList(
[decoder_layer_cls(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
)
if config.use_last_layernorm:
self.norm = LlamaRMSNorm(config.hidden_size, config.rms_norm_eps)
# Optionally, we use a smaller vocab table for eagle module
if config.draft_vocab_size != config.vocab_size or config.has_lm_head:
# Need an extra lm_head for eagle module since vocab size is reduced.
assert config.draft_vocab_size <= config.vocab_size, (
"EAGLE module's vocab size should be <= base model vocab size!"
)
# Initialize the buffers to zero.
# Their values depend on specific tokenizer and calibration dataset, and should be set in training script.
if config.draft_vocab_size < config.vocab_size:
self.register_buffer("d2t", torch.zeros(config.draft_vocab_size, dtype=torch.int64))
self.lm_head = nn.Linear(
config.hidden_size,
config.draft_vocab_size,
bias=False,
)
if config.use_aux_hidden_state:
# In EAGLE-3, the FC concentrate hidden states from multiple base model layers
self.fc = nn.Linear(
len(config.eagle_aux_hidden_state_layer_ids) * config.hidden_size,
config.hidden_size,
bias=bias,
)
first_layer_attn = self.layers[0].self_attn
# Expand first attn input dim since it accepts cat(input_embeds, hidden_states)
self._expand_first_attn_in_dim(first_layer_attn)
# EAGLE-3's first attention require [input_layernorm_output, aux_hidden_states]
first_layer_attn.register_forward_pre_hook(
self._eagle3_attention_forward_pre_hook, with_kwargs=True
)
# In EAGLE-3, input_embeds and hidden_states are normalized separately before concatenation.
self.layers[0].input_layernorm = LlamaRMSNorm(
config.hidden_size, eps=config.rms_norm_eps
)
self.layers[0].hidden_norm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
def _maybe_init_rope(self, device=None):
if self.config.eagle_decoder_type == "llama" and not hasattr(self, "rotary_emb"):
self.rotary_emb = LlamaRotaryEmbedding(config=self.config, device=device)
def _expand_first_attn_in_dim(self, first_layer_attn):
"""Modify qkv projection in first layer to accept 2h hidden size."""
# Find Linear modules to expand
eagle_attn_type = type(first_layer_attn)
if eagle_attn_type.__name__ == "LlamaAttention":
expand_modules = ["q_proj", "k_proj", "v_proj"]
elif eagle_attn_type.__name__ == "DeepseekV3Attention":
if first_layer_attn.q_lora_rank is None:
expand_modules = ["q_proj", "kv_a_proj_with_mqa"]
else:
expand_modules = ["q_a_proj", "kv_a_proj_with_mqa"]
else:
raise ValueError(f"Unsupported attention type: {eagle_attn_type}")
# Replace Linear with 2x input dim
for module in expand_modules:
original_linear = getattr(first_layer_attn, module)
assert isinstance(original_linear, nn.Linear), f"Module {module} is not a Linear"
setattr(
first_layer_attn,
module,
nn.Linear(
original_linear.in_features * 2,
original_linear.out_features,
bias=first_layer_attn.config.attention_bias,
),
)
def _eagle3_attention_forward_pre_hook(self, module, args, kwargs):
"""Concat input_embeds and hidden_states for EAGLE-3's first attention layer."""
if "hidden_states" not in kwargs:
raise ValueError("hidden_states not found in kwargs")
if self._input_embeds is None:
raise ValueError("self._input_embeds is None")
input_embeds = self._input_embeds
self._input_embeds = None
kwargs["hidden_states"] = torch.cat(
(input_embeds, self.layers[0].hidden_norm(kwargs["hidden_states"])), dim=-1
)
return args, kwargs
def forward(
self,
hidden_states: torch.Tensor,
inputs_embeds: torch.Tensor,
attention_mask: torch.Tensor,
position_ids: torch.LongTensor | None = None,
past_key_values: Cache | None = None,
use_cache: bool | None = None,
output_attentions: bool | None = False,
):
"""Forward function for EagleModule."""
batch_size, seq_length, _ = hidden_states.shape
seq_length_with_past = seq_length
past_key_values_length = 0
if past_key_values is not None:
past_key_values_length = past_key_values.get_seq_length()
seq_length_with_past = seq_length_with_past + past_key_values_length
if position_ids is None:
device = hidden_states.device if hidden_states is not None else inputs_embeds.device
position_ids = torch.arange(
past_key_values_length,
seq_length + past_key_values_length,
dtype=torch.long,
device=device,
)
position_ids = position_ids.unsqueeze(0).view(-1, seq_length)
else:
position_ids = position_ids.view(-1, seq_length).long()
inputs_embeds = inputs_embeds.to(hidden_states.dtype).to(hidden_states.device)
# In EAGLE-3, we save input embeddings to attribute, and use it in first decoder layer by hook function
# Also, we normalize input embeddings and hidden states before concatenating them.
# The default input norm in first layer attn will be disabled.
self._input_embeds = self.layers[0].input_layernorm(inputs_embeds)
if self.config.eagle_decoder_type == "llama":
# rotary_emb must be pre-initialized by the caller (see HFEagleModel);
# lazy init here would allocate inv_freq inside the torch.compile/CUDAGraph
# capture region and get overwritten by subsequent runs.
position_embeddings = self.rotary_emb(hidden_states, position_ids)
else:
position_embeddings = None
for decoder_layer in self.layers:
layer_outputs = decoder_layer(
hidden_states,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_values,
output_attentions=output_attentions,
use_cache=use_cache,
position_embeddings=position_embeddings,
)
# For HF>= 4.54.0, the layer_outputs is a tensor, for older, it is a tuple.
if isinstance(layer_outputs, tuple):
hidden_states = layer_outputs[0]
else:
hidden_states = layer_outputs
pre_norm_h = hidden_states
post_norm_h = self.norm(hidden_states) if hasattr(self, "norm") else hidden_states
return post_norm_h, pre_norm_h, past_key_values
@dataclass
class EagleBaseModelOutput:
"""Output container for base model forward pass in EAGLE training."""
out_hiddens: torch.Tensor
aux_hiddens: torch.Tensor | None = None
logits: torch.Tensor | None = None
input_embeds: torch.Tensor | None = None
loss: torch.Tensor | None = None
@classmethod
def from_offline_dict(cls, d: dict):
"""Construct from a dict of pre-computed base model outputs (offline training)."""
return cls(
out_hiddens=d.get("base_model_hidden_states"),
aux_hiddens=d.get("aux_hidden_states"),
logits=d.get("base_model_logits"),
input_embeds=d.get("base_model_input_embeds"),
loss=None,
)
File diff suppressed because it is too large Load Diff
+3 -3
View File
@@ -554,14 +554,14 @@ def get_ttt_msk_func(seq_length, ttt_step):
@contextlib.contextmanager
def enable_cp_ttt_patch():
"""Context manager to enable CP TTT patch."""
import modelopt.torch.speculative.plugins.transformers
import modelopt.torch.speculative.plugins.hf_eagle
modelopt.torch.speculative.plugins.transformers.ENABLE_CP_TTT_PATCH = True
modelopt.torch.speculative.plugins.hf_eagle.ENABLE_CP_TTT_PATCH = True
with sdpa_kernel([SDPBackend.CUDNN_ATTENTION, SDPBackend.MATH]):
try:
yield
finally:
modelopt.torch.speculative.plugins.transformers.ENABLE_CP_TTT_PATCH = False
modelopt.torch.speculative.plugins.hf_eagle.ENABLE_CP_TTT_PATCH = False
def load_vlm_or_llm(
@@ -37,6 +37,9 @@ def _make_exporter(
model = MagicMock()
model.eagle_config.eagle_decoder_type = "llama"
model.eagle_config.rope_scaling = {"rope_type": rope_type, "rope_theta": rope_theta}
# rope_theta lives inside rope_scaling in transformers 5.x; clear the top-level attr
# so the fallback path is exercised instead of MagicMock's auto-attr.
model.eagle_config.rope_theta = None
model.eagle_export_rope_scaling = eagle_export_rope_scaling
model._draft_model_config = None
model.config.rope_scaling = None
@@ -55,16 +58,18 @@ def test_yarn_rope_injected_with_correct_config():
assert config["rope_scaling"] == DEFAULT_ROPE_SCALING
def test_rope_not_injected_when_non_default_training_rope():
"""rope_scaling is not overridden when training rope_type is not 'default'."""
def test_rope_not_overridden_when_non_default_training_rope():
"""Export override is not applied when training rope_type is not 'default';
rope_scaling falls through to the training config."""
config = _make_exporter(rope_type="llama3")._export_config()
assert config.get("rope_scaling") is None
assert config["rope_scaling"] == {"rope_type": "llama3", "rope_theta": 10000}
def test_rope_not_injected_when_eagle_export_rope_scaling_is_empty():
"""rope_scaling is not injected when eagle_export_rope_scaling is empty dict."""
def test_rope_not_overridden_when_eagle_export_rope_scaling_is_empty():
"""Export override is not applied when eagle_export_rope_scaling is empty;
rope_scaling falls through to the training config."""
config = _make_exporter(eagle_export_rope_scaling={})._export_config()
assert config.get("rope_scaling") is None
assert config["rope_scaling"] == {"rope_type": "default", "rope_theta": 10000}
def test_rope_theta_fallback_from_rope_scaling():