""" 霜云(Shimokumo) - 推理引擎模块 实现自回归文本生成,支持: - Top-K / Top-P采样 - KV Cache增量推理优化 - 流式生成(yield逐token输出) - 多轮对话管理 - FP16混合精度推理 """ import time from dataclasses import dataclass, field from typing import Dict, Generator, List, Optional, Tuple import torch import torch.nn.functional as F from config import ShimokumoConfig from model.shimokumo_model import ShimokumoModel from model.tokenizer import ShimokumoTokenizer from utils.logger import get_logger logger = get_logger("Shimokumo.Inference") @dataclass class DialogueMessage: """对话消息数据类""" role: str # "user", "assistant", "system" content: str timestamp: float = field(default_factory=time.time) token_ids: Optional[List[int]] = field(default=None, repr=False) def __post_init__(self): """自动记录创建时间""" if self.timestamp == 0: self.timestamp = time.time() class DialogueHistory: """多轮对话历史管理器""" def __init__(self, max_turns: int = 20, max_tokens: int = 4096): """ 初始化对话历史。 Args: max_turns: 最大对话轮数 max_tokens: 最大token总数 """ self.max_turns = max_turns self.max_tokens = max_tokens self.messages: List[DialogueMessage] = [] self.system_prompt: Optional[DialogueMessage] = None def set_system_prompt(self, prompt: str) -> None: """设置系统提示词""" self.system_prompt = DialogueMessage(role="system", content=prompt) def add_message(self, role: str, content: str) -> None: """ 添加一条对话消息。 Args: role: 角色类型 ("user" / "assistant" / "system") content: 消息内容 """ msg = DialogueMessage(role=role, content=content) self.messages.append(msg) self._trim_history() def _trim_history(self) -> None: """裁剪历史,保持不超过最大轮数""" if len(self.messages) > self.max_turns * 2: # 保留最近的对话,但保留系统提示 self.messages = self.messages[-(self.max_turns * 2):] def get_messages(self) -> List[Dict[str, str]]: """ 获取格式化的消息列表。 Returns: 包含role和content的字典列表 """ result: List[Dict[str, str]] = [] if self.system_prompt: result.append({ "role": "system", "content": self.system_prompt.content, }) for msg in self.messages: result.append({"role": msg.role, "content": msg.content}) return result def get_context_tokens( self, tokenizer: ShimokumoTokenizer, ) -> List[int]: """ 将对话历史编码为token序列。 Args: tokenizer: 分词器实例 Returns: 编码后的token ID列表 """ messages = self.get_messages() return tokenizer.encode_chat(messages, add_bos=True, add_eos=False) def clear(self) -> None: """清空对话历史""" self.messages.clear() def get_last_n_turns(self, n: int) -> List[DialogueMessage]: """获取最近n轮对话""" return self.messages[-(n * 2):] def __len__(self) -> int: """返回消息数量""" return len(self.messages) class ShimokumoInference: """霜云推理引擎 封装模型推理逻辑,提供高效的文本生成接口。 支持流式输出和多轮对话。 用法: inference = ShimokumoInference(config, model, tokenizer) # 单次生成 response = inference.generate("你好") # 流式生成 for token_text in inference.generate_stream("你好"): print(token_text, end="", flush=True) # 多轮对话 inference.dialogue.add_message("user", "你好") response = inference.chat_reply() """ def __init__( self, config: ShimokumoConfig, model: ShimokumoModel, tokenizer: ShimokumoTokenizer, ): """ 初始化推理引擎。 Args: config: 模型配置 model: 霜云模型实例 tokenizer: 分词器实例 """ self.config = config self.model = model self.tokenizer = tokenizer self.device = self._select_device() # 将模型移到目标设备 self.model = self.model.to(self.device) # 设置为评估模式 self.model.eval() # 启用FP16混合精度 if config.use_fp16 and self.device.type == "cuda": self.model = self.model.half() logger.info("已启用FP16混合精度推理") else: logger.info(f"使用FP32精度在设备 {self.device} 上推理") # 多轮对话历史 self.dialogue = DialogueHistory( max_turns=20, max_tokens=config.max_seq_len - config.max_new_tokens, ) self.dialogue.set_system_prompt(config.system_prompt) logger.info(f"推理引擎初始化完成 | 设备: {self.device} | " f"模型参数: {model.get_num_params() / 1e9:.2f}B") def _select_device(self) -> torch.device: """选择计算设备""" if self.config.device != "auto": return torch.device(self.config.device) return torch.device("cuda" if torch.cuda.is_available() else "cpu") @torch.no_grad() def generate( self, prompt: str, max_new_tokens: Optional[int] = None, temperature: Optional[float] = None, top_k: Optional[int] = None, top_p: Optional[float] = None, repetition_penalty: Optional[float] = None, stop_tokens: Optional[List[str]] = None, ) -> str: """ 生成完整回复。 Args: prompt: 输入提示词 max_new_tokens: 最大生成token数 temperature: 采样温度 top_k: Top-K采样参数 top_p: Top-P采样参数 repetition_penalty: 重复惩罚系数 stop_tokens: 停止词列表 Returns: 生成的文本 """ # 收集生成的文本 generated_text = "" for token_text in self.generate_stream( prompt=prompt, max_new_tokens=max_new_tokens, temperature=temperature, top_k=top_k, top_p=top_p, repetition_penalty=repetition_penalty, stop_tokens=stop_tokens, ): generated_text += token_text return generated_text @torch.no_grad() def generate_stream( self, prompt: str, max_new_tokens: Optional[int] = None, temperature: Optional[float] = None, top_k: Optional[int] = None, top_p: Optional[float] = None, repetition_penalty: Optional[float] = None, stop_tokens: Optional[List[str]] = None, ) -> Generator[str, None, None]: """ 流式生成回复(逐token输出)。 Args: prompt: 输入提示词 max_new_tokens: 最大生成token数 temperature: 采样温度 top_k: Top-K采样参数 top_p: Top-P采样参数 repetition_penalty: 重复惩罚系数 stop_tokens: 停止词列表 Yields: 每次生成一个token的文本 """ # 使用配置默认值 max_new_tokens = max_new_tokens or self.config.max_new_tokens temperature = temperature if temperature is not None else self.config.temperature top_k = top_k if top_k is not None else self.config.top_k top_p = top_p if top_p is not None else self.config.top_p repetition_penalty = repetition_penalty if repetition_penalty is not None else self.config.repetition_penalty # 编码输入 input_ids = self.tokenizer.encode(prompt, add_bos=True, add_eos=False) input_tensor = torch.tensor([input_ids], dtype=torch.long, device=self.device) # 初始化KV缓存 num_layers = self.config.num_layers kv_caches: List[Optional[Tuple[torch.Tensor, torch.Tensor]]] = [None] * num_layers # 记录已生成的token,用于重复惩罚 all_generated_ids: List[int] = list(input_ids) position_offset = len(input_ids) # 生成循环 for step in range(max_new_tokens): # 前向传播(使用KV Cache) logits, kv_caches = self.model( input_ids=input_tensor, kv_caches=kv_caches, position_offset=position_offset, ) # 取最后一个位置的logits next_logits = logits[:, -1, :].float() # 转为FP32进行采样 # 应用温度 if temperature > 0: next_logits = next_logits / temperature # 应用重复惩罚 if repetition_penalty != 1.0 and all_generated_ids: for token_id in set(all_generated_ids): if next_logits[0, token_id] > 0: next_logits[0, token_id] /= repetition_penalty else: next_logits[0, token_id] *= repetition_penalty # Top-K过滤 if top_k > 0: top_k_values, _ = torch.topk(next_logits, min(top_k, next_logits.size(-1))) threshold = top_k_values[:, -1].unsqueeze(-1) next_logits = torch.where( next_logits < threshold, torch.full_like(next_logits, float("-inf")), next_logits, ) # Top-P(核采样)过滤 if top_p < 1.0: sorted_logits, sorted_indices = torch.sort(next_logits, descending=True) sorted_probs = F.softmax(sorted_logits, dim=-1) cumulative_probs = torch.cumsum(sorted_probs, dim=-1) # 移除累积概率超过top_p的token sorted_mask = cumulative_probs > top_p # 保留第一个超过阈值的token sorted_mask[..., 1:] = sorted_mask[..., :-1].clone() sorted_mask[..., 0] = False # 将掩码应用回原始顺序 mask = sorted_mask.scatter(1, sorted_indices, sorted_mask) next_logits = torch.where( mask, torch.full_like(next_logits, float("-inf")), next_logits, ) # 计算概率并采样 probs = F.softmax(next_logits, dim=-1) next_token = torch.multinomial(probs, num_samples=1).squeeze(-1) next_token_id = next_token.item() # 解码当前token token_text = self.tokenizer.decode([next_token_id], skip_special=True) # 检查是否遇到特殊token(如等) if next_token_id in self.tokenizer.SPECIAL_TOKEN_IDS: special_token = self.tokenizer.SPECIAL_TOKEN_IDS[next_token_id] if special_token == "": break # 遇到结束符则停止生成 # 跳过其他特殊token的输出 token_text = "" # 检查停止词 if stop_tokens and token_text: should_stop = False for stop in stop_tokens: if stop in token_text: # 截取到停止词之前的部分 idx = token_text.index(stop) token_text = token_text[:idx] should_stop = True break if should_stop: if token_text: yield token_text break # 产出token文本 if token_text: yield token_text # 更新状态 all_generated_ids.append(next_token_id) input_tensor = next_token.unsqueeze(0) position_offset += 1 @torch.no_grad() def chat_reply( self, user_message: Optional[str] = None, **kwargs, ) -> str: """ 多轮对话回复。 Args: user_message: 用户新消息(为None则使用对话历史中的最后一条用户消息) **kwargs: 传给generate的额外参数 Returns: 助手回复文本 """ if user_message: self.dialogue.add_message("user", user_message) # 构建完整的prompt messages = self.dialogue.get_messages() prompt_parts: List[str] = [] for msg in messages: if msg["role"] == "system": prompt_parts.append(f"[系统提示]\n{msg['content']}\n") elif msg["role"] == "user": prompt_parts.append(f"[用户]\n{msg['content']}\n") elif msg["role"] == "assistant": prompt_parts.append(f"[霜云]\n{msg['content']}\n") prompt_parts.append("[霜云]\n") full_prompt = "\n".join(prompt_parts) # 生成回复 response = self.generate(prompt=full_prompt, **kwargs) # 将助手回复加入历史 self.dialogue.add_message("assistant", response) return response @torch.no_grad() def chat_stream( self, user_message: Optional[str] = None, **kwargs, ) -> Generator[str, None, None]: """ 多轮对话流式回复。 Args: user_message: 用户新消息 **kwargs: 传给generate_stream的额外参数 Yields: 逐token的回复文本 """ if user_message: self.dialogue.add_message("user", user_message) # 构建完整prompt(同chat_reply) messages = self.dialogue.get_messages() prompt_parts: List[str] = [] for msg in messages: if msg["role"] == "system": prompt_parts.append(f"[系统提示]\n{msg['content']}\n") elif msg["role"] == "user": prompt_parts.append(f"[用户]\n{msg['content']}\n") elif msg["role"] == "assistant": prompt_parts.append(f"[霜云]\n{msg['content']}\n") prompt_parts.append("[霜云]\n") full_prompt = "\n".join(prompt_parts) # 流式生成并收集完整回复 full_response = "" for token_text in self.generate_stream(prompt=full_prompt, **kwargs): full_response += token_text yield token_text # 将完整回复加入历史 if full_response: self.dialogue.add_message("assistant", full_response) def clear_dialogue(self) -> None: """清空对话历史""" self.dialogue.clear() self.dialogue.set_system_prompt(self.config.system_prompt) logger.info("对话历史已清空") def get_dialogue_history(self) -> List[Dict[str, str]]: """获取对话历史""" return self.dialogue.get_messages()