基于分组中位数的Pandas DataFrame缺失值填充(支持训测分离)
问题:按Country分组计算中位数填充缺失值,并实现带fit/transform方法的类对象
需求说明
需要按Country字段分组计算各列中位数,仅基于训练集计算中位数,再用该中位数填充训练集和测试集对应分组的缺失值,最终实现一个包含fit、transform方法的对象,适配后续模型训练流程。
模拟数据与初始代码
import numpy as np import pandas as pd data = [['A', 10, 20, np.nan, np.nan, 50, 30], ['A', 2, 1, 5, np.nan, 34, 35], ['A', 13, 212, 3, 6, np.nan, 37], ['B', 120, 230, 53, np.nan, 63, 23], ['B', 22, 115, 15, 61, 4, 15], ['B', np.nan, 22, 12, np.nan, np.nan, 31], ['C', 105, 120, np.nan, 22, 520, 3], ['C', 26, 11, 15, np.nan, 34, 3], ['C', 13, np.nan, 13, 234, np.nan, 10], ['D', 101, 220, 654, 143, 634, 123], ['D', 32, 21, 61, 24, np.nan, 32], ['D', 11, 72, 23, np.nan, 534, 30] ] df = pd.DataFrame(data, columns=['Country','col1','col2','col3','col4','col5','col6']) # 初始分组计算中位数 median_data = df.groupby('Country').median().reset_index()
现有方法的问题
循环填充逻辑错误:
原循环代码会因索引对齐问题导致填充值混乱,例如Country A的col4应全部填充为6,但实际被后续循环的Country B的col4中位数61覆盖部分值。df_new = df.copy() for country in median_data.Country: country_data = median_data[median_data.Country == country].copy() for col in median_data.columns[2:]: df_new[col] = df_new[col].fillna(country_data[col])merge+update方法的顾虑:
该方法需要重置数据集索引,担心会干扰后续模型训练的索引关联逻辑。median_data = df.groupby('Country').median() df.update(df[['Country']].merge(median_data, on='Country', how='left'), overwrite=False)
解决方案:自定义Transformer类
实现一个兼容scikit-learn接口的自定义类,仅在fit阶段计算训练集的分组中位数,transform阶段用预计算的中位数填充缺失值,且无需修改原始索引。
完整代码实现
from sklearn.base import BaseEstimator, TransformerMixin class GroupMedianImputer(BaseEstimator, TransformerMixin): def __init__(self, group_col='Country'): self.group_col = group_col self.median_dict = {} # 存储每个分组各列的中位数 def fit(self, X, y=None): # 仅基于训练集计算分组中位数 median_df = X.groupby(self.group_col).median().reset_index() # 转换为字典:key是Country,value是各列中位数的字典 self.median_dict = median_df.set_index(self.group_col).to_dict('index') return self def transform(self, X): X_transformed = X.copy() # 遍历所有需要填充的列(排除分组列) cols_to_impute = [col for col in X_transformed.columns if col != self.group_col] for col in cols_to_impute: # 为每个样本匹配对应分组的中位数 median_values = X_transformed[self.group_col].map(lambda c: self.median_dict[c][col]) # 填充缺失值 X_transformed[col] = X_transformed[col].fillna(median_values) return X_transformed
使用示例
# 拆分训练集和测试集(示例拆分) train_df = df.iloc[:9].copy() test_df = df.iloc[9:].copy() # 初始化并训练填充器 imputer = GroupMedianImputer(group_col='Country') imputer.fit(train_df) # 填充训练集和测试集 train_imputed = imputer.transform(train_df) test_imputed = imputer.transform(test_df) # 验证结果:查看Country A的col4填充情况 print(train_imputed[train_imputed['Country'] == 'A']['col4'])
说明
fit方法:仅处理训练集,计算每个Country分组下各列的中位数并存储为字典,确保后续填充用的是训练集的统计量,避免数据泄露。transform方法:对输入数据(训练/测试集),按Country匹配预存的中位数填充缺失值,全程保留原始索引,不影响后续模型训练。- 兼容scikit-learn的Pipeline流程,可以直接加入模型训练流水线中。
内容的提问来源于stack exchange,提问作者MushyMush
相关产品推荐
相关产品推荐

