You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于分组中位数的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()

现有方法的问题

  1. 循环填充逻辑错误:
    原循环代码会因索引对齐问题导致填充值混乱,例如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])
    
  2. 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.17 00:31:39