""" 多头注意力机制模块 包含MultiHeadAttention类,支持KV cache """ import numpy as np from .layers import Linear from .utils import softmax class MultiHeadAttention: """多头注意力机制(支持KV cache)""" def __init__(self, d_model, num_heads): self.d_model = d_model # 模型维度 self.num_heads = num_heads # 注意力头数 self.d_k = d_model // num_heads # 每个头的维度 # 定义Q、K、V的线性变换层 self.W_q = Linear(d_model, d_model) # Query变换 self.W_k = Linear(d_model, d_model) # Key变换 self.W_v = Linear(d_model, d_model) # Value变换 self.W_o = Linear(d_model, d_model) # 输出变换 # KV cache:用于推理时缓存历史Key和Value self.k_cache = None # 缓存的Key [batch_size, num_heads, seq_len, d_k] self.v_cache = None # 缓存的Value [batch_size, num_heads, seq_len, d_k] def clear_cache(self): """清除KV cache""" self.k_cache = None self.v_cache = None def forward(self, Q, K, V, mask=None, use_cache=False): """ 前向传播 Q: Query张量 [batch_size, seq_len, d_model] K: Key张量 [batch_size, seq_len, d_model] V: Value张量 [batch_size, seq_len, d_model] mask: 注意力掩码 use_cache: 是否使用KV cache(推理时设为True) """ batch_size = Q.shape[0] seq_len = Q.shape[1] # 线性变换并分割成多头 Q = self.W_q.forward(Q).reshape(batch_size, seq_len, self.num_heads, self.d_k).transpose(0, 2, 1, 3) K = self.W_k.forward(K).reshape(batch_size, seq_len, self.num_heads, self.d_k).transpose(0, 2, 1, 3) V = self.W_v.forward(V).reshape(batch_size, seq_len, self.num_heads, self.d_k).transpose(0, 2, 1, 3) # KV cache逻辑 if use_cache: if self.k_cache is None: # Prefill阶段:首次计算,缓存完整的K和V self.k_cache = K self.v_cache = V else: # Decode阶段:拼接历史缓存和新的K、V self.k_cache = np.concatenate([self.k_cache, K], axis=2) self.v_cache = np.concatenate([self.v_cache, V], axis=2) # 使用完整的K和V进行注意力计算 K = self.k_cache V = self.v_cache # 计算缩放点积注意力分数 # Q: [batch, heads, q_len, d_k] # K: [batch, heads, kv_len, d_k](decode时kv_len > q_len) attn_scores = Q @ K.transpose(0, 1, 3, 2) / np.sqrt(self.d_k) # 应用掩码(如果提供) if mask is not None: attn_scores = np.where(mask == 0, -1e9, attn_scores) # Softmax得到注意力权重 attn_probs = softmax(attn_scores, axis=-1) # 加权求和 attn_output = attn_probs @ V # 拼接多头输出 attn_output = attn_output.transpose(0, 2, 1, 3).reshape(batch_size, seq_len, self.d_model) # 最终线性变换 output = self.W_o.forward(attn_output) return output