LayoutLMv3后处理方法仅返回前512 Token结果问题排查求助
LayoutLMv3后处理分块数据异常修复
问题背景
在LayoutLMv3的预处理-推理-后处理流程中出现异常:
- 预处理阶段对超过512 Token的序列设置128步长分块,
overflow_to_sample_mapping显示生成4个对应同一样本的chunk - 推理阶段成功处理全部2048个Token并返回预测结果
- 后处理方法仅能处理并返回前512个Token的结果,无法处理后续分块的Token数据
问题定位
后处理方法存在以下核心问题:
- 仅遍历单个chunk的bbox数据:原代码中
self._processed_data['bbox'].tolist()[page]只取了第page个chunk,但实际每个原始样本对应多个chunk,需通过overflow_to_sample_mapping找到所有属于当前样本的chunk索引 - 预测结果索引不匹配:推理输出是按所有chunk顺序排列的(每个chunk对应512个预测),原代码直接用
inference_output[j]只对应第一个chunk的预测,未关联后续chunk的结果 - 结果过滤参数错误:原代码中
filtered_docs = self.filterLabels([{'output': output_spans}])错误地传入了单个样本的结果,而非全部处理后的docs
修复后代码
def postprocess(self, inference_output): try: docs = [] k = 0 if isinstance(inference_output[0], list): inference_output = [item for sublist in inference_output for item in sublist] # 遍历每个原始样本 for page, doc_words in enumerate(self._raw_input_data['words']): doc_list = [] width, height = self._images_size[page] # 获取当前原始样本对应的所有chunk索引 chunk_indices = [i for i, idx in enumerate(self._overflow_to_sample_mapping.tolist()) if idx == page] for i, doc_word in enumerate(doc_words, start=0): word_tagging = None word_labels = [] word = dict() word['id'] = k k += 1 word['text'] = doc_word word['pageNum'] = page + 1 word['box'] = self._raw_input_data['bboxes'][page][i] _normalized_box = normalize_box(self._raw_input_data['bboxes'][page][i], width, height) # 遍历当前样本对应的所有chunk for chunk_idx in chunk_indices: # 获取当前chunk的bbox和offset mapping chunk_bboxes = self._processed_data['bbox'].tolist()[chunk_idx] chunk_offsets = self._offset_mapping.tolist()[chunk_idx] # 遍历chunk内的每个token,跳过padding的token(offset为[0,0]) for j, (box, offset) in enumerate(zip(chunk_bboxes, chunk_offsets)): if offset == [0, 0]: continue if compare_boxes(box, _normalized_box): # 计算当前token对应的推理输出索引:chunk_idx*512 + j pred_idx = chunk_idx * 512 + j label = self.model.config.id2label[inference_output[pred_idx]] if label != 'O': word_labels.append(label[2:]) else: word_labels.append('other') if word_labels: # 优先取非other的标签,若都是other则取最后一个 non_other = [l for l in word_labels if l != 'other'] word_tagging = non_other[0] if non_other else word_labels[-1] else: word_tagging = 'other' word['label'] = word_tagging word['pageSize'] = {'width': width, 'height': height} if word['label'] != 'other': doc_list.append(word) # 合并相邻实体的逻辑保持不变 spans = [] def adjacents(entity): return [adj for adj in doc_list if adjacent(entity, adj)] output_test_tmp = doc_list[:] for entity in doc_list: if not adjacents(entity): spans.append([entity]) output_test_tmp.remove(entity) while output_test_tmp: span = [output_test_tmp[0]] output_test_tmp = output_test_tmp[1:] while output_test_tmp and adjacent(span[-1], output_test_tmp[0]): span.append(output_test_tmp[0]) output_test_tmp.remove(output_test_tmp[0]) spans.append(span) output_spans = [] for span in spans: if len(span) == 1: output_span = { "text": span[0]['text'], "label": span[0]['label'], "words": [{ 'id': span[0]['id'], 'box': span[0]['box'], 'text': span[0]['text'] }] } else: output_span = { "text": ' '.join([entity['text'] for entity in span]), "label": span[0]['label'], "words": [{ 'id': entity['id'], 'box': entity['box'], 'text': entity['text'] } for entity in span] } output_spans.append(output_span) docs.append({'output': output_spans}) logger.debug(f"post-processing results: {docs}") # 修复参数错误,传入全部docs filtered_docs = self.filterLabels(docs) cleaned_docs = self.validate_fields(filtered_docs) ordered_docs = self.order_data_by_position(cleaned_docs) logger.info(f"Post-processing completed. {len(ordered_docs)} documents processed.") return [json.dumps(ordered_docs, ensure_ascii=False)] except Exception as e: logger.error(f"Error in postprocess: {e}") traceback.print_exc() raise e
修复关键点说明
- 关联所有对应chunk:通过
overflow_to_sample_mapping获取当前原始样本对应的所有chunk索引,确保遍历所有分块数据 - 匹配正确的预测结果:每个chunk对应512个预测结果,通过
chunk_idx * 512 + j计算当前token对应的推理输出索引 - 过滤padding token:利用
offset_mapping跳过offset为[0,0]的padding token,避免无效的bbox匹配 - 修复结果过滤参数:将
filtered_docs的参数改为处理后的完整docs列表,确保所有样本结果都被过滤处理
内容的提问来源于stack exchange,提问作者j3ws3r
相关产品推荐
相关产品推荐

