如何基于微调后的LayoutLMv3对新图像进行真实推理?
LayoutLMv3 关键实体抽取真实推理实现步骤
1. 环境与模型准备
- 确保已安装依赖库:
transformers、torch、Pillow、pytesseract(用于OCR)、layoutparser(可选,用于结果可视化) - 加载微调后的模型与tokenizer:
from transformers import LayoutLMv3ForTokenClassification, LayoutLMv3TokenizerFast import torch # 替换为你的微调模型保存路径 model_dir = "./fine-tuned-layoutlmv3" tokenizer = LayoutLMv3TokenizerFast.from_pretrained(model_dir) model = LayoutLMv3ForTokenClassification.from_pretrained(model_dir) model.eval() # 切换至推理模式
2. 新图像预处理
提取图像中的文本、边界框及像素信息,直接使用tokenizer内置的OCR能力:
from PIL import Image # 加载目标图像 image = Image.open("./new-document.jpg").convert("RGB") # 调用tokenizer完成OCR与编码 encoding = tokenizer( image, return_tensors="pt", truncation=True, padding="max_length", max_length=512 # 与训练时的max_length保持一致 ) input_ids = encoding["input_ids"] bbox = encoding["bbox"] pixel_values = encoding["pixel_values"]
3. 执行推理计算
关闭梯度计算节省内存,获取模型预测结果:
with torch.no_grad(): outputs = model(input_ids=input_ids, bbox=bbox, pixel_values=pixel_values) logits = outputs.logits # 将logits转换为预测标签ID predictions = torch.argmax(logits, dim=-1)
4. 解析与整理结果
将预测的标签ID映射回实体名称,并处理子词拼接:
# 替换为你训练时的标签映射(需与LabelStudio标注的类别一致) label2id = {"O": 0, "PERSON": 1, "ADDRESS": 2, "PHONE": 3} id2label = {v: k for k, v in label2id.items()} # 转换token并过滤特殊符号 tokens = tokenizer.convert_ids_to_tokens(input_ids[0]) predicted_labels = [id2label[p.item()] for p in predictions[0]] # 整理实体结果,处理BPE子词 final_entities = [] current_entity = None for token, label in zip(tokens, predicted_labels): if token in ["<s>", "</s>", "<pad>"]: continue # 处理以##开头的子词 if token.startswith("##"): if current_entity: current_entity["text"] += token[2:] else: if current_entity: final_entities.append(current_entity) current_entity = {"text": token, "entity_type": label} if current_entity: final_entities.append(current_entity) # 输出整理后的结果 for entity in final_entities: print(f"实体类型: {entity['entity_type']}, 内容: {entity['text']}")
5. 结果可视化(可选)
用layoutparser在原图上标注实体区域,直观查看推理效果:
import layoutparser as lp # 生成边界框与标签 boxes = [] for i in range(len(tokens)): if tokens[i] in ["<s>", "</s>", "<pad>"]: continue box = lp.BBox(*bbox[0][i].tolist(), label=predicted_labels[i]) boxes.append(box) # 绘制标注图像并保存 annotated_image = lp.draw_box(image, boxes, box_width=2, text_color="red") annotated_image.save("./annotated-result.jpg")
内容的提问来源于stack exchange,提问作者Montassar Jaziri
相关产品推荐
相关产品推荐

