mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[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:
@@ -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|
|
||||
|
||||
@@ -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 *
|
||||
|
||||
@@ -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
@@ -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():
|
||||
|
||||
Reference in New Issue
Block a user