shimokumo / src /model /inference.py
Shimokumo's picture
Upload folder using huggingface_hub
94fd0b0 verified
Raw
History Blame Contribute Delete
15.4 kB
"""
霜云(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(如</think_end>等)
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的输出
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()