Pandas按filename分组统计pred与gt不匹配次数(transform实现)
按分组统计不匹配行数的解决方案
需求
按filename列分组,计算每组中pred与gt不相等的行数,并将该统计结果对应到每组的每一行中。比如f1.wav只有1行不匹配,该行及同组其他行的统计值都是1;f2.wav有3行不匹配,同组所有行的统计值都是3。
错误代码问题分析
你之前的尝试代码:
df.groupby('filename').transform(lambda x: x['pred'].ne(x['gt']).sum(), axis=1)
触发报错:TypeError: Transform function invalid for data types
问题出在两点:
axis=1参数误用:按filename分组是对行进行分组,transform默认按列(axis=0)处理即可,指定axis=1会让pandas尝试按列处理分组,逻辑完全错误。- 未指定目标列就直接操作整个分组:
groupby('filename').transform会默认对每列单独执行函数,而你的lambda是操作整个分组的DataFrame,导致数据类型不兼容。
正确实现方法
1. 先标记不匹配行
先新增一列,标记每行是否存在pred != gt的情况:
df['mismatch'] = df['pred'] != df['gt']
这一步会生成布尔值列,True代表该行不匹配,False代表匹配。
2. 分组统计并广播结果
对filename分组后,针对mismatch列使用transform('sum')——sum会把布尔值转为0/1求和,得到每组的不匹配行数,transform会把这个统计值广播到该组的每一行:
df['mismatch_count'] = df.groupby('filename')['mismatch'].transform('sum')
完整示例代码
import pandas as pd # 构造示例DataFrame df = pd.DataFrame([ {'pred': 0, 'gt': 0, 'filename': 'f1.wav'}, {'pred': 1, 'gt': 1, 'filename': 'f1.wav'}, {'pred': 0, 'gt': 1, 'filename': 'f1.wav'}, # 不匹配 {'pred': 1, 'gt': 0, 'filename': 'f2.wav'}, # 不匹配 {'pred': 1, 'gt': 0, 'filename': 'f2.wav'}, # 不匹配 {'pred': 0, 'gt': 1, 'filename': 'f2.wav'}, # 不匹配 {'pred': 0, 'gt': 0, 'filename': 'f2.wav'}, ]) # 标记不匹配行 df['mismatch'] = df['pred'] != df['gt'] # 分组统计每组不匹配行数并广播到每行 df['mismatch_count'] = df.groupby('filename')['mismatch'].transform('sum') print(df)
预期输出
pred gt filename mismatch mismatch_count 0 0 0 f1.wav False 1 1 1 1 f1.wav False 1 2 0 1 f1.wav True 1 3 1 0 f2.wav True 3 4 1 0 f2.wav True 3 5 0 1 f2.wav True 3 6 0 0 f2.wav False 3
内容的提问来源于stack exchange,提问作者Kenan
相关产品推荐
相关产品推荐

