首页
/ Transformers 中的 BROS:面向文档关键信息抽取的布局感知预训练模型完全指南

Transformers 中的 BROS:面向文档关键信息抽取的布局感知预训练模型完全指南

2026-09-06 18:54:06作者:卓炯娓

BROS(BERT Relying On Spatiality)是专门为文档关键信息抽取(KIE)设计的 Encoder-only 布局感知 Transformer,其核心思路是用相对空间位置编码替代绝对坐标,结合 2D 区域掩码语言建模(AMLM)自监督目标,仅凭文本与版面信息即可在 FUNSD、SROIE、CORD、SciTSR 等 KIE 基准上取得有竞争力的结果。本文以官方模型文档 docs/source/en/model_doc/bros.md 为主体,结合本仓库 src/transformers/models/bros 下的完整源码,为你讲透 BROS 的建模原理、四个可加载模块的差异,以及 bbox 归一化、box_first_token_mask 构造、推理与损失计算等一整套可直接落地的实操细节。


1. 概览:什么是 BROS

BROS 于论文 BROS: A Pre-trained Language Model Focusing on Text and Layout for Better Key Information Extraction from Documents 中提出,由 jinho8345 于 2023-09-15 贡献进本仓库(对应 checkpoint 为 jinho8345/bros-base-uncased)。

模型名的含义是 BERT Relying On Spatiality。它是一个 Encoder-only Transformer:

  • 输入:一串 token 序列 + 每个 token 对应的边界框(bounding box)
  • 输出:一串隐藏状态(hidden states),可接各种下游头;
  • 核心思想:不像 LayoutLM 等模型那样直接使用绝对空间信息,而是编码文本在二维空间中的相对位置

在真实文档中,文本天然分布在二维页面上,传统的从左到右序列化会丢失版面结构,这正是文档 KIE 最大的难点之一。BROS 通过在 Transformer 自注意力中注入两两 token 之间的相对空间关系,让模型在阅读文本的同时感知"谁在谁的右边、谁在谁的下面",从而更好地完成实体抽取与实体间关系抽取。

论文摘要本身对这一动机作了精炼概括:从文档图像中做 KIE,需要理解文本在二维空间中的上下文语义与空间语义;BROS 从"文本与版面的有效结合"这一基本问题出发,用区域掩码策略从未标注文档上学习,在 FUNSD、SROIE、CORD、SciTSR 四个 KIE 基准上无需依赖视觉特征即可达到与既有方法相当或更好的效果,并直面 KIE 的两大真实挑战——文本排序错误带来的误差、以及下游标注样本稀少下的高效学习。

注意:BROS 不需要卷积视觉编码器抽取的图像特征,bounding box 通常由外部 OCR 系统(如 Tesseract、云 OCR 等)提供。


2. 两个预训练目标:TMLM 与 AMLM

BROS 使用两个目标联合预训练(模型文档概览):

  1. Token Masked Language Modeling(TMLM):即 BERT 经典的随机掩码语言建模。随机遮蔽一部分 token,模型利用空间信息与其余未掩码 token 来预测被掩码 token。空间信息在这里是关键差异点——被掩码 token 的 bbox 依然可见,模型需要学会"这个位置上的词最可能是什么"。

  2. Area Masked Language Modeling(AMLM):TMLM 的 2D 版本。它同样随机遮蔽文本 token 并用相同信息预测,但掩码的单位是**文本块(区域/area)**而非单个 token。AMLM 通常连续遮蔽落在同一文本区域内的多个 token,强迫模型在"整块区域消失"的情况下,借助周围文本与版面关系重建内容,这对单据、发票等由密集字段块构成的文档非常有效。

简单概括:TMLM 是 1D 的(token 粒度),AMLM 是 2D 的(区域粒度);两者的共同点是始终保留空间信息作为可利用的线索


3. 仓库中的模型结构:BrosModel 与三个任务头

在本仓库 src/transformers/models/bros/modeling_bros.py 中,BROS 被拆成 1 个主干 + 3 个任务头,均可从 checkpoint 直接加载:

3.1 BrosModel(主干)

