tencent/HunyuanOCR
823733k
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)