diff --git a/src/gui/panels/step9_ml_predict_panel.py b/src/gui/panels/step9_ml_predict_panel.py index dfe7217..f69e93f 100644 --- a/src/gui/panels/step9_ml_predict_panel.py +++ b/src/gui/panels/step9_ml_predict_panel.py @@ -233,6 +233,14 @@ class Step9MlPredictPanel(QWidget): return self.external_model_dir = dir_path + + # ★ 读取用户选择的模型评估指标 + _metric_key = self.metric.currentData() # "test_r2" / "test_rmse" / "test_mae" + _metric_display = self.metric.currentText() # "R² (决定系数)" / ... + _is_r2 = 'r2' in _metric_key.lower() + _is_rmse = 'rmse' in _metric_key.lower() + _is_mae = 'mae' in _metric_key.lower() + models_found = {} errors = [] @@ -243,26 +251,53 @@ class Step9MlPredictPanel(QWidget): if not subentry.is_dir(): continue subdir_name = subentry.name - joblib_files = [ + joblib_files = sorted([ f for f in os.scandir(subentry.path) if f.is_file() and f.name.lower().endswith(".joblib") - ] + ], key=lambda x: x.name) if not joblib_files: continue - joblib_path = joblib_files[0].path - try: - loaded = joblib.load(joblib_path) - if isinstance(loaded, dict) and "model" in loaded: - # ★ 保留完整 dict(含 metadata / train_wavelengths), - # 推理端需要 train_wavelengths 做光谱重采样 - models_found[subdir_name] = loaded - elif hasattr(loaded, "predict"): - models_found[subdir_name] = loaded - else: - errors.append(f"{subdir_name}: 无法识别的格式 {type(loaded).__name__}") + + # ★ 遍历该子目录下全部 .joblib 文件,按 metric 选出最佳模型 + best_score = -float('inf') if _is_r2 else float('inf') + best_entry = None + best_fname = None + + for f_entry in joblib_files: + try: + data = joblib.load(f_entry.path) + if not isinstance(data, dict) or "model" not in data: + continue + meta = data.get('metadata', {}) + # 按指标优先级:test_xxx → train_xxx + score = None + if _is_r2: + score = meta.get('test_r2', meta.get('train_r2', None)) + elif _is_rmse: + score = meta.get('test_rmse', meta.get('train_rmse', None)) + elif _is_mae: + score = meta.get('test_mae', meta.get('train_mae', None)) + + if score is None: + continue + + if _is_r2 and score > best_score: + best_score = score + best_entry = data + best_fname = f_entry.name + elif not _is_r2 and score < best_score: + best_score = score + best_entry = data + best_fname = f_entry.name + except Exception: continue - except Exception as e: - errors.append(f"{subdir_name}: {type(e).__name__}: {e}") + + if best_entry is not None: + models_found[subdir_name] = best_entry + print(f"[模型加载] 目标 {subdir_name} 根据 {_metric_display} " + f"自动选择最佳模型: {best_fname}, 分数: {best_score:.6f}") + else: + errors.append(f"{subdir_name}: 无可评估的 .joblib 文件") except Exception as e: QMessageBox.warning(