标准 Encoder-only Transformer,结构上由三部分组成(见 modeling_bros.py):

  • BrosTextEmbeddings:与 BERT 一致的 word / position / token_type 三类 embedding 之和;
  • BrosBboxEmbeddings:BROS 特有的相对边界框位置编码器(详见第 4 节);
  • BrosEncoder + 可选的 BrosPooler

前向时必须同时提供 input_idsbbox,缺一不可——源码在 modeling_bros.py 处显式抛错:

  • input_idsinputs_embeds 二者恰好只给一个,抛出 "You must specify exactly one of input_ids or inputs_embeds";
  • bbox is None,抛出 "You have to specify bbox"。

3.2 BrosForTokenClassification(平面序列标注头)

BrosModel 之上加一个简单的线性分类层,对每个 token 预测标签(如每 token 的实体类型标签)。对应源码实现位于 modeling_bros.py,其分类器为:

self.dropout = nn.Dropout(classifier_dropout)          # classifier_dropout 或 hidden_dropout_prob
self.classifier = nn.Linear(config.hidden_size, config.num_labels)

3.3 BrosSpadeEEForTokenClassification(SPADE 实体抽取头)

SPADE(Spatial DAWG-based Entity extraction)思路下的实体抽取头,在 BrosModel 之上有两个组件(见 modeling_bros.py):

  • initial_token_classifier:用于预测每个实体的首 token,实现为两层 MLP(Dropout → Linear → Dropout → Linear);
  • subsequent_token_classifier:即 BrosRelationExtractor,用于在给定实体首 token 后预测实体内部的下一个 token

BrosForTokenClassificationBrosSpadeEEForTokenClassification 本质干的是同一件事(实体抽取),但差异关键:

  • BrosForTokenClassification 假设输入 token 已被完美序列化(在 2D 空间中排序,本身就很困难),一旦 OCR/阅读顺序出错,标签序列就被打乱;
  • BrosSpadeEEForTokenClassification 不依赖完美序列化——它以"实体首 token → 下一个 token → 再下一个 token"的链式指针方式预测实体,因此对序列化错误更鲁棒(详见其类文档说明与 modeling_bros.pycustom_intro)。

BrosSpadeEEForTokenClassification 前向需要两份标签 initial_token_labelssubsequent_token_labels,总损失为二者交叉熵之和:

loss = initial_token_loss + subsequent_token_loss

输出为 BrosSpadeOutput(定义于 modeling_bros.py):

字段 形状 含义
loss (1,) 初始 token 分类损失 + 后续 token 分类损失之和
initial_token_logits (batch_size, sequence_length, config.num_labels) 每个 token 作为"实体首 token"的得分(Softmax 前)
subsequent_token_logits (batch_size, sequence_length, sequence_length + 1) 每个 token 指向"实体内部下一 token"的得分(Softmax 前)
hidden_states / attentions 各层隐藏状态与注意力权重

注意 subsequent_token_logits 最后一维是 sequence_length + 1:多出的 1 列对应"实体到此结束"的终止位,模型学到在本 token 之后不再有同实体的 token。

3.4 BrosSpadeELForTokenClassification(SPADE 实体链接头)

BrosModel 之上加一个 entity_linker(同样是 BrosRelationExtractor,见 modeling_bros.py),用于执行实体内链接(intra-entity linking):预测"这个 token 与另一个 token 是否属于存在关系的两个实体",即输出 (N, num_relations, batch, seq_len, seq_len+1) 关系矩阵,打通跨实体的关系抽取。

四个类的加载方式(来自 modeling_bros.py 各类的 docstring 示例)统一为:

from transformers import BrosModel, BrosForTokenClassification
from transformers import BrosSpadeEEForTokenClassification, BrosSpadeELForTokenClassification

model = BrosModel.from_pretrained("jinho8345/bros-base-uncased")
# model = BrosForTokenClassification.from_pretrained("jinho8345/bros-base-uncased")
# model = BrosSpadeEEForTokenClassification.from_pretrained("jinho8345/bros-base-uncased")
# model = BrosSpadeELForTokenClassification.from_pretrained("jinho8345/bros-base-uncased")

