使用HuggingFace Transformers的LayoutLMv2提取文本形式信息
LayoutLMv2 提取文本答案实现方法
你现有代码已经完成了模型推理的核心逻辑,只需要补充预测位置解码步骤就能拿到文本形式的答案,不需要额外做OCR或者检测操作,具体实现如下:
- 首先移除推理代码里手动传入的
start_positions、end_positions参数。这两个参数是模型微调阶段用来计算损失的人工标注标签,纯推理场景不需要传入,否则会额外计算训练损失,属于冗余操作。 - 从模型输出的
start_logits、end_logits中取概率最高的位置,作为预测答案在token序列中的起止索引,注意要过滤结束索引小于起始索引的非法区间。 - 利用处理器自带的token解码器,把对应索引区间的token还原成自然文本,就是你需要的提取结果。
修正后的完整可运行代码:
from transformers import LayoutLMv2Processor, LayoutLMv2ForQuestionAnswering from PIL import Image import torch processor = LayoutLMv2Processor.from_pretrained("microsoft/layoutlmv2-base-uncased") model = LayoutLMv2ForQuestionAnswering.from_pretrained("microsoft/layoutlmv2-base-uncased") # 替换成你的发票图片路径 image = Image.open("invoice_sample.png").convert("RGB") # 替换成你要提取的信息对应的问题,比如"发票金额是多少"、"销售方名称是什么" question = "what's the invoice number?" encoding = processor(image, question, return_tensors="pt") # 推理时关闭梯度计算减少内存占用 with torch.no_grad(): outputs = model(**encoding) # 拿到预测答案的起止token索引 start_idx = torch.argmax(outputs.start_logits, dim=1).item() end_idx = torch.argmax(outputs.end_logits, dim=1).item() # 处理非法区间 if end_idx < start_idx: end_idx = start_idx # 解码得到纯文本答案 answer_token_ids = encoding.input_ids[0][start_idx:end_idx+1] extracted_text = processor.tokenizer.decode(answer_token_ids, skip_special_tokens=True) print(f"提取结果:{extracted_text}")
补充说明:
- 公开的基础预训练权重没有针对发票场景做专项微调,直接拿来提取发票字段准确率会很低,要达到可用的效果需要用标注好的发票问答数据对模型做微调。
- 如果需要同时获取答案对应的坐标框来做可视化,直接读取
encoding.bbox[0][start_idx:end_idx+1]里的坐标值,还原到原图尺寸即可绘制标注框,和现有示例的可视化逻辑完全兼容。
内容的提问来源于stack exchange,提问作者An old man in the sea.
相关产品推荐
相关产品推荐

