Files
transformer_study/transformer/model.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

117 lines
4.3 KiB
Python

"""
Transformer模型模块
包含SimpleTransformer类
"""
import numpy as np
from .block import TransformerBlock
class SimpleTransformer:
"""简化的Transformer编码器模型(支持prefill/decode分离)"""
def __init__(self, vocab_size, d_model, num_heads, num_layers, d_ff, max_seq_len):
self.d_model = d_model
self.vocab_size = vocab_size
# 词嵌入层
self.embedding = np.random.randn(vocab_size, d_model) * 0.01
# 位置编码
self.positional_encoding = self._create_positional_encoding(max_seq_len, d_model)
# Transformer编码器层堆叠
self.transformer_blocks = [
TransformerBlock(d_model, num_heads, d_ff) for _ in range(num_layers)
]
def _create_positional_encoding(self, max_seq_len, d_model):
"""
创建正弦余弦位置编码
使用不同频率的正弦和余弦函数生成位置信息
"""
pe = np.zeros((max_seq_len, d_model))
position = np.arange(0, max_seq_len).reshape(-1, 1).astype(float)
div_term = np.exp(np.arange(0, d_model, 2).astype(float) * -(np.log(10000.0) / d_model))
pe[:, 0::2] = np.sin(position * div_term) # 偶数维度用正弦
pe[:, 1::2] = np.cos(position * div_term) # 奇数维度用余弦
return pe
def clear_cache(self):
"""清除所有层的KV cache"""
for block in self.transformer_blocks:
block.clear_cache()
def prefill(self, x, mask=None):
"""
Prefill阶段:处理完整的输入序列
x: 输入序列 [batch_size, seq_len]
mask: 注意力掩码
返回: 输出张量 [batch_size, seq_len, d_model]
"""
seq_len = x.shape[1]
# 词嵌入 + 缩放
x = self.embedding[x] * np.sqrt(self.d_model)
# 加上位置编码
x = x + self.positional_encoding[:seq_len, :]
# 通过所有Transformer块,使用KV cache
for block in self.transformer_blocks:
x = block.forward(x, mask, use_cache=True)
return x
def decode(self, x, position):
"""
Decode阶段:处理单个新token
x: 新token [batch_size, 1]
position: 当前token在序列中的位置
返回: 输出张量 [batch_size, 1, d_model]
"""
# 词嵌入 + 缩放
x = self.embedding[x] * np.sqrt(self.d_model)
# 加上位置编码(使用当前位置)
x = x + self.positional_encoding[position:position+1, :]
# 通过所有Transformer块,使用KV cache
for block in self.transformer_blocks:
x = block.forward(x, use_cache=True)
return x
def forward(self, x, mask=None):
"""
前向传播(不使用KV cache,用于训练)
x: 输入序列 [batch_size, seq_len]
mask: 注意力掩码
"""
seq_len = x.shape[1]
# 词嵌入 + 缩放
x = self.embedding[x] * np.sqrt(self.d_model)
# 加上位置编码
x = x + self.positional_encoding[:seq_len, :]
# 通过所有Transformer块
for block in self.transformer_blocks:
x = block.forward(x, mask)
return x
def count_parameters(self):
"""统计模型参数数量"""
count = 0
count += self.embedding.size # 嵌入层参数
for block in self.transformer_blocks:
# 注意力层参数
count += block.attention.W_q.weight.size + block.attention.W_q.bias.size
count += block.attention.W_k.weight.size + block.attention.W_k.bias.size
count += block.attention.W_v.weight.size + block.attention.W_v.bias.size
count += block.attention.W_o.weight.size + block.attention.W_o.bias.size
# 前馈网络参数
count += block.feed_forward.linear1.weight.size + block.feed_forward.linear1.bias.size
count += block.feed_forward.linear2.weight.size + block.feed_forward.linear2.bias.size
# 层归一化参数
count += block.norm1.gamma.size + block.norm1.beta.size
count += block.norm2.gamma.size + block.norm2.beta.size
return count