拆分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
+128
View File
@@ -0,0 +1,128 @@
"""
推理引擎模块
包含InferenceEngine类,管理prefill和decode阶段
"""
import numpy as np
from .utils import softmax, create_look_ahead_mask
class InferenceEngine:
"""推理引擎:管理prefill和decode阶段"""
def __init__(self, model):
"""
初始化推理引擎
model: SimpleTransformer模型实例
"""
self.model = model
self.generated_tokens = [] # 已生成的token列表
self.current_position = 0 # 当前位置
def reset(self):
"""重置推理状态"""
self.model.clear_cache()
self.generated_tokens = []
self.current_position = 0
def prefill(self, input_ids):
"""
Prefill阶段:处理完整的输入序列
input_ids: 输入token序列 [batch_size, seq_len]
返回: 下一个token的logits [batch_size, vocab_size]
"""
self.reset()
batch_size = input_ids.shape[0]
seq_len = input_ids.shape[1]
# 创建前瞻掩码(防止看到未来信息)
mask = create_look_ahead_mask(seq_len)
# 执行prefill前向传播(使用KV cache)
output = self.model.prefill(input_ids, mask)
# 获取最后一个位置的输出(用于预测下一个token)
last_output = output[:, -1, :] # [batch_size, d_model]
# 简单的logits计算:使用嵌入矩阵的转置作为输出投影
# logits = last_output @ embedding.T
logits = last_output @ self.model.embedding.T # [batch_size, vocab_size]
# 更新状态
self.current_position = seq_len
return logits
def decode_step(self, input_token):
"""
Decode阶段:处理单个新token
input_token: 新token [batch_size, 1]
返回: 下一个token的logits [batch_size, vocab_size]
"""
# 执行decode前向传播(使用KV cache)
output = self.model.decode(input_token, self.current_position)
# 获取输出(单token,所以直接取[:, 0, :])
last_output = output[:, 0, :] # [batch_size, d_model]
# 计算logits
logits = last_output @ self.model.embedding.T # [batch_size, vocab_size]
# 更新位置
self.current_position += 1
return logits
def generate(self, input_ids, max_new_tokens=50, temperature=1.0):
"""
自回归生成文本
input_ids: 输入token序列 [batch_size, seq_len]
max_new_tokens: 最大生成token数
temperature: 温度参数(控制随机性)
返回: 完整的生成序列 [batch_size, seq_len + max_new_tokens]
"""
batch_size = input_ids.shape[0]
# 保存原始输入
generated = input_ids.copy()
# Prefill阶段
logits = self.prefill(input_ids)
# 采样第一个生成的token
next_token = self._sample(logits, temperature)
generated = np.concatenate([generated, next_token], axis=1)
# Decode阶段:逐个生成token
for i in range(max_new_tokens - 1):
logits = self.decode_step(next_token)
next_token = self._sample(logits, temperature)
generated = np.concatenate([generated, next_token], axis=1)
# 检查是否生成了结束符(这里用0作为结束符)
if np.all(next_token == 0):
break
return generated
def _sample(self, logits, temperature=1.0):
"""
从logits中采样
logits: [batch_size, vocab_size]
temperature: 温度参数
返回: 采样的token [batch_size, 1]
"""
# 应用温度缩放
logits = logits / temperature
# 计算概率分布
probs = softmax(logits, axis=-1)
# 按概率采样
batch_size = logits.shape[0]
next_tokens = np.zeros((batch_size, 1), dtype=int)
for i in range(batch_size):
# 使用numpy的random.choice采样
next_tokens[i, 0] = np.random.choice(logits.shape[1], p=probs[i])
return next_tokens