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