拆分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:
@@ -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
|
||||
Reference in New Issue
Block a user