""" 推理引擎模块 包含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