4. 核心创新源码解析:相对空间位置如何进入注意力

BROS 与"把 bbox 当作额外特征拼进 embedding"的绝对位置方案不同,它把相对几何关系直接注入每一层自注意力。以 BrosSelfAttention.forward 中的关键片段(modeling_bros.py)为例:

attention_scores = torch.matmul(query_layer, key_layer.transpose(-1, -2))

# bbox positional encoding
bbox_pos_emb = bbox_pos_emb.view(seq_length, seq_length, batch_size, d_head)
bbox_pos_emb = bbox_pos_emb.permute([2, 0, 1, 3])
bbox_pos_scores = torch.einsum("bnid,bijd->bnij", (query_layer, bbox_pos_emb))

attention_scores = attention_scores + bbox_pos_scores

即:常规内容注意力分数之上,再叠加一项 query · bbox_pos_emb,而 bbox_pos_emb[i, j] 编码的是 token i 与 token j 边界框之间的相对几何关系。这一整套相对坐标编码在 modeling_bros.py 中由三个组件完成:

  1. BrosPositionalEmbedding1D:Transformer-XL 风格的 1D 正弦编码(源码注释直接引用 kimiyoung/transformer-xl 的实现),频率按 1 / 10000^(2k/d) 递减;
  2. BrosPositionalEmbedding2D:将 bbox 的 8 个坐标按奇偶位分别送入 x 与 y 方向的 1D 编码器——偶下标走 x_pos_emb、奇下标走 y_pos_emb
  3. BrosBboxEmbeddings:先做相对差变换 bbox_t[None, :, :, :] - bbox_t[:, None, :, :](shape 变为 (seq_len, seq_len, batch, dim_bbox)),再经正弦编码与一个无 bias 的 bbox_projection 线性层投影到与单头维度一致的向量。

关于 bbox 维度还有一个易被忽略的细节(在 BrosModel.forward 中,modeling_bros.py):

# if bbox has 2 points (4 float tensors) per token, convert it to 4 points (8 float tensors) per token
if bbox.shape[-1] == 4:
    bbox = bbox[:, :, [0, 1, 2, 1, 2, 3, 0, 3]]
scaled_bbox = bbox * self.config.bbox_scale
  • 若你传入每 token 只有 (x0, y0, x1, y1) 4 个值,模型会自动按 [x0, y0, x1, y0, x1, y1, x0, y1] 扩成 8 个值(重复左上/右上/右下/左下四个角,便于 x/y 交替编码);
  • 随后所有坐标乘以 bbox_scale(默认 100.0),把 0~1 归一化坐标放大到适合高频正弦编码的数值范围。

__init__.py 注册层面,四类均通过 base_model_prefix = "bros" 支持标准的 from_pretrained 加载与保存。


5. 使用指南:bbox 预处理与关键 mask

5.1 输入要求与坐标约定

BrosModel.forward 需要两个必需输入(modeling_bros.py):

  • input_ids:token id 序列;
  • bbox:形状 (batch_size, num_boxes, 4)(或 8)的边界框坐标,每个框为 (x0, y0, x1, y1),即左上角右下角
  • 坐标应按页面宽度归一化 x、按页面高度归一化 y,即取值范围在 0~1;
  • bbox 的来源取决于外部 OCR 系统(BrosProcessor 本身不做 OCR),需要你在送入模型前自行完成 OCR 版面解析。

模型文档给出的归一化辅助函数(原样继承,其中文档原文的循环体引用的是全局 width/height,按函数签名语义应使用入参 doc_width/doc_height,以下为修正后的可运行版本):

import numpy as np

def expand_and_normalize_bbox(bboxes, doc_width, doc_height):
    # bboxes 为 numpy array,每行形如 (x0, y0, x1, y1)
    # 归一化到 0 ~ 1
    bboxes[:, [0, 2]] = bboxes[:, [0, 2]] / doc_width
    bboxes[:, [1, 3]] = bboxes[:, [1, 3]] / doc_height
    return bboxes

5.2 推理示例

