diff --git a/src/core/modeling/modeling_batch.py b/src/core/modeling/modeling_batch.py index ceac44d..58e5b86 100644 --- a/src/core/modeling/modeling_batch.py +++ b/src/core/modeling/modeling_batch.py @@ -20,6 +20,7 @@ from sklearn.ensemble import GradientBoostingRegressor, AdaBoostRegressor, Extra from sklearn.tree import DecisionTreeRegressor from sklearn.neural_network import MLPRegressor from sklearn.pipeline import Pipeline +from sklearn.impute import SimpleImputer from joblib import parallel_backend # 第三方模型导入 # try: @@ -611,6 +612,7 @@ class WaterQualityModelingBatch: # ============ 关键:把预处理器塞进 Pipeline ============ preproc = get_preprocessing_transformer(preprocess_method) pipeline = Pipeline([ + ('imputer', SimpleImputer(strategy='median')), ('preproc', preproc), ('model', base_model), ]) diff --git a/src/core/prediction/inference_batch.py b/src/core/prediction/inference_batch.py index d4b1c68..b392bc0 100644 --- a/src/core/prediction/inference_batch.py +++ b/src/core/prediction/inference_batch.py @@ -48,9 +48,15 @@ class WaterQualityInference: # 规范化 loaded_model_data:始终为 dict,确保 ['model'] 访问不崩溃 if external_model is not None: - # 外部传入的是裸模型对象 → 包装为 dict,统一后续 .get('model') 访问 - self.loaded_model_data = {'model': external_model, 'preprocess_method': 'None'} - print(f" 外部模型已规范化: type={type(external_model).__name__}") + # ★ 外部模型可能是完整 dict(含 model + metadata + train_wavelengths), + # 也可能是裸 Pipeline 对象(旧版兼容) + if isinstance(external_model, dict) and 'model' in external_model: + self.loaded_model_data = external_model + print(f" 外部模型已规范化: dict (含 metadata)") + else: + self.loaded_model_data = {'model': external_model, + 'preprocess_method': 'None'} + print(f" 外部模型已规范化: type={type(external_model).__name__}") else: self.loaded_model_data = None diff --git a/src/gui/panels/step9_ml_predict_panel.py b/src/gui/panels/step9_ml_predict_panel.py index 0b5daf3..dfe7217 100644 --- a/src/gui/panels/step9_ml_predict_panel.py +++ b/src/gui/panels/step9_ml_predict_panel.py @@ -253,13 +253,14 @@ class Step9MlPredictPanel(QWidget): try: loaded = joblib.load(joblib_path) if isinstance(loaded, dict) and "model" in loaded: - model_obj = loaded["model"] + # ★ 保留完整 dict(含 metadata / train_wavelengths), + # 推理端需要 train_wavelengths 做光谱重采样 + models_found[subdir_name] = loaded elif hasattr(loaded, "predict"): - model_obj = loaded + models_found[subdir_name] = loaded else: errors.append(f"{subdir_name}: 无法识别的格式 {type(loaded).__name__}") continue - models_found[subdir_name] = model_obj except Exception as e: errors.append(f"{subdir_name}: {type(e).__name__}: {e}")