You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

LayoutLMv3后处理方法仅返回前512 Token结果问题排查求助

LayoutLMv3后处理分块数据异常修复

问题背景

在LayoutLMv3的预处理-推理-后处理流程中出现异常:

  • 预处理阶段对超过512 Token的序列设置128步长分块,overflow_to_sample_mapping显示生成4个对应同一样本的chunk
  • 推理阶段成功处理全部2048个Token并返回预测结果
  • 后处理方法仅能处理并返回前512个Token的结果,无法处理后续分块的Token数据

问题定位

后处理方法存在以下核心问题:

  1. 仅遍历单个chunk的bbox数据:原代码中self._processed_data['bbox'].tolist()[page]只取了第page个chunk,但实际每个原始样本对应多个chunk,需通过overflow_to_sample_mapping找到所有属于当前样本的chunk索引
  2. 预测结果索引不匹配:推理输出是按所有chunk顺序排列的(每个chunk对应512个预测),原代码直接用inference_output[j]只对应第一个chunk的预测,未关联后续chunk的结果
  3. 结果过滤参数错误:原代码中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

修复关键点说明

  1. 关联所有对应chunk:通过overflow_to_sample_mapping获取当前原始样本对应的所有chunk索引,确保遍历所有分块数据
  2. 匹配正确的预测结果:每个chunk对应512个预测结果,通过chunk_idx * 512 + j计算当前token对应的推理输出索引
  3. 过滤padding token:利用offset_mapping跳过offset为[0,0]的padding token,避免无效的bbox匹配
  4. 修复结果过滤参数:将filtered_docs的参数改为处理后的完整docs列表,确保所有样本结果都被过滤处理

内容的提问来源于stack exchange,提问作者j3ws3r

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.13 13:53:09