示例
核心逻辑
-
支持4种缩放方法:z-score标准化、最小-最大归一化、鲁棒缩放、最大绝对值缩放
-
训练-测试分离:在训练集上拟合参数,然后用相同参数转换测试集(避免数据泄露)
-
指定列处理 :只对
num_cols指定的数值列进行转换import pandas as pd
from sklearn.preprocessing import StandardScaler, MinMaxScaler, RobustScaler, MaxAbsScalerdef normalize(X_train, X_test, num_cols, method='zscore', copy=True, return_scaler=False):
"""
标准化/归一化数值列Parameters: ----------- X_train, X_test : pd.DataFrame 训练集和测试集 num_cols : list 需要处理的数值列名列表 method : str 'zscore', 'minmax', 'robust', 'maxabs' copy : bool 是否返回副本(避免修改原数据) return_scaler : bool 是否返回训练好的scaler对象 Returns: -------- X_train, X_test : pd.DataFrame 处理后的数据 scaler : object (optional) 训练好的scaler """ # 参数验证 if not isinstance(X_train, pd.DataFrame) or not isinstance(X_test, pd.DataFrame): raise TypeError("X_train and X_test must be pandas DataFrames") valid_methods = {'zscore', 'minmax', 'robust', 'maxabs'} if method not in valid_methods: raise ValueError(f"method must be one of {valid_methods}, got {method}") if not num_cols: return (X_train.copy(), X_test.copy()) if copy else (X_train, X_test) # 列存在性验证 missing_cols = set(num_cols) - set(X_train.columns) if missing_cols: raise ValueError(f"Missing columns in training data: {missing_cols}") # 拷贝数据 if copy: X_train = X_train.copy() X_test = X_test.copy() # 选择scaler scaler_map = { 'zscore': StandardScaler(), 'minmax': MinMaxScaler(), 'robust': RobustScaler(), 'maxabs': MaxAbsScaler() } scaler = scaler_map[method] # 检查缺失值 if X_train[num_cols].isnull().any().any(): print("Warning: Missing values detected in training data.") # 拟合并转换 X_train[num_cols] = scaler.fit_transform(X_train[num_cols]) X_test[num_cols] = scaler.transform(X_test[num_cols]) if return_scaler: return X_train, X_test, scaler return X_train, X_test
调用示例
import pandas as pd
import numpy as np
from sklearn.model_selection import train_test_split
# 创建示例数据集
np.random.seed(42)
data = pd.DataFrame({
'age': np.random.randint(18, 65, 100),
'salary': np.random.randint(30000, 150000, 100),
'experience': np.random.randint(0, 40, 100),
'score': np.random.uniform(60, 100, 100),
'city': np.random.choice(['北京', '上海', '深圳'], 100), # 分类变量
'target': np.random.randint(0, 2, 100)
})
# 划分训练集和测试集
X = data.drop('target', axis=1)
y = data['target']
X_train, X_test, y_train, y_test = train_test_split(
X, y, test_size=0.2, random_state=42
)
print("训练集形状:", X_train.shape)
print("测试集形状:", X_test.shape)
print("\n原始数据预览:")
print(X_train.head())
# 定义需要标准化的数值列
num_cols = ['age', 'salary', 'experience', 'score']
# 调用函数
X_train_norm, X_test_norm = normalize(X_train, X_test, num_cols, method='zscore')
print("Z-score标准化后:")
print(X_train_norm.head())
print(f"\n均值: {X_train_norm['age'].mean():.2f}") # 应接近0
print(f"标准差: {X_train_norm['age'].std():.2f}") # 应接近1
四种缩放方式对比
1. Z-score标准化 (StandardScaler)
数学原理
z=x−μσz=σx−μ
其中 μμ 是均值,σσ 是标准差
特点
-
输出范围:理论上无界,通常约在 -3, 3 之间
-
中心化:均值为0,标准差为1
-
分布形状:保持原始分布形状
代码示例
import numpy as np
import pandas as pd
from sklearn.preprocessing import StandardScaler
# 示例数据
data = np.array([10, 20, 30, 40, 50, 1000]).reshape(-1, 1) # 注意有异常值
scaler = StandardScaler()
scaled = scaler.fit_transform(data)
print("Z-score标准化结果:")
print(f"原始: {data.flatten()}")
print(f"标准化后: {scaled.flatten()}")
print(f"均值: {scaled.mean():.4f}") # 约等于0
print(f"标准差: {scaled.std():.4f}") # 等于1
适用场景
✅ 最适合:
-
线性模型(线性回归、逻辑回归、SVM)
-
神经网络(避免梯度消失/爆炸)
-
K-means聚类、PCA等依赖距离的算法
-
特征量纲差异大时
❌ 不适合:
-
数据有较多异常值(会被极度放大)
-
需要严格限制输出范围时
2. 最小-最大归一化 (MinMaxScaler)
数学原理
xnorm=x−xminxmax−xminxnorm=xmax−xminx−xmin
特点
-
输出范围:严格在 0, 1 之间
-
受异常值影响:极大(一个异常值可压缩所有正常数据)
-
保持相对大小:数据间的比例关系保持不变
代码示例
from sklearn.preprocessing import MinMaxScaler
data = np.array([10, 20, 30, 40, 50, 1000]).reshape(-1, 1)
scaler = MinMaxScaler()
scaled = scaler.fit_transform(data)
print("Min-Max归一化结果:")
print(f"原始: {data.flatten()}")
print(f"归一化后: {scaled.flatten()}")
print(f"最小值: {scaled.min():.4f}") # 0
print(f"最大值: {scaled.max():.4f}") # 1
# 特殊用法:指定范围 [a, b]
scaler_range = MinMaxScaler(feature_range=(-1, 1))
scaled_range = scaler_range.fit_transform(data)
print(f"\n缩放到[-1,1]: {scaled_range.flatten()}")
适用场景
✅ 最适合:
-
图像像素值处理(0-255 → 0-1)
-
神经网络激活函数(如sigmoid需要0-1输入)
-
需要边界约束的优化问题
-
特征有明确边界(如百分比数据)
❌ 不适合:
-
数据有明显异常值
-
数据分布不对称(会拉伸稀疏区域)
3. 鲁棒缩放 (RobustScaler)
数学原理
xrobust=x−Q2Q3−Q1xrobust=Q3−Q1x−Q2
其中 Q1Q1 是下四分位数(25%),Q2Q2 是中位数(50%),Q3Q3 是上四分位数(75%)
特点
-
使用中位数和IQR:对异常值不敏感
-
稳健性:极端值影响较小
-
输出范围:受数据分布影响,无固定边界
代码示例
from sklearn.preprocessing import RobustScaler
# 对比三种方法对异常值的敏感度
data_normal = np.array([10, 20, 30, 40, 50]).reshape(-1, 1)
data_outlier = np.array([10, 20, 30, 40, 50, 1000]).reshape(-1, 1)
# 无异常值时的表现
scaler_robust = RobustScaler()
scaled_normal = scaler_robust.fit_transform(data_normal)
print("正常数据(无异常值):")
print(f"原始: {data_normal.flatten()}")
print(f"Robust: {scaled_normal.flatten()}")
print(f"中位数: {scaled_normal.mean():.4f}")
# 有异常值时的表现
scaler_robust2 = RobustScaler()
scaled_outlier = scaler_robust2.fit_transform(data_outlier)
print("\n含异常值数据:")
print(f"原始: {data_outlier.flatten()}")
print(f"Robust: {scaled_outlier.flatten()}")
# 对比Z-score的敏感性
scaler_std = StandardScaler()
scaled_std = scaler_std.fit_transform(data_outlier)
print(f"Z-score: {scaled_std.flatten()}")
print("注意: Z-score中异常值被极大放大!")
适用场景
✅ 最适合:
-
数据包含明显的异常值(财务数据、传感器数据)
-
数据分布非正态(偏态分布)
-
业务指标(如收入分布,常见长尾分布)
-
需要稳健统计的场景
❌ 不适合:
-
数据量很小(分位数估计不稳定)
-
需要严格边界约束的场景
4. 最大绝对值缩放 (MaxAbsScaler)
数学原理
xmaxabs=x∣x∣maxxmaxabs=∣x∣maxx
其中 ∣x∣max∣x∣max 是最大绝对值
特点
-
保持正负号:保留数据的符号信息
-
输出范围:-1, 1 之间
-
稀疏性保持:不破坏稀疏数据结构(0值保持为0)
-
对称处理:正负值对称缩放
代码示例
from sklearn.preprocessing import MaxAbsScaler
# 包含正负值的数据
data = np.array([-100, -50, 0, 25, 50, 200]).reshape(-1, 1)
scaler = MaxAbsScaler()
scaled = scaler.fit_transform(data)
print("MaxAbs缩放结果:")
print(f"原始: {data.flatten()}")
print(f"缩放后: {scaled.flatten()}")
print(f"最大绝对值: {scaler.max_abs_}") # 200
print(f"范围: [{scaled.min():.2f}, {scaled.max():.2f}]")
# 稀疏矩阵示例
from scipy.sparse import csr_matrix
sparse_data = csr_matrix([[0, 0, 5], [0, 3, 0], [0, 0, 0]])
scaler_sparse = MaxAbsScaler()
sparse_scaled = scaler_sparse.fit_transform(sparse_data)
print("\n稀疏矩阵处理:")
print(f"原始非零值: [5, 3]")
print(f"缩放后非零值: {sparse_scaled.data}")
print("注意: 0值保持为0,适合稀疏数据")
适用场景
✅ 最适合:
-
稀疏数据(文本向量、One-Hot编码)
-
已经有中心化的数据
-
需要保持0值不变
-
需要保持符号信息
❌ 不适合:
-
数据不对称(正负值分布不均)
-
有极端异常值