HuggingFace Transformers NER模型单样本预测标签转换问题
问题根因
你加载模型时使用了通用的AutoModel类,这个类的输出是DistilBert最后一层的隐藏状态特征,你看到的768就是模型的隐藏层维度,不是NER任务的分类结果。所以你对768维取最大值索引,得到的自然是0-767之间的数值,和你要的0-8标签完全不匹配。
修复步骤
1. 替换模型加载类
改用专为令牌分类任务设计的AutoModelForTokenClassification类加载模型,它会自动读取你训练时保存的分类头,输出对应9个标签的预测结果:
# 替换原有AutoModel导入和加载逻辑 from transformers import AutoModelForTokenClassification model_checkpoint = "distilbert-base-uncased" tokenizer = AutoTokenizer.from_pretrained(model_checkpoint) loaded_model = AutoModelForTokenClassification.from_pretrained( './my_model_own_custom_training.pth', from_tf=False )
修改完成后,你得到的predictions形状会变成(1,13,9),此时执行predictions.argmax(axis=2)就能得到0-8范围内的标签索引。
2. 标签映射与实体对齐
首先你需要把配置文件里默认的LABEL_0这类占位符,替换成你训练时实际对应的业务标签,示例如下:
# 按你训练时的实际标签定义修改即可 id2label = { 0: "O", 1: "B-PER", 2: "I-PER", 3: "B-LOC", 4: "I-LOC", 5: "B-ORG", 6: "I-ORG", 7: "B-MISC", 8: "I-MISC" }
接下来做子词和原单词的对齐,避免子词拆分导致的标签重复:
# 处理输入得到预测标签 input_sentence = "John Nash is a great mathematician, he lives in France" tokenized_input = tokenizer( input_sentence, truncation=True, return_tensors='pt' ) outputs = loaded_model(**tokenized_input) predictions = outputs.logits.argmax(dim=2)[0].tolist() pred_labels = [id2label[p] for p in predictions] # 对齐原单词和标签 tokens = tokenizer.convert_ids_to_tokens(tokenized_input["input_ids"][0]) word_ids = tokenized_input.word_ids() prev_word_id = None entity_result = [] for token, label, word_id in zip(tokens, pred_labels, word_ids): # 跳过[CLS]、[SEP]这类特殊标记 if word_id is None: continue # 同一个单词仅保留第一个子词的标签,合并子词 if word_id != prev_word_id: if token.startswith("##"): entity_result[-1] = (entity_result[-1][0] + token[2:], entity_result[-1][1]) else: entity_result.append((token, label)) prev_word_id = word_id
最终entity_result就是你要的[('John', 'B-PER'), ('Nash', 'I-PER'), ..., ('France', 'B-LOC')]格式的结果。
3. 简化方案(可选)
你可以直接使用pipeline接口自动完成标签映射、子词合并的逻辑,不用手动处理:
from transformers import pipeline ner_pipe = pipeline( task="ner", model=loaded_model, tokenizer=tokenizer, aggregation_strategy="simple" ) print(ner_pipe(input_sentence))
直接输出合并完成的实体组、对应标签和置信度。
内容的提问来源于stack exchange,提问作者Timbus Calin
相关产品推荐
相关产品推荐

