swarm-backend / core /local_inference.py
lk080424's picture
Upload folder using huggingface_hub
17fba62 verified
Raw
History Blame Contribute Delete
10.1 kB
#!/usr/bin/env python3
"""
虫群本地推理后端
- 加载训练好的SwarmModel模型
- 提供统一的推理接口
- 支持多模型管理和切换
- 与MOA引擎集成
"""
import os
import sys
import time
import json
import torch
import threading
from typing import Dict, List, Optional, Tuple
sys.path.insert(0, "/home/admin/swarm")
from training.model import SwarmModel
from training.tokenizer import SwarmTokenizer
class LocalModelBackend:
"""
本地模型推理后端
管理多个本地小模型的加载、推理和卸载。
按需加载模型到内存,支持LRU淘汰。
"""
def __init__(self, models_dir: str = "/home/admin/swarm/training/models", max_loaded: int = 3):
self.models_dir = models_dir
self.max_loaded = max_loaded
# 已加载的模型: model_id -> (model, tokenizer)
self._loaded: Dict[str, Tuple[SwarmModel, SwarmTokenizer]] = {}
self._load_times: Dict[str, float] = {} # LRU时间戳
self._lock = threading.Lock()
# 可用模型列表
self._available = self._scan_models()
# 设备
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def _scan_models(self) -> Dict[str, dict]:
"""扫描可用模型"""
available = {}
if not os.path.exists(self.models_dir):
return available
# 共享分词器路径
shared_tok = os.path.join(self.models_dir, "shared_tokenizer.json")
# v2分词器
v2_tok = os.path.join(self.models_dir, "tokenizer_v2.json")
for name in os.listdir(self.models_dir):
model_dir = os.path.join(self.models_dir, name)
if not os.path.isdir(model_dir):
continue
model_path = os.path.join(model_dir, "model.pt")
meta_path = os.path.join(model_dir, "meta.json")
if os.path.exists(model_path):
meta = {}
if os.path.exists(meta_path):
with open(meta_path, "r") as f:
meta = json.load(f)
# 优先用模型目录下的分词器,否则用共享分词器或v2分词器
tok_path = os.path.join(model_dir, "tokenizer.json")
if not os.path.exists(tok_path) and os.path.exists(shared_tok):
tok_path = shared_tok
if not os.path.exists(tok_path) and os.path.exists(v2_tok):
tok_path = v2_tok
# 检测量化模型
quantized = os.path.exists(os.path.join(model_dir, "quantized_int8"))
available[name] = {
"path": model_dir,
"model_path": model_path,
"tokenizer_path": tok_path,
"meta": meta,
"quantized": quantized,
}
return available
def list_available(self) -> List[str]:
"""列出可用模型"""
return list(self._available.keys())
def get_model_info(self, model_id: str) -> Optional[dict]:
"""获取模型信息"""
if model_id in self._available:
return self._available[model_id].get("meta", {})
if model_id in self._loaded:
model, _ = self._loaded[model_id]
return model.get_info()
return None
def load(self, model_id: str) -> bool:
"""加载模型到内存"""
if model_id in self._loaded:
self._load_times[model_id] = time.time()
return True
if model_id not in self._available:
print(f" ❌ 模型不存在: {model_id}")
return False
# LRU淘汰
with self._lock:
while len(self._loaded) >= self.max_loaded:
oldest = min(self._load_times, key=self._load_times.get)
self.unload(oldest)
try:
info = self._available[model_id]
# 加载分词器
tokenizer = SwarmTokenizer.load(info["tokenizer_path"])
# 加载模型
checkpoint = torch.load(info["model_path"], map_location=self.device, weights_only=False)
saved_cfg = checkpoint.get("config", {})
# 从权重推断实际max_len(保存的max_len可能不准)
actual_max_len = 256
state_dict = checkpoint.get("model_state_dict", {})
if "pos_emb.weight" in state_dict:
actual_max_len = state_dict["pos_emb.weight"].shape[0]
elif "max_len" in checkpoint:
actual_max_len = checkpoint["max_len"]
model = SwarmModel(
vocab_size=checkpoint.get("vocab_size", tokenizer.vocab_size_actual),
d_model=saved_cfg.get("d_model", 192 if model_id == "tiny" else 320),
n_heads=saved_cfg.get("n_heads", 6 if model_id == "tiny" else 8),
n_layers=saved_cfg.get("n_layers", 6 if model_id == "tiny" else 8),
d_ff=saved_cfg.get("d_model", 192) * 4, # 从d_model推算
max_len=actual_max_len,
)
model.load_state_dict(checkpoint["model_state_dict"])
model = model.to(self.device)
model.eval()
with self._lock:
self._loaded[model_id] = (model, tokenizer)
self._load_times[model_id] = time.time()
param_count = model.get_info()["params_M"]
print(f" ✅ 模型已加载: {model_id} ({param_count}M参数)")
return True
except Exception as e:
print(f" ❌ 加载失败: {model_id}, {e}")
return False
def unload(self, model_id: str) -> bool:
"""卸载模型"""
with self._lock:
if model_id in self._loaded:
del self._loaded[model_id]
del self._load_times[model_id]
# 释放GPU内存
if torch.cuda.is_available():
torch.cuda.empty_cache()
print(f" 🗑️ 模型已卸载: {model_id}")
return True
return False
def infer(
self,
model_id: str,
query: str,
max_new_tokens: int = 128,
temperature: float = 0.7,
top_k: int = 40,
) -> Dict:
"""
推理接口
返回:
{
"response": str,
"model_id": str,
"latency_ms": float,
"tokens_generated": int,
"success": bool,
}
"""
# 确保模型已加载
if model_id not in self._loaded:
if not self.load(model_id):
return {
"response": "",
"model_id": model_id,
"latency_ms": 0,
"tokens_generated": 0,
"success": False,
}
model, tokenizer = self._loaded[model_id]
start_time = time.perf_counter()
try:
# 编码输入(只加BOS,不加EOS——EOS是生成终止标记)
bos_id = tokenizer.bos_id
input_ids = [bos_id] + tokenizer.encode(query, add_special=False)
input_tensor = torch.tensor([input_ids], dtype=torch.long).to(self.device)
# 生成(启用重复惩罚,小模型必备)
with torch.no_grad():
output_ids = model.generate(
input_tensor,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_k=top_k,
eos_id=tokenizer.eos_id,
repetition_penalty=1.3,
)
# 解码
output_ids = output_ids[0].cpu().tolist()
# 去掉输入部分
generated_ids = output_ids[len(input_ids):]
response = tokenizer.decode(generated_ids)
# 清理
response = response.replace("[UNK]", "").strip()
latency_ms = (time.perf_counter() - start_time) * 1000
return {
"response": response,
"model_id": model_id,
"latency_ms": round(latency_ms, 1),
"tokens_generated": len(generated_ids),
"success": True,
}
except Exception as e:
latency_ms = (time.perf_counter() - start_time) * 1000
return {
"response": f"推理错误: {str(e)[:100]}",
"model_id": model_id,
"latency_ms": round(latency_ms, 1),
"tokens_generated": 0,
"success": False,
}
def get_status(self) -> Dict:
"""后端状态"""
return {
"available_models": list(self._available.keys()),
"loaded_models": list(self._loaded.keys()),
"max_loaded": self.max_loaded,
"device": str(self.device),
}
# 全局单例
_backend = None
_backend_lock = threading.Lock()
def get_local_backend() -> LocalModelBackend:
"""获取本地推理后端单例"""
global _backend
with _backend_lock:
if _backend is None:
_backend = LocalModelBackend()
return _backend
if __name__ == "__main__":
# 测试
backend = LocalModelBackend()
print(f"可用模型: {backend.list_available()}")
print(f"状态: {backend.get_status()}")
for model_id in backend.list_available():
print(f"\n--- 测试 {model_id} ---")
result = backend.infer(model_id, "你好")
print(f" 回复: {result['response'][:100]}")
print(f" 延迟: {result['latency_ms']}ms")
print(f" 成功: {result['success']}")