Files
transformer_study/transformer/inference.py
T
XingfenD 7856f5bf4c 拆分transformer.py为独立模块
按Transformer架构组件拆分:
- utils.py: 激活函数和掩码工具
- layers.py: 基础层(Linear, LayerNorm)
- attention.py: 多头注意力机制
- feedforward.py: 前馈网络
- block.py: Transformer编码器块
- model.py: SimpleTransformer模型
- inference.py: 推理引擎
- 删除旧的transformer.py文件
2026-07-17 11:23:11 +08:00

128 lines
4.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
推理引擎模块
包含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