如何从含6类标签的DataFrame中每类移除10行并拆分数据集
解决方法
你可以利用DataFrame的索引快速移除测试集的行,因为你抽取的df_tester保留了原DataFrame的索引信息,直接通过这些索引就能完成删除操作:
方法一:基于现有代码直接删除
import pandas as pd # 你的原有代码:抽取每类10行作为测试集 df_tester = pd.concat(g.sample(10) for idx, g in df.groupby('Label')) # 从原DataFrame中删除测试集对应的行 df = df.drop(df_tester.index)
方法二:先收集索引再拆分(更直观)
如果想更清晰地控制采样过程,可以先收集每个分组的采样索引,再生成测试集并删除原数据:
test_indices = [] # 遍历每个标签分组,收集10行的索引 for _, group in df.groupby('Label'): sampled_idx = group.sample(10).index test_indices.extend(sampled_idx) # 生成测试集 df_tester = df.loc[test_indices].copy() # 删除原数据中的对应行 df = df.drop(test_indices)
验证结果
执行完后用len(df)检查原DataFrame的行数,应该是650-60=590,和预期一致。
额外说明:如果你的DataFrame存在重复索引,上述方法可能会误删多行,这种情况下可以用反连接方式拆分:
df = df.merge(df_tester, how='left', indicator=True) df = df[df['_merge'] == 'left_only'].drop(columns='_merge')
内容的提问来源于stack exchange,提问作者Sam
相关产品推荐
相关产品推荐

