多标签分类去除classes_中多余空格、逗号类别的方法
问题根因
出现多余标签的核心原因有两个:
- 你当前使用的
LabelBinarizer是单标签多分类场景的二值化工具,不适用于多标签分类任务 - 数据集的
tags列原始值为字符串格式(例如"1, 2"、" 3 "这类带逗号、空格的拼接字符串),没有预先解析为标签列表,二值化工具会把单个字符、符号识别为独立标签,因此混入了空格、逗号等无效类别。
解决步骤
- 第一步:预处理标签列,将字符串格式的标签清洗为仅含有效类别的列表
def parse_tags(tag_str): # 按逗号分割标签,去除每个标签前后的空白字符 tag_list = [t.strip() for t in str(tag_str).split(',')] # 仅保留1/2/3/4四类有效标签,过滤空值和无效值 return [t for t in tag_list if t in {'1','2','3','4'}] # 对标签列做清洗 df['tags'] = df['tags'].apply(parse_tags) y = df['tags']
- 第二步:使用
MultiLabelBinarizer做多标签二值化,可手动指定有效类别从根源避免混入无效标签
# 手动指定classes为你需要的4个标签,不需要再额外过滤 multilabel = MultiLabelBinarizer(classes=['1','2','3','4']) y = multilabel.fit_transform(y) # 此时输出的classes仅会包含['1','2','3','4'] print(multilabel.classes_)
后续代码调整
原来的代码中你直接用了LabelBinarizer输出的y做数据集分割,替换为上述处理后的二值化y即可,其余模型训练、评估逻辑不需要调整。
内容的提问来源于stack exchange,提问作者hideonbush
相关产品推荐
相关产品推荐

