CoolFace
Modelpublic

tencent/HunyuanOCR

sourceHugging Faceotherupdated 28d agoView on Hugging Face
823likes733kdownloads
dflash.py188 linesDownload Raw Back to dflash
1from typing import Optional, Callable2from typing_extensions import Unpack, Tuple3import torch4from torch import nn5from transformers.models.qwen3.modeling_qwen3 import (6    Qwen3RMSNorm,7    Qwen3RotaryEmbedding,8    Qwen3Config,9    Qwen3PreTrainedModel,10    Qwen3MLP,11    GradientCheckpointingLayer,12    FlashAttentionKwargs,13    rotate_half,14    eager_attention_forward,15    ALL_ATTENTION_FUNCTIONS,16)17from transformers.modeling_outputs import CausalLMOutputWithPast18from transformers.cache_utils import Cache19 20def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1):21    cos = cos.unsqueeze(unsqueeze_dim)22    sin = sin.unsqueeze(unsqueeze_dim)23    q_len = q.size(-2)24    q_embed = (q * cos[..., -q_len:, :]) + (rotate_half(q) * sin[..., -q_len:, :])25    k_embed = (k * cos) + (rotate_half(k) * sin)26    return q_embed, k_embed27 28class Qwen3DFlashAttention(nn.Module):29    """Multi-headed attention from 'Attention Is All You Need' paper"""30 31    def __init__(self, config: Qwen3Config, layer_idx: int):32        super().__init__()33        self.config = config34        self.layer_idx = layer_idx35        self.head_dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)36        self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads37        self.scaling = self.head_dim**-0.538        self.attention_dropout = config.attention_dropout39        self.is_causal = False  40        self.q_proj = nn.Linear(41            config.hidden_size, config.num_attention_heads * self.head_dim, bias=config.attention_bias42        )43        self.k_proj = nn.Linear(44            config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias45        )46        self.v_proj = nn.Linear(47            config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias48        )49        self.o_proj = nn.Linear(50            config.num_attention_heads * self.head_dim, config.hidden_size, bias=config.attention_bias51        )52        self.q_norm = Qwen3RMSNorm(self.head_dim, eps=config.rms_norm_eps)53        self.k_norm = Qwen3RMSNorm(self.head_dim, eps=config.rms_norm_eps)54        self.sliding_window = config.sliding_window if config.layer_types[layer_idx] == "sliding_attention" else None55 56    def forward(57        self,58        hidden_states: torch.Tensor,59        target_hidden: torch.Tensor,60        position_embeddings: tuple[torch.Tensor, torch.Tensor],61        attention_mask: Optional[torch.Tensor],62        past_key_values: Optional[Cache] = None,63        cache_position: Optional[torch.LongTensor] = None,64        **kwargs: Unpack[FlashAttentionKwargs],65    ) -> tuple[torch.Tensor, Optional[torch.Tensor]]:66        bsz, q_len = hidden_states.shape[:-1]67        ctx_len = target_hidden.shape[1]68        q = self.q_proj(hidden_states)69        q = q.view(bsz, q_len, -1, self.head_dim)70        q = self.q_norm(q).transpose(1, 2)71        k_ctx = self.k_proj(target_hidden)72        k_noise = self.k_proj(hidden_states)73        v_ctx = self.v_proj(target_hidden)74        v_noise = self.v_proj(hidden_states)75        k = torch.cat([k_ctx, k_noise], dim=1).view(bsz, ctx_len + q_len, -1, self.head_dim)76        v = torch.cat([v_ctx, v_noise], dim=1).view(bsz, ctx_len + q_len, -1, self.head_dim)77        k = self.k_norm(k).transpose(1, 2)78        v = v.transpose(1, 2)79        cos, sin = position_embeddings80        q, k = apply_rotary_pos_emb(q, k, cos, sin)81        if past_key_values is not None:82            cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position}83            k, v = past_key_values.update(k, v, self.layer_idx, cache_kwargs)84        attn_fn: Callable = eager_attention_forward85        if self.config._attn_implementation != "eager":86            attn_fn = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]87        attn_output, attn_weights = attn_fn(88            self,89            q,90            k,91            v,92            attention_mask,93            dropout=0.0 if not self.training else self.attention_dropout,94            scaling=self.scaling,95            sliding_window=self.sliding_window,96            **kwargs,97        )98        attn_output = attn_output.reshape(bsz, q_len, -1)99        attn_output = self.o_proj(attn_output)100        return attn_output, attn_weights101 102class Qwen3DFlashDecoderLayer(GradientCheckpointingLayer):103    def __init__(self, config: Qwen3Config, layer_idx: int):104        super().__init__()105        self.hidden_size = config.hidden_size106        self.self_attn = Qwen3DFlashAttention(config=config, layer_idx=layer_idx)107        self.mlp = Qwen3MLP(config)108        self.input_layernorm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)109        self.post_attention_layernorm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)110 111    def forward(112        self,113        target_hidden: Optional[torch.Tensor] = None,114        hidden_states: Optional[torch.Tensor] = None,115        attention_mask: Optional[torch.Tensor] = None,116        position_ids: Optional[torch.LongTensor] = None,117        past_key_value: Optional[Cache] = None,118        output_attentions: Optional[bool] = False,119        use_cache: Optional[bool] = False,120        cache_position: Optional[torch.LongTensor] = None,121        position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,  # necessary, but kept here for BC122        **kwargs: Unpack[FlashAttentionKwargs],123    ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:124        residual = hidden_states125        hidden_states = self.input_layernorm(hidden_states)126        hidden_states = self.self_attn(127            hidden_states=hidden_states,128            target_hidden=target_hidden,129            attention_mask=attention_mask,130            position_ids=position_ids,131            past_key_values=past_key_value,132            output_attentions=output_attentions,133            use_cache=use_cache,134            cache_position=cache_position,135            position_embeddings=position_embeddings,136            **kwargs,137        )[0]138        hidden_states = residual + hidden_states139        residual = hidden_states140        hidden_states = self.post_attention_layernorm(hidden_states)141        hidden_states = self.mlp(hidden_states)142        hidden_states = residual + hidden_states143        return hidden_states144 145class DFlashDraftModel(Qwen3PreTrainedModel):146    config_class = Qwen3Config147    _no_split_modules = ["Qwen3DFlashDecoderLayer"]148 149    def __init__(self, config) -> None:150        super().__init__(config)151        self.config = config152        self.layers = nn.ModuleList(153            [Qwen3DFlashDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]154        )155        self.target_layer_ids = self.config.dflash_config.get("target_layer_ids", None)156        self.norm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)157        self.rotary_emb = Qwen3RotaryEmbedding(config)158        self.fc = nn.Linear(len(self.target_layer_ids) * config.hidden_size, config.hidden_size, bias=False)159        self.hidden_norm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)160        self.block_size = config.block_size161        self.mask_token_id = self.config.dflash_config.get("mask_token_id", None)162        self.post_init()163 164    def forward(165        self,166        position_ids: torch.LongTensor,167        attention_mask: Optional[torch.Tensor] = None,168        noise_embedding: Optional[torch.Tensor] = None,169        target_hidden: Optional[torch.Tensor] = None,170        past_key_values: Optional[Cache] = None,171        use_cache: bool = False,172        **kwargs,173    ) -> CausalLMOutputWithPast:174        hidden_states = noise_embedding175        target_hidden = self.hidden_norm(self.fc(target_hidden))176        position_embeddings = self.rotary_emb(hidden_states, position_ids)177        for layer in self.layers:178            hidden_states = layer(179                hidden_states=hidden_states,180                target_hidden=target_hidden,181                attention_mask=attention_mask,182                position_ids=position_ids,183                past_key_value=past_key_values,184                use_cache=use_cache,185                position_embeddings=position_embeddings,186                **kwargs,187            )188        return self.norm(hidden_states)