SimpleTransformers的LayoutLM模型预测报错,如何正确定义输入数据?
错误原因
你当前使用的是LayoutLM文档理解模型,该模型需要同时输入文本、文本对应在文档中的边界框坐标、原始文档图像三类数据才能正常推理,你仅传入了文本字符串列表,不符合输入格式要求,框架解析输入时类型匹配失败触发了报错。
正确输入定义步骤
- 首先准备待预测的文档图像,以及该图像的OCR识别结果(包含每个文本块的内容、对应像素坐标)
- 对OCR得到的像素坐标做归一化处理,所有坐标值映射到0~1000的整数范围
- 按要求构造字典格式的输入样本
完整代码示例
from PIL import Image import pytesseract # 可选,你也可以用任意其他OCR工具获取文本和坐标 # 加载待预测文档图像 test_image = Image.open("你的文档路径.png").convert("RGB") img_width, img_height = test_image.size # ========== 以下为OCR获取文本和坐标示例,替换为你自己的实际数据 ========== ocr_data = pytesseract.image_to_data(test_image, output_type=pytesseract.Output.DICT) text_list = [] bbox_list = [] for i in range(len(ocr_data["text"])): text = ocr_data["text"][i].strip() if not text: continue text_list.append(text) # 归一化坐标到0~1000 x0 = int(ocr_data["left"][i] * 1000 / img_width) y0 = int(ocr_data["top"][i] * 1000 / img_height) x1 = int((ocr_data["left"][i] + ocr_data["width"][i]) * 1000 / img_width) y1 = int((ocr_data["top"][i] + ocr_data["height"][i]) * 1000 / img_height) bbox_list.append([x0, y0, x1, y1]) # ======================================================================= # 构造符合要求的预测输入 predict_input = [ { "text": " ".join(text_list), "bboxes": bbox_list, "image": test_image } ] # 调用预测 predictions, raw_outputs = model.predict(predict_input)
替代方案(无需布局能力时可选)
如果你仅需要做纯文本分类,不需要用到LayoutLM的文档布局、图像特征能力,可以直接将模型初始化时的模型类型从layoutlm改为bert,即可直接使用你原来的字符串列表格式输入:
model = ClassificationModel( "bert", "bert-base-uncased", num_labels=2, use_cuda=True, cuda_device = 0 ) # 原输入格式可正常运行 predictions, raw_outputs = model.predict(['test data abc'])
内容的提问来源于stack exchange,提问作者K..
相关产品推荐
相关产品推荐