注意:BROS 的预训练权重在 [CLS] 等特殊 token 上同样需要 bbox,因此一个朴素但合法的做法是把整页的单位框 (0, 0, 1, 1) 广播给所有 token(源码 docstring 与 tests/models/bros/test_modeling_bros.py 的测试用例都采用了此类合法 bbox 构造)。完整推理流程:

import torch
from transformers import BrosProcessor, BrosModel

processor = BrosProcessor.from_pretrained("jinho8345/bros-base-uncased")
model = BrosModel.from_pretrained("jinho8345/bros-base-uncased")

encoding = processor("Hello, my dog is cute", add_special_tokens=False, return_tensors="pt")
bbox = torch.tensor([[[0, 0, 1, 1]]]).repeat(1, encoding["input_ids"].shape[-1], 1)
encoding["bbox"] = bbox

outputs = model(**encoding)
last_hidden_states = outputs.last_hidden_state

BrosForTokenClassificationBrosSpadeEEForTokenClassificationBrosSpadeELForTokenClassification 的前向签名与之相同,区别仅在于分别额外接收 labels / (initial_token_labels, subsequent_token_labels) / labels

5.3 训练损失关键:box_first_token_mask

BrosForTokenClassification.forwardBrosSpadeEEForTokenClassification.forwardBrosSpadeELForTokenClassification.forward 在做损失计算时不仅需要 input_idsbbox,还需要 box_first_token_maskmodeling_bros.py 对它的定义):

  • 1:该 token 是某个 bbox 的第一个 token(未被掩掉);
  • 0:非首 token。

为什么需要它?因为在文档场景中,一个"词/文本块"常被 tokenizer 切成多个子词 token(如 tokenizationtoken + ##ization),而实体标签通常标注在词的粒度。计算损失时我们希望只对每个 bbox 的首 token 计算交叉熵,避免对子词重复计损。

该 mask 无法由模型自行推导,必须在由 words 构造 input_ids 时记录每个 bbox 的起始 token 下标。模型文档给出了如下构造函数(原样继承;使用前需 import itertools):

import itertools
import numpy as np

def make_box_first_token_mask(bboxes, words, tokenizer, max_seq_length=512):
    box_first_token_mask = np.zeros(max_seq_length, dtype=np.bool_)

    # encode(tokenize) each word from words (list[str])
    input_ids_list = [tokenizer.encode(e, add_special_tokens=False) for e in words]

    # get the length of each box
    tokens_length_list = [len(l) for l in input_ids_list]

    box_end_token_indices = np.array(list(itertools.accumulate(tokens_length_list)))
    box_start_token_indices = box_end_token_indices - np.array(tokens_length_list)

    # filter out the indices that are out of max_seq_length
    box_end_token_indices = box_end_token_indices[box_end_token_indices < max_seq_length - 1]
    if len(box_start_token_indices) > len(box_end_token_indices):
        box_start_token_indices = box_start_token_indices[: len(box_end_token_indices)]

    # set box_start_token_indices to True
    box_first_token_mask[box_start_token_indices] = True

    return box_first_token_mask

工作原理分四步:

  1. 对每个 word 单独 tokenizer.encode(..., add_special_tokens=False),得到每个 box 的 token 数;
  2. itertools.accumulate 累计出每个 box 的结束下标,减去各自长度即得起始下标
  3. 截断超过 max_seq_length - 1 的索引(预留特殊 token 位置),并对齐两端长度;
  4. 在这些起始下标处置 True

为什么各头损失里要用它(源码行为对照)

  • BrosForTokenClassificationmodeling_bros.py):提供 mask 时,交叉熵只在 logits[bbox_first_token_mask]labels[bbox_first_token_mask] 上计算;不提供则对所有 token 计损;
  • BrosSpadeEEForTokenClassificationmodeling_bros.py):initial_token_loss 只在首 token 上计算;subsequent_token_loss 则对全部有效 token 计算,同时会用 attention_mask 屏蔽 padding,并用单位矩阵对角线屏蔽"自己指向自己";
  • BrosSpadeELForTokenClassificationmodeling_bros.py):只以首 token 作为 query去链接其他实体 token,因此损失按 bbox_first_token_mask 过滤 query 行,同时对非首 token 的 key 列与自环对角线做 masked_fill(填 -inf)。

