feat: 推理端自动波长回填+linspace子采样(旧模型兼容)

This commit is contained in:
duxin
2026-07-27 18:08:37 +08:00
parent de9d6f0835
commit c16cddf384

View File

@ -375,6 +375,58 @@ class WaterQualityInference:
print(f" 模型名称: {self.loaded_model_data['model_name']}") print(f" 模型名称: {self.loaded_model_data['model_name']}")
print(f" 模型类型: {type(self.loaded_model_data['model'])}") print(f" 模型类型: {type(self.loaded_model_data['model'])}")
def _auto_detect_train_wavelengths(self):
"""自动获取训练波长:优先级 json/txt > 工作目录 CSV > None"""
import os as _os, json as _json, glob as _glob
_dir = str(self.artifacts_dir)
# 搜集所有可能的搜索根目录
_roots = [_dir, _os.path.dirname(_dir), _os.getcwd()]
# 找所有 work_dir 级别的父目录
for _r in list(_roots):
_p = _os.path.dirname(_r)
if _p and _p not in _roots:
_roots.append(_p)
# 优先级 1: 显式 wavelength 文件JSON/TXT
for _root in _roots:
for _fname in ('train_wavelengths.json', 'train_wavelengths.txt'):
_p = _os.path.join(_root, _fname)
if not _os.path.isfile(_p):
continue
try:
if _fname.endswith('.json'):
wl = _json.load(open(_p))
if wl:
print(f"[波长回填] 来源: {_p}")
return wl
else:
wl = [float(x) for x in open(_p).read().strip().split()]
if wl:
print(f"[波长回填] 来源: {_p}")
return wl
except Exception:
pass
# 优先级 2: training_spectra.csv / sampling_spectra.csv
for _root in _roots:
for _sub in ('4_sampling', '6_Spectral_Feature_Extraction', ''):
_d = _os.path.join(_root, _sub) if _sub else _root
for _pat in ('sampling_spectra.csv', 'training_spectra.csv'):
_p = _os.path.join(_d, _pat)
if not _os.path.isfile(_p):
continue
try:
_df = pd.read_csv(_p, nrows=0)
_wl = [float(c) for c in _df.columns
if c.replace('.','').lstrip('-').isdigit()]
if _wl:
print(f"[波长回填] 来源: {_p} ({len(_wl)} 波长)")
return _wl
except Exception:
pass
return None
def preprocess_spectra(self, spectra: pd.DataFrame) -> np.ndarray: def preprocess_spectra(self, spectra: pd.DataFrame) -> np.ndarray:
""" """
对光谱数据进行预处理 + 跨传感器光谱重采样。 对光谱数据进行预处理 + 跨传感器光谱重采样。
@ -418,6 +470,7 @@ class WaterQualityInference:
model = self.loaded_model_data['model'] model = self.loaded_model_data['model']
metadata = self.loaded_model_data.get('metadata', {}) metadata = self.loaded_model_data.get('metadata', {})
train_wavelengths = metadata.get('train_wavelengths', None) train_wavelengths = metadata.get('train_wavelengths', None)
# 旧模型无 train_wavelengths → 不做波长匹配,走下方分支 B 的 linspace 路径
train_columns = metadata.get('train_columns', None) train_columns = metadata.get('train_columns', None)
# ═══════════════════════════════════════════════════════════ # ═══════════════════════════════════════════════════════════