拆分transformer.py为独立模块
按Transformer架构组件拆分: - utils.py: 激活函数和掩码工具 - layers.py: 基础层(Linear, LayerNorm) - attention.py: 多头注意力机制 - feedforward.py: 前馈网络 - block.py: Transformer编码器块 - model.py: SimpleTransformer模型 - inference.py: 推理引擎 - 删除旧的transformer.py文件
This commit is contained in:
@@ -0,0 +1,84 @@
|
||||
"""
|
||||
多头注意力机制模块
|
||||
包含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
|
||||
Reference in New Issue
Block a user