Files
transformer_study/transformer/attention.py
T
XingfenD 7856f5bf4c 拆分transformer.py为独立模块
按Transformer架构组件拆分:
- utils.py: 激活函数和掩码工具
- layers.py: 基础层(Linear, LayerNorm)
- attention.py: 多头注意力机制
- feedforward.py: 前馈网络
- block.py: Transformer编码器块
- model.py: SimpleTransformer模型
- inference.py: 推理引擎
- 删除旧的transformer.py文件
2026-07-17 11:23:11 +08:00

84 lines
3.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
多头注意力机制模块
包含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