如何从HuggingFace Token分类模型输出中获取正确预测文本?
问题原因与修正方法
你的核心错误是用tokenizer解码NER标签ID——模型输出的predicted_labels是实体标签的索引,不是分词后的token ID,直接用tokenizer.batch_decode()会把这些索引对应到tokenizer词汇表中的未使用占位符(也就是[unusedXX]),完全不符合预期。
正确的处理流程是:
- 用模型自带的
id2label映射表,把标签ID转换成对应的实体标签(比如B-PERSON、I-DATE这类) - 结合原文本的分词结果,识别敏感实体并替换,完成去标识
修正后的完整代码
import torch from transformers import AutoTokenizer, AutoModelForTokenClassification tokenizer = AutoTokenizer.from_pretrained("obi/deid_bert_i2b2", do_lower_case=True) model = AutoModelForTokenClassification.from_pretrained("obi/deid_bert_i2b2") id2label = model.config.id2label # 获取标签ID到实体标签的映射 text = "Patient John Doe visited the hospital on 01/05/2023 with complaints of chest pain." # 编码文本,同时保留原始token的位置信息(处理subword拆分) encoded_input = tokenizer(text, padding=True, return_tensors='pt', return_offsets_mapping=True) offsets = encoded_input.pop('offset_mapping').squeeze() # 移除offset_mapping,避免传入模型 outputs = model(**encoded_input) predicted_labels = torch.argmax(outputs.logits, dim=2).squeeze() # 将标签ID转换为实体标签 predicted_entities = [id2label[label_id.item()] for label_id in predicted_labels] # 构建去标识后的文本 deidentified_text = [] current_entity = None current_start = 0 for idx, (entity, offset) in enumerate(zip(predicted_entities, offsets)): start, end = offset # 跳过特殊token(比如[CLS]、[SEP]) if start == 0 and end == 0: continue # 处理实体开始(B-开头的标签) if entity.startswith("B-"): # 如果之前有未处理的普通文本,先添加 if current_start < start: deidentified_text.append(text[current_start:start]) # 记录当前实体类型 current_entity = entity.split("-")[1] current_start = end # 处理实体中间(I-开头的标签) elif entity.startswith("I-") and current_entity is not None: current_start = end # 处理非实体或实体结束 else: if current_entity is not None: # 添加实体占位符 deidentified_text.append(f"[{current_entity.upper()}]") current_entity = None # 添加普通文本 deidentified_text.append(text[current_start:end]) current_start = end # 处理最后一段未添加的内容 if current_entity is not None: deidentified_text.append(f"[{current_entity.upper()}]") elif current_start < len(text): deidentified_text.append(text[current_start:]) # 拼接最终结果 deidentified_text = "".join(deidentified_text) print(deidentified_text)
代码说明
return_offsets_mapping=True:获取每个token在原文本中的起始和结束位置,解决BERT分词的subword拆分问题(避免把一个实体拆成多个subword后无法合并)id2label映射:直接用模型配置中的映射表,把标签ID转成B-PERSON这类标准NER标签- 实体合并与替换:遍历每个token的实体标签,将连续的实体(B-开头+后续I-开头)合并,替换成
[PERSON]、[DATE]这类占位符,非实体内容直接保留
预期输出
Patient [PERSON] visited the hospital on [DATE] with complaints of chest pain.
内容的提问来源于stack exchange,提问作者learner
相关产品推荐
相关产品推荐

