Edge0/ARK-ASR-3B
1049.3k
1from typing import Any, Optional, Tuple, Union2 3import torch4from torch import Tensor, nn5from torch.nn.functional import scaled_dot_product_attention6from transformers import WhisperConfig7from transformers.modeling_outputs import BaseModelOutputWithPastAndCrossAttentions8from transformers.models.whisper.modeling_whisper import WhisperEncoder, WhisperEncoderLayer9from transformers.utils import logging10 11logger = logging.get_logger(__name__)12 13# ==========================================14# 1. Rotary Embedding 核心组件15# ==========================================16 17class RotaryEmbedding(nn.Module):18 def __init__(self, dim, rope_ratio=1):19 super().__init__()20 self.dim = dim21 self.rope_ratio = rope_ratio22 23 @torch.no_grad()24 def get_emb(self, seq_len: int, dtype: torch.dtype, device: torch.device, base: int = 10000):25 """生成 RoPE 缓存"""26 base = base * self.rope_ratio27 # 计算频率 theta28 inv_freq = 1.0 / (base ** (torch.arange(0, self.dim, 2, dtype=torch.float, device=device) / self.dim))29 30 # 生成位置索引31 t = torch.arange(seq_len, device=device, dtype=torch.float)32 freqs = torch.outer(t, inv_freq) # [seq_len, dim/2]33 34 # 构造 cos 和 sin 缓存35 # 形状: [seq_len, dim/2, 2]36 emb = torch.stack([torch.cos(freqs), torch.sin(freqs)], dim=-1)37 38 if dtype in (torch.float16, torch.bfloat16):39 emb = emb.to(dtype)40 return emb41 42def apply_rotary_pos_emb(x: torch.Tensor, rope_cache: torch.Tensor) -> torch.Tensor:43 """44 x: [batch, num_heads, seq_len, head_dim]45 rope_cache: [1, seq_len, dim/2, 2]46 """47 b, nh, sq, hd = x.shape48 rot_dim = rope_cache.shape[-2] * 249 50 # 将 x 分为旋转部分和不旋转部分51 x_rot, x_pass = x[..., :rot_dim], x[..., rot_dim:]52 53 # 调整 x_rot 形状以匹配 rope_cache: [b, nh, sq, rot_dim/2, 2]54 x_shaped = x_rot.reshape(b, nh, sq, rot_dim // 2, 2)55 56 # 计算旋转: (a+bi)(c+di) = (ac-bd) + (ad+bc)i57 cos = rope_cache[..., 0] # [1, sq, rot_dim/2]58 sin = rope_cache[..., 1] # [1, sq, rot_dim/2]59 60 # 增加 head 维度61 cos = cos.unsqueeze(1) # [1, 1, sq, rot_dim/2]62 sin = sin.unsqueeze(1) # [1, 1, sq, rot_dim/2]63 64 x_out = torch.stack([65 x_shaped[..., 0] * cos - x_shaped[..., 1] * sin,66 x_shaped[..., 1] * cos + x_shaped[..., 0] * sin67 ], dim=-1)68 69 x_out = x_out.flatten(3) # 合并最后两维到 rot_dim70 return torch.cat([x_out, x_pass], dim=-1)71 72# ==========================================73# 2. 基于 SDPA 的 RoPE Attention74# ==========================================75 76class WhisperRoPESdpaAttention(nn.Module):77 """78 使用 PyTorch 原生 scaled_dot_product_attention 替代 WhisperFlashAttention2。79 """80 def __init__(self, config: WhisperConfig, embed_dim: int, num_heads: int, dropout: float = 0.0):81 super().__init__()82 self.config = config83 self.embed_dim = embed_dim84 self.num_heads = num_heads85 self.dropout = dropout86 self.head_dim = embed_dim // num_heads87 88 # Whisper 标准投影层89 self.q_proj = nn.Linear(embed_dim, embed_dim)90 self.k_proj = nn.Linear(embed_dim, embed_dim, bias=False)91 self.v_proj = nn.Linear(embed_dim, embed_dim)92 self.out_proj = nn.Linear(embed_dim, embed_dim)93 94 self.is_causal = False95 96 def forward(97 self,98 hidden_states: torch.Tensor,99 attention_mask: Optional[torch.Tensor] = None,100 layer_head_mask: Optional[torch.Tensor] = None,101 output_attentions: bool = False,102 rotary_pos_emb: Optional[torch.Tensor] = None,103 ) -> Tuple[torch.Tensor, Optional[torch.Tensor], None]:104 105 bsz, q_len, _ = hidden_states.size()106 107 # 1. 投影映射108 query_states = self.q_proj(hidden_states)109 key_states = self.k_proj(hidden_states)110 value_states = self.v_proj(hidden_states)111 112 # 2. 变形为 [batch, heads, seq, dim] 并确保内存连续113 query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()114 key_states = key_states.view(bsz, -1, self.num_heads, self.head_dim).transpose(1, 2).contiguous()115 value_states = value_states.view(bsz, -1, self.num_heads, self.head_dim).transpose(1, 2).contiguous()116 117 # 3. 应用 RoPE118 if rotary_pos_emb is not None:119 query_states = apply_rotary_pos_emb(query_states, rotary_pos_emb)120 key_states = apply_rotary_pos_emb(key_states, rotary_pos_emb)121 122 # 4. 数据类型对齐 (处理 fp32 LayerNorm 带来的类型不匹配)123 target_dtype = self.q_proj.weight.dtype124 query_states = query_states.to(target_dtype)125 key_states = key_states.to(target_dtype)126 value_states = value_states.to(target_dtype)127 128 # 5. SDPA 计算 (关键:不要手动乘以 scaling, SDPA 内部会自动处理)129 # 注意: 如果传入了 4D attention_mask,SDPA 会正确应用它130 attn_output = scaled_dot_product_attention(131 query_states,132 key_states,133 value_states,134 attn_mask=attention_mask,135 dropout_p=self.dropout if self.training else 0.0,136 is_causal=self.is_causal,137 )138 139 # 6. 恢复形状并输出投影140 attn_output = attn_output.transpose(1, 2).contiguous()141 attn_output = attn_output.reshape(bsz, q_len, self.embed_dim)142 attn_output = self.out_proj(attn_output)143 144 return attn_output, None, None145 146# ==========================================147# 3. 封装好的 Encoder 层和 Encoder148# ==========================================149 150class WhisperSpecialEncoderLayer(WhisperEncoderLayer):151 def __init__(self, config: WhisperConfig):152 super().__init__(config)153 # 替换 Self-Attention 为我们的 RoPE SDPA 版本154 self.self_attn = WhisperRoPESdpaAttention(155 config=config,156 embed_dim=self.embed_dim,157 num_heads=config.encoder_attention_heads,158 dropout=config.attention_dropout,159 )160 161 def forward(162 self,163 hidden_states: torch.Tensor,164 attention_mask: Optional[torch.Tensor] = None,165 layer_head_mask: Optional[torch.Tensor] = None,166 output_attentions: bool = False,167 rotary_pos_emb: Optional[torch.Tensor] = None,168 position_ids: Optional[torch.Tensor] = None,169 ) -> Tuple[torch.Tensor, Any]:170 171 residual = hidden_states172 hidden_states = self.self_attn_layer_norm(hidden_states)173 174 hidden_states, attn_weights, _ = self.self_attn(175 hidden_states=hidden_states,176 attention_mask=attention_mask,177 layer_head_mask=layer_head_mask,178 output_attentions=output_attentions,179 rotary_pos_emb=rotary_pos_emb,180 )181 182 hidden_states = nn.functional.dropout(hidden_states, p=self.dropout, training=self.training)183 hidden_states = residual + hidden_states184 185 residual = hidden_states186 hidden_states = self.final_layer_norm(hidden_states)187 hidden_states = self.activation_fn(self.fc1(hidden_states))188 hidden_states = nn.functional.dropout(hidden_states, p=self.activation_dropout, training=self.training)189 hidden_states = self.fc2(hidden_states)190 hidden_states = nn.functional.dropout(hidden_states, p=self.dropout, training=self.training)191 hidden_states = residual + hidden_states192 193 if hidden_states.dtype == torch.float16:194 clamp_value = torch.finfo(hidden_states.dtype).max - 1000195 hidden_states = torch.clamp(hidden_states, min=-clamp_value, max=clamp_value)196 197 return (hidden_states, None) # 保持与 Whisper 接口一致的 tuple 长度198 199class WhisperSpecialEncoder(WhisperEncoder):200 def __init__(self, config: WhisperConfig, use_rope=True, rope_ratio=1):201 super().__init__(config)202 self.use_rope = use_rope203 # 覆盖父类的层列表204 self.layers = nn.ModuleList(205 [WhisperSpecialEncoderLayer(config) for _ in range(config.encoder_layers)]206 )207 208 if use_rope:209 # 计算 RoPE 维度: 通常是 head_dim 的一部分210 head_dim = config.d_model // config.encoder_attention_heads211 self.rotary_embedding = RotaryEmbedding(head_dim // 2, rope_ratio)212 213 def forward(214 self,215 input_features,216 attention_mask=None,217 head_mask=None,218 output_attentions=None,219 output_hidden_states=None,220 return_dict=None,221 position_ids=None,222 ):223 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions224 output_hidden_states = output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states225 return_dict = return_dict if return_dict is not None else self.config.use_return_dict226 227 # Whisper 卷积特征提取228 inputs_embeds = nn.functional.gelu(self.conv1(input_features))229 inputs_embeds = nn.functional.gelu(self.conv2(inputs_embeds))230 inputs_embeds = inputs_embeds.permute(0, 2, 1) # [B, T_down, D]231 232 if self.use_rope:233 # 生成旋转编码缓存234 rotary_embs = self.rotary_embedding.get_emb(235 seq_len=inputs_embeds.shape[1],236 dtype=inputs_embeds.dtype,237 device=inputs_embeds.device238 )239 # 形状调整为 [1, seq_len, dim/2, 2] 以便广播240 rotary_embs = rotary_embs.unsqueeze(0)241 hidden_states = inputs_embeds 242 else:243 rotary_embs = None244 # 回退到绝对位置编码245 embed_pos = self.embed_positions.weight[:inputs_embeds.shape[1]]246 hidden_states = inputs_embeds + embed_pos247 248 hidden_states = nn.functional.dropout(hidden_states, p=self.dropout, training=self.training)249 250 encoder_states = () if output_hidden_states else None251 all_attentions = () if output_attentions else None252 253 for idx, encoder_layer in enumerate(self.layers):254 if output_hidden_states:255 encoder_states = encoder_states + (hidden_states,)256 257 if self.gradient_checkpointing and self.training:258 layer_outputs = self._gradient_checkpointing_func(259 encoder_layer.__call__,260 hidden_states,261 None, # attention_mask262 (head_mask[idx] if head_mask is not None else None),263 output_attentions,264 rotary_embs,265 position_ids,266 )267 else:268 layer_outputs = encoder_layer(269 hidden_states,270 attention_mask=None,271 layer_head_mask=(head_mask[idx] if head_mask is not None else None),272 output_attentions=output_attentions,273 rotary_pos_emb=rotary_embs,274 position_ids=position_ids,275 )276 277 hidden_states = layer_outputs[0]278 279 if output_attentions:280 all_attentions = all_attentions + (layer_outputs[2],)281 282 hidden_states = self.layer_norm(hidden_states)283 if output_hidden_states:284 encoder_states = encoder_states + (hidden_states,)285 286 if not return_dict:287 return tuple(v for v in [hidden_states, encoder_states, all_attentions] if v is not None)288 289 return BaseModelOutputWithPastAndCrossAttentions(290 last_hidden_state=hidden_states,291 hidden_states=encoder_states,292 attentions=all_attentions,293 )