Files
ReinLoopTest/ReinLoop/core/model_manager.py
T
2026-08-03 14:05:27 +08:00

138 lines
4.8 KiB
Python

# model_manager.py
"""RL 模型管理器:从云端扫描和加载强化学习模型。
纯业务逻辑,无 UI 依赖。通过回调与 UI 层通信。
"""
import threading
import io
import requests
import torch
from stable_baselines3 import SAC
# 关键:禁用 PyTorch 内部多线程,防止在 PyInstaller daemon 线程中 segfault
torch.set_num_threads(1)
from api import device_post, the_folder
class ModelManager:
"""管理 RL 模型的云端扫描与加载"""
def __init__(self):
self.model_file_map = {} # 文件名 → fileID 映射
self.rl_model = None # 加载的 SAC 模型实例
self._on_log = None
self._on_models_loaded = None
self._on_load_complete = None
# ---- 回调设置 ----
def set_log_callback(self, callback):
"""设置日志回调: callback(message: str)"""
self._on_log = callback
def set_models_loaded_callback(self, callback):
"""设置模型列表加载完成回调: callback(file_names: list)"""
self._on_models_loaded = callback
def set_load_complete_callback(self, callback):
"""设置模型加载完成回调: callback(success: bool, message: str)"""
self._on_load_complete = callback
def log(self, message):
if self._on_log:
self._on_log(message)
# ---- 模型扫描 ----
def scan_models(self):
"""异步扫描云端模型文件夹,完成后回调通知"""
def fetch_models():
try:
payload = {"type": "listModels", "folder": f"{the_folder}/model_config"}
result = device_post(payload, timeout=10)
if result.get("success"):
files = result.get("files", [])
file_list = result.get("fileList", [])
self.model_file_map = {
item.get("fileName"): item.get("fileID")
for item in file_list if item.get("fileName")
}
self.log("模型列表刷新成功")
if self._on_models_loaded:
self._on_models_loaded(files)
else:
err = result.get('errMsg', '未知错误')
self.log(f"获取模型列表失败: {err}")
except Exception as e:
self.log(f"扫描模型异常: {str(e)}")
threading.Thread(target=fetch_models, daemon=True).start()
# ---- 模型加载 ----
def load_model(self, model_name: str):
"""异步从云端加载指定的 RL 模型
Args:
model_name: 模型文件名
"""
if not model_name or model_name == "无模型文件":
self.log("错误:请先选择一个有效的模型")
return
def download_and_load():
try:
file_id = self.model_file_map.get(model_name)
if not file_id:
self.log("模型加载失败: 缺少 fileID,请先刷新模型列表")
return
self.log(f"正在加载模型: {model_name}...")
# 获取临时下载 URL
payload = {"type": "downloadModel", "fileID": file_id}
result = device_post(payload, timeout=15)
if not result.get("success"):
err = result.get('errMsg', '未知错误')
self.log(f"模型加载异常: {err}")
return
url = result['url']
# 下载模型文件
model_resp = requests.get(url, timeout=30)
if model_resp.status_code != 200:
self.log(f"模型加载异常: HTTP {model_resp.status_code}")
return
model_bytes = model_resp.content
# 直接加载到内存
model_stream = io.BytesIO(model_bytes)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.rl_model = SAC.load(model_stream, device=device)
self.log(f"成功加载模型: {model_name}")
if self._on_load_complete:
self._on_load_complete(True, f"成功加载模型: {model_name}")
except Exception as e:
msg = str(e)
self.log(f"加载模型失败: {msg}")
if self._on_load_complete:
self._on_load_complete(False, f"加载失败: {msg}")
threading.Thread(target=download_and_load, daemon=True).start()
def is_model_loaded(self) -> bool:
"""检查模型是否已加载"""
return self.rl_model is not None