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