# 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 base_url, data_record_url, 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"} resp = requests.post(data_record_url, json=payload, timeout=10) result = resp.json() 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} resp = requests.post(data_record_url, json=payload, timeout=15) result = resp.json() 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