按Transformer架构组件拆分: - utils.py: 激活函数和掩码工具 - layers.py: 基础层(Linear, LayerNorm) - attention.py: 多头注意力机制 - feedforward.py: 前馈网络 - block.py: Transformer编码器块 - model.py: SimpleTransformer模型 - inference.py: 推理引擎 - 删除旧的transformer.py文件
117 lines
4.3 KiB
Python
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 |