137 lines
4.8 KiB
Python
137 lines
4.8 KiB
Python
# model_manager.py
|
|
"""RL 模型管理器:从云端扫描和加载强化学习模型。
|
|
|
|
纯业务逻辑,无 UI 依赖。通过回调与 UI 层通信。
|
|
"""
|
|
|
|
import threading
|
|
import io
|
|
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
|