全局修正

This commit is contained in:
DXC
2026-06-29 16:16:55 +08:00
parent 2788fb3fe1
commit e337f01312
7 changed files with 437 additions and 104 deletions

View File

@ -1,5 +1,6 @@
import numpy as np
from scipy import signal
from sklearn.base import BaseEstimator, TransformerMixin
from sklearn.linear_model import LinearRegression
from sklearn.preprocessing import MinMaxScaler, StandardScaler
import pandas as pd
@ -174,3 +175,223 @@ def Preprocessing(method, input_spectrum, save_path=None):
print("No such method of preprocessing!")
output_spectrum = input_spectrum.values
return output_spectrum
# ============================================================================
# sklearn Pipeline 兼容的 Transformer 包装
# ----------------------------------------------------------------------------
# 设计目的:让 12 种预处理方法都能塞进 sklearn.pipeline.Pipeline,从而:
# 1) 训练时 scaler/MSC mean spectrum 等状态被绑定在 Pipeline 内,
# 避免传统"手动 Preprocessing(X_raw) → 拆分 → CV"的数据泄露链;
# 2) 推理时直接 pipeline.predict(X_raw),无需重新应用预处理;
# 3) .joblib 内 model 字段即为完整 Pipeline,跨进程状态自包含。
#
# 注意:MMSTransformer/SSTransformer 直接复用 sklearn 自带的 MinMaxScaler/StandardScaler,
# 不再封装(避免维护重复代码);其余 9 种自定义方法各自写一个 TransformerMixin 子类。
# ============================================================================
class _ArrayAsFloat64:
"""统一的 ndarray 入口辅助(DataFrame/np.ndarray 都吃,输出 ndarray float64)"""
@staticmethod
def _to_ndarray(X):
if isinstance(X, pd.DataFrame):
return X.values.astype(np.float64)
return np.asarray(X, dtype=np.float64)
class IdentityTransformer(TransformerMixin, BaseEstimator, _ArrayAsFloat64):
"""无预处理(None)—— 数据原样透传,shape 不变。"""
def fit(self, X, y=None):
return self
def transform(self, X):
return self._to_ndarray(X)
class CTTransformer(TransformerMixin, BaseEstimator, _ArrayAsFloat64):
"""均值中心化(CT):每行减自身均值。shape 不变。"""
def fit(self, X, y=None):
return self
def transform(self, X):
X = self._to_ndarray(X)
return X - X.mean(axis=1, keepdims=True)
class SNVTransformer(TransformerMixin, BaseEstimator, _ArrayAsFloat64):
"""标准正态变换(SNV):每行 (x - mean) / std。shape 不变。"""
def fit(self, X, y=None):
return self
def transform(self, X):
X = self._to_ndarray(X)
row_mean = X.mean(axis=1, keepdims=True)
row_std = X.std(axis=1, keepdims=True)
row_std = np.where(row_std == 0, 1.0, row_std)
return (X - row_mean) / row_std
class MATransformer(TransformerMixin, BaseEstimator, _ArrayAsFloat64):
"""移动平均平滑(MA):每行卷积 np.ones(WSZ)/WSZ。shape 不变。"""
def __init__(self, wsz: int = 11):
self.wsz = wsz
def fit(self, X, y=None):
return self
def transform(self, X):
X = self._to_ndarray(X)
out = np.empty_like(X)
WSZ = self.wsz
r = np.arange(1, WSZ - 1, 2)
for i in range(X.shape[0]):
row = X[i]
out0 = np.convolve(row, np.ones(WSZ, dtype=int), 'valid') / WSZ
start = np.cumsum(row[:WSZ - 1])[::2] / r
stop = (np.cumsum(row[:-WSZ:-1])[::2] / r)[::-1]
out[i] = np.concatenate((start, out0, stop))
return out
class SGTransformer(TransformerMixin, BaseEstimator, _ArrayAsFloat64):
"""Savitzky-Golay 平滑(SG):每行调用 signal.savgol_filter。shape 不变。"""
def __init__(self, w: int = 15, p: int = 2):
self.w = w
self.p = p
def fit(self, X, y=None):
return self
def transform(self, X):
X = self._to_ndarray(X)
return signal.savgol_filter(X, self.w, self.p, axis=1)
class MSCTransformer(TransformerMixin, BaseEstimator, _ArrayAsFloat64):
"""多元散射校正(MSC):fit 阶段计算训练集平均光谱;transform 阶段对每行
以平均光谱为参考做线性回归 (k, b),输出 (x - b) / k。shape 不变。
注意:fit 阶段对每行分别拟合一次回归取 (k, b) 仅用于兼容旧实现,标准 MSC
只存储 mean_spectrum_。这里为了与原代码行为一致,保留 per-row 拟合路径。
"""
def fit(self, X, y=None):
X = self._to_ndarray(X)
self.mean_spectrum_ = X.mean(axis=0)
return self
def transform(self, X):
X = self._to_ndarray(X)
mean = self.mean_spectrum_
out = np.empty_like(X)
lr = LinearRegression()
for i in range(X.shape[0]):
y = X[i]
lr.fit(mean.reshape(-1, 1), y.reshape(-1, 1))
k = lr.coef_[0, 0]
b = lr.intercept_[0]
out[i] = (y - b) / (k if k != 0 else 1.0)
return out
class D1Transformer(TransformerMixin, BaseEstimator, _ArrayAsFloat64):
"""一阶导数(D1):每行 np.diff。shape 从 (n, p) → (n, p-1)。"""
def fit(self, X, y=None):
return self
def transform(self, X):
X = self._to_ndarray(X)
return np.diff(X, axis=1)
class D2Transformer(TransformerMixin, BaseEstimator, _ArrayAsFloat64):
"""二阶导数(D2):每行二次 np.diff。shape 从 (n, p) → (n, p-2)。"""
def fit(self, X, y=None):
return self
def transform(self, X):
X = self._to_ndarray(X)
return np.diff(X, n=2, axis=1)
class DTTransformer(TransformerMixin, BaseEstimator, _ArrayAsFloat64):
"""趋势校正(DT):每行对自身索引做线性回归,减去趋势线。shape 不变。"""
def fit(self, X, y=None):
return self
def transform(self, X):
X = self._to_ndarray(X)
n_cols = X.shape[1]
x = np.asarray(range(n_cols), dtype=np.float32).reshape(-1, 1)
out = np.empty_like(X)
lr = LinearRegression()
for i in range(X.shape[0]):
row = X[i]
lr.fit(x, row.reshape(-1, 1))
trend = (x @ lr.coef_.T + lr.intercept_).ravel()
out[i] = row - trend
return out
class WVAETransformer(TransformerMixin, BaseEstimator, _ArrayAsFloat64):
"""小波变换(WVAE):每行调用 pywt 阈值去噪重构。shape 可能略有变化。"""
def fit(self, X, y=None):
return self
def transform(self, X):
X = self._to_ndarray(X)
w = pywt.Wavelet('db8')
maxlev = pywt.dwt_max_level(X.shape[1], w.dec_len)
out = np.empty_like(X)
for i in range(X.shape[0]):
row = X[i]
coeffs = pywt.wavedec(row, 'db8', level=maxlev)
for ci in range(1, len(coeffs)):
coeffs[ci] = pywt.threshold(coeffs[ci], 0.04 * max(np.abs(coeffs[ci])) if coeffs[ci].size else 1.0)
reconstructed = pywt.waverec(coeffs, 'db8')
# waverec 可能比原信号长 1 元素(边界效应),裁剪对齐
out[i] = reconstructed[:X.shape[1]]
return out
# ============================================================================
# 工厂函数:根据方法名返回对应的 sklearn 兼容 Transformer(None 表示无预处理)
# ============================================================================
_PREPROCESSING_TRANSFORMERS = {
'None': IdentityTransformer,
'MMS': MinMaxScaler, # sklearn 自带
'SS': StandardScaler, # sklearn 自带
'CT': CTTransformer,
'SNV': SNVTransformer,
'MA': MATransformer,
'SG': SGTransformer,
'MSC': MSCTransformer,
'D1': D1Transformer,
'D2': D2Transformer,
'DT': DTTransformer,
'WVAE': WVAETransformer,
}
def get_preprocessing_transformer(method: str):
"""根据预处理方法名返回 sklearn 兼容的 Transformer 实例。
- method 为 "None" 或 None:返回 IdentityTransformer(等价于无处理)
- method 为 "MMS"/"SS":直接返回 sklearn 自带 MinMaxScaler/StandardScaler
- method 为 "CT"/"SNV"/"MA"/"SG"/"MSC"/"D1"/"D2"/"DT"/"WVAE":返回对应包装类
- method 不识别:返回 IdentityTransformer + 打印警告(与原 Preprocessing 行为一致)
Args:
method: 预处理方法名(大小写敏感,与 Preprocessing() 一致)
Returns:
sklearn 兼容的 Transformer 实例(可直接放入 Pipeline)
"""
if method is None:
return IdentityTransformer()
if method not in _PREPROCESSING_TRANSFORMERS:
print(f"未知预处理方法 '{method}',回退为 IdentityTransformer")
return IdentityTransformer()
cls = _PREPROCESSING_TRANSFORMERS[method]
return cls()