如何解决Pandas DataFrame多列分层抽样的非精确问题?
解决多标签精确分层抽样问题
针对你遇到的IterativeStratification分层不精确、fold大小不一致的问题,可以通过直接对标签组合做精确分配的方式解决,确保每个标签组合在各fold中的数量差异≤1,且所有fold大小完全一致。
核心思路
- 将多列标签合并为唯一的标签组合标识,方便按组合分组处理。
- 对每个标签组合,计算其在各fold中应分配的基础数量(
总样本数 // 折数)和余数(总样本数 % 折数)。 - 随机将余数部分分配到不同的fold(每个fold最多多1个该组合的样本),保证组合分布差异≤1。
- 最终所有fold的样本数自动对齐(因为总样本数是折数的整数倍)。
代码实现
假设你的DataFrame标签列为['c', 'f', 'h'],总样本300,分10折:
import pandas as pd import numpy as np # 1. 生成标签组合列(替换为你的实际标签列) label_cols = ['c', 'f', 'h'] df['label_combination'] = df[label_cols].apply(lambda x: ','.join(x.astype(str)), axis=1) n_folds = 10 # 2. 统计每个标签组合的样本量,计算分配方案 comb_stats = df['label_combination'].value_counts().reset_index() comb_stats.columns = ['combination', 'total'] comb_stats['per_fold_base'] = comb_stats['total'] // n_folds comb_stats['remainder'] = comb_stats['total'] % n_folds # 3. 初始化fold容器 folds = [[] for _ in range(n_folds)] # 4. 按组合分配样本 for _, row in comb_stats.iterrows(): comb = row['combination'] base_count = row['per_fold_base'] remainder = row['remainder'] # 获取当前组合的所有样本索引并打乱 comb_indices = df[df['label_combination'] == comb].index.tolist() np.random.shuffle(comb_indices) # 随机选择需要多分配1个样本的fold extra_folds = np.random.choice(range(n_folds), size=remainder, replace=False) if remainder > 0 else [] idx = 0 for fold_idx in range(n_folds): # 确定当前fold的分配数量 assign_count = base_count + 1 if fold_idx in extra_folds else base_count folds[fold_idx].extend(comb_indices[idx:idx + assign_count]) idx += assign_count # 5. 将分配结果映射回原DataFrame df['fold'] = -1 for fold_idx, indices in enumerate(folds): df.loc[indices, 'fold'] = fold_idx
结果验证
运行以下代码验证效果:
# 验证各fold样本数(应全为30) print("各fold样本量:") print(df['fold'].value_counts().sort_index()) # 验证每个标签组合的fold分布差异(应≤1) print("\n标签组合分布检查:") for comb in comb_stats['combination']: fold_counts = df[df['label_combination'] == comb]['fold'].value_counts().reindex(range(n_folds), fill_value=0) max_cnt = fold_counts.max() min_cnt = fold_counts.min() if max_cnt - min_cnt > 1: print(f"⚠️ 组合 {comb} 分布偏差:{max_cnt - min_cnt}") else: print(f"✅ 组合 {comb} 符合要求:{fold_counts.tolist()}")
为什么这个方法有效
- 针对每个标签组合做精确的数量分配,避免了贪心算法的累积误差,确保每个组合在各fold的数量差异最多为1。
- 总样本数为折数的整数倍时,所有fold的样本量会自动对齐(因为余数总和是折数的整数倍,每个fold被分配到的额外样本数相同)。
- 随机打乱和随机分配余数,保证了分层的随机性,避免引入偏差。
内容的提问来源于stack exchange,提问作者nemo
相关产品推荐
相关产品推荐