6. 三个 Token 分类头的选型对比

上层结构 任务语义 对序列化错误的容忍度 所需标签
BrosForTokenClassification 单层 Linear 每个 token 的实体类别(平面 NER) 低:假定输入已完美按阅读顺序序列化 labels
BrosSpadeEEForTokenClassification initial_token_classifier + BrosRelationExtractor 判定实体首 token,并从首 token 出发链式预测实体内部下一个 token 高:通过"下一 token 指针"绕过排序误差(详见源码 modeling_bros.py initial_token_labels + subsequent_token_labels
BrosSpadeELForTokenClassification BrosRelationExtractorentity_linker 实体间关系 / 实体内链接:预测两实体 token 是否存在某关系 高:按实体对打分 labels(关系矩阵)

官方模型文档原话强调:BrosForTokenClassificationBrosSpadeEEForTokenClassification 干的是同一份活,但后者允许在序列化出错时拥有更大灵活性——这正是真实发票、表格识别中无法保证 OCR 阅读顺序时选择 SPADE 头的理由。

需要留意的是,EE 与 EL 头共享同一个 BrosRelationExtractor 结构(定义于 modeling_bros.py),内部含 n_relations 个关系维度的 query/key 线性层与一个可学习的 dummy 节点self.dummy_node = nn.Parameter(torch.zeros(...)))。dummy 节点在 key 侧被拼接到序列末尾,承担"空/结束"占位角色——EE 中它表示实体到此结束,EL 中它表示没有可链接的关系。


7. BrosConfig 配置参数

BrosConfig 定义于 src/transformers/models/bros/configuration_bros.py,除继承 PreTrainedConfig 的标准 BERT 类参数外,新增了三个 BROS 专属字段。以下是该文件声明的全部默认值(预训练 checkpoint 为 jinho8345/bros-base-uncased,即 BERT-base 体量):

7.1 常规 Transformer 参数(默认值)

参数 默认值 说明
vocab_size 30522 词表大小(uncased BERT 词表)
hidden_size 768 隐藏层维度
num_hidden_layers 12 Transformer 层数
num_attention_heads 12 注意力头数
intermediate_size 3072 FFN 中间层维度
hidden_act "gelu" FFN 激活函数
hidden_dropout_prob 0.1 隐藏层 dropout
attention_probs_dropout_prob 0.1 注意力 dropout
max_position_embeddings 512 最大序列长度(含特殊 token)
type_vocab_size 2 segment 类型数
initializer_range 0.02 参数初始化标准差
layer_norm_eps 1e-12 LayerNorm epsilon
pad_token_id 0 padding token id
classifier_dropout_prob 0.1 分类头 dropout

7.2 BROS 专属参数(建模关键)

参数 默认值 说明
dim_bbox 8 bbox 坐标维度。官方注释表述为 (x0, y1, x1, y0, x1, y1, x0, y1) 共 8 个值,即"四点坐标"版本;传入 4 点 (x0, y0, x1, y1) 时模型会自动展开(见第 4 节)
bbox_scale 100.0 bbox 坐标放大系数,归一化后的 0~1 坐标先乘以该值再进入正弦编码
n_relations 1 SPADE 的 EE/EL 头支持的关系种类数(如"上下""左右"等多类关系时调大)

7.3 __post_init__ 自动推导的派生维度

BROS 不需要使用者手动设置的三个内部维度在 BrosConfig.__post_init__configuration_bros.py)中由默认架构参数推导:

self.dim_bbox_sinusoid_emb_2d = self.hidden_size // 4           # 768//4 = 192:2D 正弦编码总维度
self.dim_bbox_sinusoid_emb_1d = self.dim_bbox_sinusoid_emb_2d // self.dim_bbox   # 192//8 = 24:单个方向/坐标通道维度
self.dim_bbox_projection = self.hidden_size // self.num_attention_heads         # 768//12 = 64:投影到单头维度

