编写ID标签标准化函数遇异常:多数标签未被正确应用
标签标准化函数异常修复
我编写了一个函数,旨在按指定ID分组,用每组最常用的标签标准化标签列;若没有多数标签(即排名前两位的标签计数相同),则取该组首个观测值作为默认标准。函数大部分场景运行正常,但遇到某ID的标签存在趋势变化时出现异常:例如ID=222的数据中,"LA Metro"出现3次,是绝对多数标签,但部分行的标准化标签却显示为"Los Angeles Metro",期望所有行的标准化标签统一为"LA Metro"。
原函数代码:
def standardize_labels(df, id_col, label_col): # Function to find the most common label or the first one if there's a tie def most_common_label(group): labels = group[label_col].value_counts() # Check if the top two labels have the same count if len(labels) > 1 and labels.iloc[0] == labels.iloc[1]: return group[label_col].iloc[0] return labels.idxmax() # Group by the ID column and apply the most_common_label function common_labels = df.groupby(id_col).apply(most_common_label) # Map the IDs in the original DataFrame to their common labels df['standardized_label'] = df[id_col].map(common_labels) return df
异常示例数据:
| ID | raw_label | standardized_label |
|---|---|---|
| 222 | LA Metro | LA Metro |
| 222 | LA Metro | LA Metro |
| 222 | Los Angeles Metro | Los Angeles Metro |
| 222 | LA Metro | Los Angeles Metro |
| 222 | Los Angeles Metro | Los Angeles Metro |
问题原因分析
原函数使用groupby(id_col).apply(most_common_label)对整个分组DataFrame处理,虽逻辑正确,但在部分场景下(如分组数据结构差异)可能导致聚合结果异常;另外value_counts()未显式指定排序规则,存在版本差异导致的排序不确定性,进而影响idxmax()的返回值。
修复后的函数
方案一:优化聚合逻辑,直接处理标签列
import pandas as pd def standardize_labels(df, id_col, label_col): def get_standard_label(group): # 强制按计数降序排列标签 label_counts = group.value_counts(ascending=False) # 检查是否存在平局 if len(label_counts) > 1 and label_counts.iloc[0] == label_counts.iloc[1]: return group.iloc[0] # 返回计数最高的标签 return label_counts.idxmax() # 按ID分组聚合得到标准标签映射 standard_label_map = df.groupby(id_col)[label_col].agg(get_standard_label) # 映射生成标准化列 df['standardized_label'] = df[id_col].map(standard_label_map) return df
方案二:利用mode()函数简化实现
import pandas as pd def standardize_labels(df, id_col, label_col): def get_standard_label(group): # 获取出现次数最多的标签(可能多个) mode_labels = group.mode() if len(mode_labels) > 1: # 平局时返回组内首个标签 return group.iloc[0] # 非平局时返回唯一众数 return mode_labels.iloc[0] standard_label_map = df.groupby(id_col)[label_col].agg(get_standard_label) df['standardized_label'] = df[id_col].map(standard_label_map) return df
验证效果
对于ID=222的分组,两种方案都会计算出"LA Metro"为计数最高的标签,映射后所有行的standardized_label将统一为"LA Metro",符合预期。
内容的提问来源于stack exchange,提问作者JLuu
相关产品推荐
相关产品推荐

