拆分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:
2026-07-17 11:23:11 +08:00
parent c57a9bf9e6
commit 7856f5bf4c
10 changed files with 509 additions and 439 deletions
+84
View File
@@ -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