这三个派生值直接决定 BrosBboxEmbeddingsbbox_projection 的输出维度必须与每个注意力头维度一致,从而保证第 4 节的相对位置偏置能逐元素加到注意力分数上。

实例化与访问方式(来自配置文档的 docstring 示例):

from transformers import BrosConfig, BrosModel

configuration = BrosConfig()          # 初始化一个 jinho8345/bros-base-uncased 风格的配置
model = BrosModel(configuration)      # 从该配置初始化模型
configuration = model.config          # 访问模型配置

8. BrosProcessor:tokenizer 的轻量封装

BrosProcessor 定义于 src/transformers/models/bros/processing_bros.py,它继承 ProcessorMixin,本质是对底层 tokenizer 的标准化封装,本身不承担 OCR 与 bbox 生成职责:

  • 构造时必须传入 tokenizer,否则直接抛 ValueError("You need to specify a tokenizer.")
  • 通过 BrosProcessorKwargs 显式声明 text_kwargs 默认值,其中 add_special_tokens=Truepadding=Falsestride=0return_overflowing_tokens=False 等(processing_bros.py);
  • 使用方式与所有 AutoProcessor 一致:processor("...", add_special_tokens=False, return_tensors="pt") 得到可直接送入模型的 input_ids 等字段,再由用户自行把 bbox 写入 encoding["bbox"]

由于仓库代码风格上模型体系命名带 Bros 前缀,from transformers import BrosProcessor 也可经由 __init__.py 的导出直接导入。


9. 配套资源:权重转换脚本与测试

仓库在 BROS 模块下提供了可直接研读/复用的工程资源:

  • 权重转换脚本 src/transformers/models/bros/convert_bros_to_pytorch.py:负责把原版实现(import bros,即原代码仓库风格)的 checkpoint 重命名为 Hugging Face 版状态字典。其中明确移除了 embeddings.bbox_sinusoid_emb.inv_freq 一类由构造函数自动生成的 Buffer(其 inv_freq_init_weights 中重建),并把 embeddings.bbox_projection.weight 映射为 bbox_embeddings.bbox_projection.weight 等,可供复现/迁移自定义 checkpoint 时参考;
  • 单元测试 tests/models/bros/test_modeling_bros.py:覆盖 BrosConfig 的配置一致性测试(ConfigTester)、BrosModel/三个任务头的 ModelTesterMixin 通用前向与梯度测试、PipelineTesterMixin pipeline 测试等。测试侧的 BrosModelTester.prepare_config_and_inputs 是一个很有价值的参照:它演示了如何构造合法 bbox——对随机 bbox 强制执行 x1 >= x0y1 >= y0 的坐标约束(第 93-102 行),并生成全 1 的 bbox_first_token_mask 作为宽松默认值。做数据管线时建议参考其坐标合法性校验逻辑。

10. 小结

BROS 给文档 KIE 提供了一条"回到基本面"的路径:不用花哨的视觉编码器,而是在**文本 + 版面(相对空间)**的组合上做到极致。在 Transformers 中使用 BROS 的完整链路可以归纳为:

  1. OCR 拿版面:外部 OCR 产出 word 文本与 (x0, y0, x1, y1),按页宽高归一化到 0~1;
  2. 分词对齐:每个 word 独立 encode,构造 input_ids 的同时记录 box 起始下标,生成 box_first_token_mask
  3. 选择任务头:训练数据为干净序列化文本选 BrosForTokenClassification;OCR 排序不可靠时选 BrosSpadeEEForTokenClassification;需要跨实体关系时再加 BrosSpadeELForTokenClassification
  4. 前向与训练BrosModel 只要求 input_ids + bbox;各任务头在计算损失时按需传入 bbox_first_token_mask 与对应标签。

如果你想进一步深入,推荐按如下顺序阅读本仓库源码:configuration_bros.py(参数与派生维度)→ modeling_bros.py(相对位置注意力与各任务头)→ processing_bros.py(processor 约定)→ test_modeling_bros.py(合法 bbox 与 mask 的构造范式),即可完整掌握这一布局感知模型从原理到落地的全部细节。

登录后查看全文
热门项目推荐
相关项目推荐