| """ |
| 霜云(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 |
| 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() |
|
|
| |
| 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) |
|
|
| |
| num_layers = self.config.num_layers |
| kv_caches: List[Optional[Tuple[torch.Tensor, torch.Tensor]]] = [None] * num_layers |
|
|
| |
| all_generated_ids: List[int] = list(input_ids) |
| position_offset = len(input_ids) |
|
|
| |
| for step in range(max_new_tokens): |
| |
| logits, kv_caches = self.model( |
| input_ids=input_tensor, |
| kv_caches=kv_caches, |
| position_offset=position_offset, |
| ) |
|
|
| |
| next_logits = logits[:, -1, :].float() |
|
|
| |
| 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 |
|
|
| |
| 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, |
| ) |
|
|
| |
| 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) |
|
|
| |
| sorted_mask = cumulative_probs > top_p |
| |
| 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_text = self.tokenizer.decode([next_token_id], skip_special=True) |
|
|
| |
| if next_token_id in self.tokenizer.SPECIAL_TOKEN_IDS: |
| special_token = self.tokenizer.SPECIAL_TOKEN_IDS[next_token_id] |
| if special_token == "<EOS>": |
| break |
| |
| 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 |
|
|
| |
| 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) |
|
|
| |
| 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) |
|
|
| |
| 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() |
|
|