如何从LayoutLMV1模型获取问题为键答案为值的字典键值对输出
LayoutLM推理输出问答键值对实现方案
场景说明
- 基于FUNSD数据集微调的LayoutLMForTokenClassification模型做文档结构化推理,需要输出问题为键、对应答案为值的结构化键值对结果
- 参考实现为LayoutLM token分类任务在FUNSD数据集上的微调教程
- 原有调试代码仅做了标签、坐标、文本的逐元素打印,未做实体聚合和问答配对:
layout_details = [] for prediction, box in zip(true_predictions, true_boxes): predicted_label = iob_to_label(prediction).lower() layout_details.append((predicted_label, prediction, box, label2color[predicted_label])) for i, j in zip(words[0], layout_details[1:-1]): print(i, j)
实现逻辑
FUNSD数据集的实体标签分为question/answer/header/other四类,要得到标准键值对需要两步:
- 把相邻、标签一致的token拼接成完整的问题/答案文本块,同时记录每个文本块的坐标范围
- 基于文档排版规则匹配问答对:常规表单排版中,问题和答案y轴坐标接近(同一行或相邻行),答案位于问题右侧,x轴间隔在合理范围内
参考实现代码
def aggregate_entities(words, pred_ids, boxes): """拼接相邻同标签token,得到完整的问题、答案实体块""" entities = [] current_entity = None for word, pred_id, box in zip(words, pred_ids, boxes): label = iob_to_label(pred_id).lower() # 跳过非问答类实体 if label not in ["question", "answer"]: if current_entity is not None: entities.append(current_entity) current_entity = None continue # 遇到新类型实体就保存上一个实体,新建当前实体 if current_entity is None or current_entity["label"] != label: if current_entity is not None: entities.append(current_entity) current_entity = { "label": label, "text": word, "bbox": box } # 同标签相邻token合并文本,更新边界框范围 else: current_entity["text"] += " " + word current_entity["bbox"] = [ min(current_entity["bbox"][0], box[0]), min(current_entity["bbox"][1], box[1]), max(current_entity["bbox"][2], box[2]), max(current_entity["bbox"][3], box[3]) ] # 加入最后一个实体 if current_entity is not None: entities.append(current_entity) return entities def match_qa_pairs(entities, y_tolerance=15, x_max_gap=200): """基于坐标位置匹配问题和对应的答案,返回键值对字典""" qa_dict = {} questions = [e for e in entities if e["label"] == "question"] answers = [e for e in entities if e["label"] == "answer"] for q in questions: q_mid_y = (q["bbox"][1] + q["bbox"][3]) / 2 best_match = None min_gap = float("inf") for a in answers: a_mid_y = (a["bbox"][1] + a["bbox"][3]) / 2 # y轴偏差在阈值内视为同一行组 if abs(q_mid_y - a_mid_y) <= y_tolerance: # 答案在问题右侧时计算x轴间隔 if a["bbox"][0] >= q["bbox"][2]: x_gap = a["bbox"][0] - q["bbox"][2] if x_gap < x_max_gap and x_gap < min_gap: min_gap = x_gap best_match = a if best_match: qa_dict[q["text"].strip()] = best_match["text"].strip() answers.remove(best_match) # 避免答案重复匹配 return qa_dict
调用方式
推理时取单样本数据,过滤掉首尾特殊token(和原有代码[1:-1]的切片逻辑一致)后调用即可:
# 替换为推理得到的单样本对应变量 sample_words = words[0] # 过滤特殊token对应的预测结果和坐标 sample_preds = true_predictions[1:-1] sample_boxes = true_boxes[1:-1] # 生成最终问答键值对 entities = aggregate_entities(sample_words, sample_preds, sample_boxes) qa_result = match_qa_pairs(entities) print(qa_result)
可根据实际文档排版调整
y_tolerance(y轴偏差阈值)和x_max_gap(x轴最大间隔)两个参数,提升配对准确率。
内容的提问来源于stack exchange,提问作者Jyoti yadav
相关产品推荐
相关产品推荐

