Download dflash.py from NaiveAI/Naive-N0.5-Flash-FP8-Draft: direct link, hf CLI and curl.
- Browser
- Download file 27.1 kB
-
https://huggingface.co/NaiveAI/Naive-N0.5-Flash-FP8-Draft/resolve/b2b8ee9f5d6b3fd1dfba113d3a363138e37c83b0/dflash.py
- Command line
-
hf download hf://NaiveAI/Naive-N0.5-Flash-FP8-Draft@b2b8ee9f5d6b3fd1dfba113d3a363138e37c83b0/dflash.py
-
curl -L -o dflash.py https://huggingface.co/NaiveAI/Naive-N0.5-Flash-FP8-Draft/resolve/b2b8ee9f5d6b3fd1dfba113d3a363138e37c83b0/dflash.py
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 | |
| 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) | |
| 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 | |