""" 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