全局修正
This commit is contained in:
@ -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()
|
||||
|
||||
Reference in New Issue
Block a user