Weiyun1025's picture
Upload folder using huggingface_hub
201c300 verified
Raw History Blame
27.1 kB
from typing import Callable, Optional
import torch
from torch import nn
from transformers import DynamicCache
from transformers.cache_utils import Cache
from transformers.modeling_outputs import CausalLMOutputWithPast
from transformers.models.qwen3.modeling_qwen3 import (
ALL_ATTENTION_FUNCTIONS,
FlashAttentionKwargs,
GradientCheckpointingLayer,
Qwen3Config,
Qwen3MLP,
Qwen3PreTrainedModel,
Qwen3RMSNorm,
Qwen3RotaryEmbedding,
eager_attention_forward,
rotate_half,
)
from typing_extensions import Tuple, Unpack
def sample(logits: torch.Tensor, temperature: float = 0.0) -> torch.Tensor:
if temperature < 1e-5:
return torch.argmax(logits, dim=-1)
bsz, seq_len, vocab_size = logits.shape
logits = logits.view(-1, vocab_size)
logits = logits / temperature
probs = torch.softmax(logits, dim=-1)
return torch.multinomial(probs, num_samples=1).view(bsz, seq_len)
def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1):
cos = cos.unsqueeze(unsqueeze_dim)
sin = sin.unsqueeze(unsqueeze_dim)
rotary_dim = cos.size(-1)
if rotary_dim > q.size(-1) or rotary_dim > k.size(-1):
raise ValueError(
f"RoPE dim ({rotary_dim}) exceeds q/k dim ({q.size(-1)}, {k.size(-1)})."
)
q_rot, q_pass = q[..., :rotary_dim], q[..., rotary_dim:]
k_rot, k_pass = k[..., :rotary_dim], k[..., rotary_dim:]
q_len = q.size(-2)
q_rot = (q_rot * cos[..., -q_len:, :]) + (
rotate_half(q_rot) * sin[..., -q_len:, :]
)
k_rot = (k_rot * cos) + (rotate_half(k_rot) * sin)
q_embed = torch.cat((q_rot, q_pass), dim=-1)
k_embed = torch.cat((k_rot, k_pass), dim=-1)
return q_embed, k_embed
def apply_rotary_single(x, cos, sin, unsqueeze_dim=1):
"""Apply (partial) RoPE to a single tensor whose seq length matches cos/sin.
Used by the ``use_target_kv`` path, where only the draft's own (query and
in-block noise-key) tokens need the draft RoPE — the target-provided context
K is already rotated in the target's space and must be left untouched.
``x`` is ``[b, heads, L, head_dim]``; ``cos``/``sin`` are ``[b, L, rotary_dim]``.
"""
cos = cos.unsqueeze(unsqueeze_dim)
sin = sin.unsqueeze(unsqueeze_dim)
rotary_dim = cos.size(-1)
if rotary_dim > x.size(-1):
raise ValueError(
f"RoPE dim ({rotary_dim}) exceeds tensor dim ({x.size(-1)})."
)
x_rot, x_pass = x[..., :rotary_dim], x[..., rotary_dim:]
x_rot = (x_rot * cos) + (rotate_half(x_rot) * sin)
return torch.cat((x_rot, x_pass), dim=-1)
class Qwen3DFlashAttention(nn.Module):
"""Multi-headed attention from 'Attention Is All You Need' paper"""
def __init__(self, config: Qwen3Config, layer_idx: int):
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_key_value_groups = (
config.num_attention_heads // config.num_key_value_heads
)
self.scaling = self.head_dim**-0.5
self.attention_dropout = config.attention_dropout
self.is_causal = False
self.q_proj = nn.Linear(
config.hidden_size,
config.num_attention_heads * self.head_dim,
bias=config.attention_bias,
)
self.k_proj = nn.Linear(
config.hidden_size,
config.num_key_value_heads * self.head_dim,
bias=config.attention_bias,
)
self.v_proj = nn.Linear(
config.hidden_size,
config.num_key_value_heads * self.head_dim,
bias=config.attention_bias,
)
self.o_proj = nn.Linear(
config.num_attention_heads * self.head_dim,
config.hidden_size,
bias=config.attention_bias,
)
self.q_norm = Qwen3RMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.k_norm = Qwen3RMSNorm(self.head_dim, eps=config.rms_norm_eps)
# target-KV consumption mode (see DFlashDraftModel): when target_kv is
# supplied, "inject" adds it as a residual on top of the draft's own
# k_proj/v_proj context K/V, whereas the default "replace" (use_target_kv)
# uses the target's KV directly. Read per-attention so forward can branch.
_dflash_cfg = getattr(config, "dflash_config", {}) or {}
self.use_target_kv_inject = bool(_dflash_cfg.get("use_target_kv_inject", False))
layer_types = getattr(config, "layer_types", None)
is_sliding_layer = (
isinstance(layer_types, (list, tuple))
and layer_idx < len(layer_types)
and layer_types[layer_idx] == "sliding_attention"
)
self.sliding_window = config.sliding_window if is_sliding_layer else None
def forward(
self,
hidden_states: torch.Tensor,
target_hidden: torch.Tensor,
position_embeddings: tuple[torch.Tensor, torch.Tensor],
attention_mask: Optional[torch.Tensor],
past_key_values: Optional[Cache] = None,
cache_position: Optional[torch.LongTensor] = None,
target_kv: Optional[tuple[torch.Tensor, torch.Tensor]] = None,
**kwargs: Unpack[FlashAttentionKwargs],
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
bsz, q_len = hidden_states.shape[:-1]
cos, sin = position_embeddings
q = self.q_proj(hidden_states)
q = q.view(bsz, q_len, -1, self.head_dim)
q = self.q_norm(q).transpose(1, 2)
if target_kv is not None and not self.use_target_kv_inject:
# ---- use_target_kv (REPLACE) ------------------------------------
# Context K/V come straight from the *target* model's own KV: the
# provided K is already k_norm'd + RoPE'd (in the target's space,
# i.e. exactly what sits in the target KV cache) and V is the raw
# projected value. Only the in-block draft (noise) tokens are keyed
# by the draft's own k_proj/v_proj here, since the target has no KV
# for not-yet-generated block tokens.
k_ctx, v_ctx = target_kv # each [bsz, ctx_len, num_kv_heads, head_dim]
k_ctx = k_ctx.transpose(1, 2) # [bsz, nkv, ctx_len, head_dim]
v_ctx = v_ctx.transpose(1, 2)
k_noise = self.k_proj(hidden_states).view(bsz, q_len, -1, self.head_dim)
k_noise = self.k_norm(k_noise).transpose(1, 2) # [bsz, nkv, q_len, hd]
v_noise = (
self.v_proj(hidden_states)
.view(bsz, q_len, -1, self.head_dim)
.transpose(1, 2)
)
# The draft (noise) tokens live at the last q_len position ids; the
# target context K is already rotated, so only rotate q and k_noise.
cos_draft, sin_draft = cos[:, -q_len:, :], sin[:, -q_len:, :]
q = apply_rotary_single(q, cos_draft, sin_draft)
k_noise = apply_rotary_single(k_noise, cos_draft, sin_draft)
k = torch.cat([k_ctx, k_noise], dim=2) # [bsz, nkv, ctx_len+q_len, hd]
v = torch.cat([v_ctx, v_noise], dim=2)
else:
# ---- baseline, and use_target_kv_inject (baseline + residual) ----
ctx_len = target_hidden.shape[1]
k_ctx = self.k_proj(target_hidden)
k_noise = self.k_proj(hidden_states)
v_ctx = self.v_proj(target_hidden)
v_noise = self.v_proj(hidden_states)
k = torch.cat([k_ctx, k_noise], dim=1).view(
bsz, ctx_len + q_len, -1, self.head_dim
)
v = torch.cat([v_ctx, v_noise], dim=1).view(
bsz, ctx_len + q_len, -1, self.head_dim
)
k = self.k_norm(k).transpose(1, 2)
v = v.transpose(1, 2)
q, k = apply_rotary_pos_emb(q, k, cos, sin)
if target_kv is not None:
# INJECT: add the target's own K/V into the context slice as a
# residual on top of the draft's projected+normed+roped context
# K/V. (k/v are [bsz, nkv, ctx_len+q_len, hd]; the context is the
# leading ctx_len keys. target K is already roped in the target
# space, matching the draft's roped context via inherited RoPE.)
k_inj, v_inj = target_kv # each [bsz, ctx_len, nkv, hd]
ctxL = k_inj.shape[1]
k[:, :, :ctxL, :] = k[:, :, :ctxL, :] + k_inj.transpose(1, 2)
v[:, :, :ctxL, :] = v[:, :, :ctxL, :] + v_inj.transpose(1, 2)
if past_key_values is not None:
cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position}
k, v = past_key_values.update(k, v, self.layer_idx, cache_kwargs)
attn_fn: Callable = eager_attention_forward
if self.config._attn_implementation != "eager":
attn_fn = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]
attn_output, attn_weights = 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,
**kwargs,
)
attn_output = attn_output.reshape(bsz, q_len, -1)
attn_output = self.o_proj(attn_output)
return attn_output, attn_weights
class Qwen3DFlashDecoderLayer(GradientCheckpointingLayer):
def __init__(self, config: Qwen3Config, layer_idx: int):
super().__init__()
self.hidden_size = config.hidden_size
self.self_attn = Qwen3DFlashAttention(config=config, layer_idx=layer_idx)
self.mlp = Qwen3MLP(config)
self.input_layernorm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.post_attention_layernorm = Qwen3RMSNorm(
config.hidden_size, eps=config.rms_norm_eps
)
def forward(
self,
target_hidden: Optional[torch.Tensor] = None,
hidden_states: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.LongTensor] = None,
past_key_value: Optional[Cache] = None,
output_attentions: Optional[bool] = False,
use_cache: Optional[bool] = False,
cache_position: Optional[torch.LongTensor] = None,
position_embeddings: Optional[
Tuple[torch.Tensor, torch.Tensor]
] = None, # necessary, but kept here for BC
target_kv: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
**kwargs: Unpack[FlashAttentionKwargs],
) -> Tuple[
torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]
]:
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
hidden_states = self.self_attn(
hidden_states=hidden_states,
target_hidden=target_hidden,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_value,
output_attentions=output_attentions,
use_cache=use_cache,
cache_position=cache_position,
position_embeddings=position_embeddings,
target_kv=target_kv,
**kwargs,
)[0]
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
def build_target_layer_ids(num_target_layers: int, num_draft_layers: int):
if num_draft_layers == 1:
return [(num_target_layers // 2)]
start = 1
end = num_target_layers - 3
span = end - start
target_layer_ids = [
int(round(start + (i * span) / (num_draft_layers - 1)))
for i in range(num_draft_layers)
]
return target_layer_ids
def extract_context_feature(
hidden_states: list[torch.Tensor],
layer_ids: Optional[list[int]],
) -> torch.Tensor:
offset = 1
selected_states = []
for layer_id in layer_ids:
selected_states.append(hidden_states[layer_id + offset])
target_hidden = torch.cat(selected_states, dim=-1)
return target_hidden
class DFlashDraftModel(Qwen3PreTrainedModel):
config_class = Qwen3Config
_no_split_modules = ["Qwen3DFlashDecoderLayer"]
def __init__(self, config) -> None:
super().__init__(config)
self.config = config
self.layers = nn.ModuleList(
[
Qwen3DFlashDecoderLayer(config, layer_idx)
for layer_idx in range(config.num_hidden_layers)
]
)
dflash_config = getattr(config, "dflash_config", {}) or {}
self.target_layer_ids = dflash_config.get(
"target_layer_ids",
build_target_layer_ids(config.num_target_layers, config.num_hidden_layers),
)
self.norm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.rotary_emb = Qwen3RotaryEmbedding(config)
# When use_target_kv is on, the draft's context K/V come directly from the
# target model's own per-layer KV (draft layer i ← target layer
# target_layer_ids[i]), so the fc/hidden_norm fusion of the aux hidden
# concat is not used and is not built. Draft layer i attends into the KV
# of target_layer_ids[i]; the list length already matches num draft layers.
self.use_target_kv = bool(dflash_config.get("use_target_kv", False))
# Alternative target-KV mode: instead of REPLACING the draft's context K/V
# with the target's KV, keep the draft's k_proj/v_proj(fc(target_hidden))
# and ADD the target KV as a residual injection. Keeps fc/hidden_norm.
self.use_target_kv_inject = bool(
dflash_config.get("use_target_kv_inject", False)
)
# Third mode: fuse ALL captured target layers' K/V (like the baseline fuses
# the aux HIDDEN layers) into one [B,S,H] feature via fc, then let each draft
# layer's own k_proj/v_proj project it — i.e. the target KV cache replaces
# the aux hidden as the fc input. Keeps the baseline attention path (each
# layer sees target_kv=None; context K/V = k_proj/v_proj of the fused KV).
self.use_target_kv_fuse = bool(dflash_config.get("use_target_kv_fuse", False))
if (
sum(
[
self.use_target_kv,
self.use_target_kv_inject,
self.use_target_kv_fuse,
]
)
> 1
):
raise ValueError(
"use_target_kv (replace) / use_target_kv_inject (residual add) / "
"use_target_kv_fuse (fuse KV as fc input) are mutually exclusive; "
"enable at most one."
)
if self.use_target_kv:
self.fc = None
self.hidden_norm = None
elif self.use_target_kv_fuse:
# fc input = concat over captured layers of flattened (K, V):
# len(target_layer_ids) * 2 * num_kv_heads * head_dim.
_hd = getattr(
config, "head_dim", config.hidden_size // config.num_attention_heads
)
_kv_feat = 2 * config.num_key_value_heads * _hd
self.fc = nn.Linear(
len(self.target_layer_ids) * _kv_feat,
config.hidden_size,
bias=False,
)
self.hidden_norm = Qwen3RMSNorm(
config.hidden_size, eps=config.rms_norm_eps
)
else:
self.fc = nn.Linear(
len(self.target_layer_ids) * config.hidden_size,
config.hidden_size,
bias=False,
)
self.hidden_norm = Qwen3RMSNorm(
config.hidden_size, eps=config.rms_norm_eps
)
self.block_size = config.block_size
self.mask_token_id = dflash_config.get("mask_token_id", None)
# Optional MiMo-style learned mask embedding. When enabled, the masked
# (to-be-predicted) block positions use this trained vector instead of
# the frozen target embed_tokens(mask_token_id). It is a normal draft
# parameter (optimized by the draft optimizer and saved into the draft
# checkpoint), and is additionally exported as `mask_embedding.pt` for
# the sglang DFlash worker, which injects it at each block's masked
# slots (noise_embedding[:, 1:, :]). Toggle via
# dflash_config["use_mask_embedding"] (CLI: --use-mask-embedding).
self.use_mask_embedding = bool(dflash_config.get("use_mask_embedding", False))
if self.use_mask_embedding:
self.mask_embedding = nn.Parameter(torch.zeros(config.hidden_size))
else:
self.register_parameter("mask_embedding", None)
self._offload_fc_input_enabled = False
self.post_init()
def set_offload_fc_input_enabled(self, enabled: bool) -> None:
self._offload_fc_input_enabled = enabled
def _pack_saved_tensor_to_cpu(self, tensor: torch.Tensor):
if tensor.device.type != "cuda":
return tensor, tensor.device
return tensor.to("cpu", non_blocking=True), tensor.device
@staticmethod
def _unpack_saved_tensor_from_cpu(packed):
cpu_tensor, device = packed
return cpu_tensor.to(device, non_blocking=True)
def forward(
self,
position_ids: torch.LongTensor,
attention_mask: Optional[torch.Tensor] = None,
noise_embedding: Optional[torch.Tensor] = None,
target_hidden: Optional[torch.Tensor] = None,
target_kv: Optional[list] = None,
past_key_values: Optional[Cache] = None,
use_cache: bool = False,
**kwargs,
) -> CausalLMOutputWithPast:
hidden_states = noise_embedding
if self.mask_embedding is not None:
# Overwrite each block's masked slots (index % block_size != 0) with
# the learned mask embedding; each block's first slot keeps the real
# anchor-token embedding. Mirrors the sglang worker, which sets
# noise_embedding[:, 1:, :] = mask_embedding per block. Works for both
# training (length = n * block_size) and spec_generate (length =
# block_size), since the anchor always sits at index % block_size == 0.
seq_len = hidden_states.shape[1]
pos = torch.arange(seq_len, device=hidden_states.device)
is_mask_pos = (pos % self.block_size) != 0
hidden_states = torch.where(
is_mask_pos.view(1, seq_len, 1),
self.mask_embedding.to(hidden_states.dtype).view(1, 1, -1),
hidden_states,
)
needs_target_kv = (
self.use_target_kv or self.use_target_kv_inject or self.use_target_kv_fuse
)
if needs_target_kv:
# All three target-KV modes require the per-draft-layer target K/V.
if target_kv is None:
raise ValueError(
"use_target_kv / use_target_kv_inject / use_target_kv_fuse is "
"enabled but no target_kv was provided to "
"DFlashDraftModel.forward."
)
if len(target_kv) != len(self.layers):
raise ValueError(
f"target_kv has {len(target_kv)} entries but the draft has "
f"{len(self.layers)} layers; expected one (k, v) pair per layer."
)
if self.use_target_kv:
# REPLACE mode: the aux-hidden fusion (fc/hidden_norm) is not used.
pass
elif self.use_target_kv_fuse:
# FUSE mode: build the context feature from the captured target K/V
# instead of the aux hidden — concat every layer's flattened (K, V),
# fc-fuse to [B,S,H], then the baseline per-layer k_proj/v_proj project
# it (attention runs the baseline path, target_kv=None). NOTE the fused
# K carries the target's RoPE and the baseline re-applies the draft RoPE
# to k_ctx (double rotation on the K component); it is a fixed function
# of position that fc/k_proj learn around, but is a known subtlety.
kv_feats = []
for k_l, v_l in target_kv: # each [B, S, num_kv_heads, head_dim]
b, s = k_l.shape[:2]
kv_feats.append(k_l.reshape(b, s, -1))
kv_feats.append(v_l.reshape(b, s, -1))
kv_concat = torch.cat(kv_feats, dim=-1) # [B, S, K*2*nkv*hd]
target_hidden = self.hidden_norm(self.fc(kv_concat))
elif self._offload_fc_input_enabled:
# Offload tensors saved by fc/norm autograd so later attention layers
# can reuse the GPU memory; they are copied back during backward.
with torch.autograd.graph.saved_tensors_hooks(
self._pack_saved_tensor_to_cpu,
self._unpack_saved_tensor_from_cpu,
):
target_hidden = self.hidden_norm(self.fc(target_hidden))
else:
# baseline AND inject: fc-fuse the aux hidden into the context feature.
target_hidden = self.hidden_norm(self.fc(target_hidden))
position_embeddings = self.rotary_emb(hidden_states, position_ids)
# target_kv is consumed by the attention (replace/inject) only; in fuse mode
# it has already been folded into target_hidden above, so the attention runs
# the plain baseline path.
attn_uses_target_kv = self.use_target_kv or self.use_target_kv_inject
for layer_idx, layer in enumerate(self.layers):
layer_attention_mask = attention_mask
if isinstance(attention_mask, dict):
layer_attention_mask = (
attention_mask["sliding"]
if layer.self_attn.sliding_window is not None
else attention_mask["full"]
)
hidden_states = layer(
hidden_states=hidden_states,
target_hidden=None if self.use_target_kv else target_hidden,
attention_mask=layer_attention_mask,
position_ids=position_ids,
past_key_value=past_key_values,
use_cache=use_cache,
position_embeddings=position_embeddings,
target_kv=target_kv[layer_idx] if attn_uses_target_kv else None,
**kwargs,
)
return self.norm(hidden_states)
@torch.inference_mode()
def spec_generate(
self,
target: nn.Module,
input_ids: torch.LongTensor,
max_new_tokens: int,
stop_token_ids: list[int],
temperature: float,
):
self.eval()
num_input_tokens = input_ids.shape[1]
max_length = num_input_tokens + max_new_tokens
block_size = self.block_size
output_ids = torch.full(
(1, max_length + block_size),
self.mask_token_id,
dtype=torch.long,
device=target.device,
)
position_ids = torch.arange(
output_ids.shape[1], device=target.device
).unsqueeze(0)
past_key_values_target = DynamicCache()
past_key_values_draft = DynamicCache()
# Prefill stage
output = target(
input_ids,
position_ids=position_ids[:, :num_input_tokens],
past_key_values=past_key_values_target,
use_cache=True,
logits_to_keep=1,
output_hidden_states=True,
)
output_ids[:, :num_input_tokens] = input_ids
output_ids[:, num_input_tokens : num_input_tokens + 1] = sample(
output.logits, temperature
)
target_hidden = extract_context_feature(
output.hidden_states, self.target_layer_ids
)
# Decode stage
acceptance_lengths = []
start = input_ids.shape[1]
while start < max_length:
block_output_ids = output_ids[:, start : start + block_size].clone()
block_position_ids = position_ids[:, start : start + block_size]
noise_embedding = target.model.embed_tokens(block_output_ids)
draft_logits = target.lm_head(
self(
target_hidden=target_hidden,
noise_embedding=noise_embedding,
position_ids=position_ids[
:, past_key_values_draft.get_seq_length() : start + block_size
],
past_key_values=past_key_values_draft,
use_cache=True,
is_causal=False,
)[:, -block_size + 1 :, :]
)
past_key_values_draft.crop(start)
block_output_ids[:, 1:] = sample(draft_logits)
output = target(
block_output_ids,
position_ids=block_position_ids,
past_key_values=past_key_values_target,
use_cache=True,
output_hidden_states=True,
)
posterior = sample(output.logits, temperature)
acceptance_length = (
(block_output_ids[:, 1:] == posterior[:, :-1])
.cumprod(dim=1)
.sum(dim=1)[0]
.item()
)
output_ids[:, start : start + acceptance_length + 1] = block_output_ids[
:, : acceptance_length + 1
]
output_ids[:, start + acceptance_length + 1] = posterior[
:, acceptance_length
]
start += acceptance_length + 1
past_key_values_target.crop(start)
target_hidden = extract_context_feature(
output.hidden_states, self.target_layer_ids
)[:, : acceptance_length + 1, :]
acceptance_lengths.append(acceptance_length + 1)
if stop_token_ids is not None and any(
stop_token_id in output_ids[:, num_input_tokens:]
for stop_token_id in stop_token_ids
):
break
output_ids = output_ids[:, :max_length]
output_ids = output_ids[:, output_ids[0] != self.mask_token_id]
if stop_token_ids is not None:
stop_token_ids = torch.tensor(stop_token_ids, device=output_ids.device)
stop_token_indices = torch.isin(
output_ids[0][num_input_tokens:], stop_token_ids
).nonzero(as_tuple=True)[0]
if stop_token_indices.numel() > 0:
output_ids = output_ids[
:, : num_input_tokens + stop_token_indices[0] + 1
]
return output_ids