BridgeTower 完整指南:借助桥接层打通视觉与文本双塔编码器(Transformers 实现解读)
本指南以 docs/source/en/model_doc/bridgetower.md 为核心主体,结合 modeling_bridgetower.py、configuration_bridgetower.py 与 processing_bridgetower.py 等仓库源码展开,帮助你在 🤗 Transformers 中完整掌握 BridgeTower 的架构原理、配置体系、Processor 用法与三类下游任务(图文对比学习、图文检索、掩码语言建模)的落地代码。
BridgeTower 是论文 BridgeTower: Building Bridges Between Encoders in Vision-Language Representative Learning 提出的视觉-语言(VL)预训练模型,也是 🤗 Transformers 官方维护的模型架构之一(model_type = "bridgetower")。其核心创新在于:在单模态编码器(视觉塔 + 文本塔)与跨模态编码器之间引入多组轻量 bridge layers(桥接层),让跨模态编码器的每一层都能直接"看到"来自双塔各语义层级的信息,从而实现自下而上的细粒度跨模态对齐与融合——仅在 400 万张图上预训练,就在 VQAv2 test-std 上达到 78.73% 的准确率,以几乎可忽略的参数与计算增量超越了此前同数据规模的 SOTA 模型 METER(1.09%)(上述指标来自论文原文)。本文将带你从论文动机一路走到可直接运行的推理代码。
一、概述:BridgeTower 要解决什么问题
传统视觉-语言模型普遍采用 TWO-TOWER(双塔)架构,并在其上演化出两类设计:
- 轻量单模态编码器 + 深层跨模态编码器:编码器本身很浅,对齐与融合完全依赖深层跨模态编码器完成,受限于单模态特征的表达能力;
- 预训练好的深层单模态编码器 + 顶层跨模态编码器:只用单模态编码器的最后一层表示喂给顶层的跨模态编码器,中间层的丰富语义被丢弃。
BridgeTower 论文指出,这两种做法都限制了视觉-语言表示学习,进而提出第三种思路——桥接:
BridgeTower introduces multiple bridge layers that build a connection between the top layers of uni-modal encoders and each layer of the crossmodal encoder.
即:用若干桥接层,将预训练单模态编码器的多个顶层与跨模态编码器的每一层两两相连,把不同语义层级的视觉、文本表示持续注入跨模态融合过程(论文称其为 effective bottom-up cross-modal alignment and fusion)。该论文发表于 AAAI'23,本文档对应的 HF paper 发布于 2022-06-17,模型于 2023-01-25 贡献给 Hugging Face Transformers(文档中给出此信息),实现由 Anahita Bhiwandiwalla、Tiep Le 与 Shaoyen Tseng 贡献。
二、架构与源码级原理:双塔如何被"桥"接起来
BridgeTower 由三大部分组成:视觉编码器、文本编码器、跨模态编码器,外加多组轻量桥接层。仓库实现对其做了明确说明(见 modeling_bridgetower.py 中 BridgeTowerModel.__init__):
self.vision_model = BridgeTowerVisionModel(vision_config):视觉塔;self.text_model = BridgeTowerTextModel(text_config):文本塔;cross_modal_image_layers/cross_modal_text_layers:跨模态编码器(双层自注意力 + 交叉注意力,代码基于 RoBERTa/BERT 结构实现);cross_modal_text_link_tower/cross_modal_image_link_tower:真正的"桥"(BridgeTowerLinkTower)。
2.1 桥接层的三种融合公式
桥接层在 modeling_bridgetower.py 中实现为 BridgeTowerLinkTower,其融合方式由配置项 link_tower_type 控制,共支持三种:
"add"(默认):LayerNorm(hidden_states + cross_modal_hidden_states),直接相加后归一化;"scaled_add":LayerNorm(hidden_states * scaled_factor + cross_modal_hidden_states),为塔侧输入引入一个可学习标量scaled_factor(初始为 1.0);"interpolate":LayerNorm(hidden_states * (1 - beta) + cross_modal_hidden_states * beta),用可学习参数beta(初始 0.5)在塔特征与跨模态特征之间插值。
从实现上看,桥接层输出的正是"单模态最新表示 + 跨模态当前表示"的归一化融合结果,随后被送入下一层跨模态自/交叉注意力,形成逐层的双向信息流动。仓库在 forward 中对这段设计有一句非常关键的注释:
Each of the top 6 layers of the visual and textual encoders is connected to each layer of the cross-modal encoder via bridge layers, which brings bottom-up alignment and fusion to the cross-modal encoder.
2.2 切分点 split_index:多少层留给单塔、多少层参与跨模态
BridgeTowerModel.forward 中计算了关键变量 split_index(见 modeling_bridgetower.py):
split_index = len(self.text_model.encoder.layer) - self.config.num_hidden_layers + 1
以默认配置(文本/视觉塔各 12 层、跨模态编码器 num_hidden_layers = 6)为例,split_index = 12 - 6 + 1 = 7。整个前向流程分三段(源码对应 modeling_bridgetower.py):
- 前 7 层:文本与图像各自独立走单模态编码器(
layer[:split_index]),不做任何跨模态交互; - 第一层跨模态层特例:由于此时还没有任何跨模态输出可供桥接,第一层直接对单模态输出做线性变换(
cross_modal_text/image_transform)、加上可学习的 token type embedding、过 LayerNorm 后进入跨模态层; - 后 5 层桥接循环:从
i = split_index开始,每推进一层单模态编码器,就调用一次text_link_tower/image_link_tower,把"该层单模态输出"与"上一层跨模态输出"融合,再送入下一层跨模态编码器,如此逐层把跨模态编码器与双塔顶层"缝"在一起。
最后 get_cls_features(modeling_bridgetower.py)分别用文本、图像两个 pooler 取出 [CLS] 表征并经 Tanh 激活后 在最后一维拼接(维度为 2 * hidden_size),作为联合表示。
2.3 输出数据结构
BridgeTowerModel.forward 返回 BridgeTowerModelOutput(modeling_bridgetower.py),字段为:
text_features:形状(batch_size, text_seq_len, hidden_size),跨模态文本输出;image_features:形状(batch_size, image_seq_len, hidden_size),跨模态图像输出;pooler_output:形状(batch_size, hidden_size * 2),文本与图像 [CLS] 拼接并经 pooler 处理后的联合表征;- 另有可选的
hidden_states、attentions供分析使用。
三、配置体系:三个 Config 类的默认值与作用
BridgeTower 的配置拆成三个类,全部定义在 configuration_bridgetower.py 中,并用 sub_configs 将文本、视觉子配置挂载到总配置下:
3.1 BridgeTowerVisionConfig(视觉塔,model_type = "bridgetower_vision_model")
对应 CLIP/ViT 风格的视觉塔,关键默认值(源码 configuration_bridgetower.py):
| 参数 | 默认值 | 含义 |
|---|---|---|
hidden_size |
768 |
隐藏维度 |
num_hidden_layers |
12 |
视觉 Transformer 层数 |
num_channels |
3 |
输入通道数(RGB) |
patch_size |
16 |
ViT patch 尺寸 |
image_size |
288 |
输入图像分辨率 |
layer_norm_eps |
1e-05 |
LayerNorm epsilon |
stop_gradient |
False |
训练时是否对该塔梯度截断 |
share_layernorm |
True |
是否共享 LayerNorm |
remove_last_layer |
False |
是否移除视觉编码器最后一层 |
3.2 BridgeTowerTextConfig(文本塔,model_type = "bridgetower_text_model")
文本塔直接沿用 RoBERTa 的配置骨架(源码 configuration_bridgetower.py):vocab_size = 50265、hidden_size = 768、num_hidden_layers = 12、num_attention_heads = 12、intermediate_size = 3072、hidden_act = "gelu"、max_position_embeddings = 514、pad/bos/eos_token_id = 1/0/2,dropout 均为 0.1。
3.3 BridgeTowerConfig(总配置,model_type = "bridgetower")
总配置负责"桥"的开关与结构(源码 configuration_bridgetower.py):
| 参数 | 默认值 | 含义 |
|---|---|---|
num_hidden_layers |
6 |
跨模态编码器层数 |
hidden_size |
768 |
跨模态隐藏维度 |
num_attention_heads |
12 |
注意力头数 |
share_cross_modal_transformer_layers |
True |
是否共享跨模态层权重 |
share_link_tower_layers |
False |
是否共享桥接层权重 |
link_tower_type |
"add" |
桥接融合方式,见 2.1 节 |
init_layernorm_from_vision_encoder |
False |
是否用视觉编码器初始化 LayerNorm |
tie_word_embeddings |
False |
是否共享词嵌入 |
text_config / vision_config |
None |
传入后自动实例化对应子 Config |
__post_init__ 中(configuration_bridgetower.py)会处理 text_config、vision_config:为 None 时用默认值初始化并打日志;为 dict 时展开为对应 Config 类。手工构造时可直接套嵌,例如:
from transformers import BridgeTowerConfig, BridgeTowerModel
config = BridgeTowerConfig(
num_hidden_layers=6,
text_config={"num_hidden_layers": 12, "hidden_size": 768},
vision_config={"num_hidden_layers": 12, "image_size": 288, "patch_size": 16},
)
model = BridgeTowerModel(config)
四、Processor:如何同时编码文本与图像
官方建议通过 [BridgeTowerProcessor] 一站式完成预处理。它在 processing_bridgetower.py 中被定义为 ProcessorMixin 子类,构造函数仅接收两个组件:
RobertaTokenizer(Fast 版)——负责把文本编码为input_ids、attention_mask;BridgeTowerImageProcessor——负责把图像转为pixel_values(以及pixel_mask)。
BridgeTowerProcessor 还带有一套默认 kwargs(源码 BridgeTowerProcessorKwargs):文本侧默认 add_special_tokens=True、padding=False、stride=0、return_overflowing_tokens=False、return_special_tokens_mask=False、return_offsets_mapping=False、return_length=False、verbose=True;图像侧默认 do_normalize=True、do_center_crop=True。不传参数时即采用这套默认行为。
4.1 图像预处理细节
BridgeTowerImageProcessor 基于 torchvision 后端实现(仓库另有 BridgeTowerImageProcessorPil 提供 PIL 后端,见 image_processing_pil_bridgetower.py),默认值定义在 image_processing_bridgetower.py:
size = {"shortest_edge": 288},crop_size = {"shortest_edge": 288}:短边先缩放到 288,再中心裁剪;resample = BICUBIC;image_mean = OPENAI_CLIP_MEAN、image_std = OPENAI_CLIP_STD:复用 OpenAI CLIP 的归一化统计量,与 CLIP/ViT 视觉塔的设计保持一致;size_divisor = 32:resize 时把宽高规整为 32 的整数倍(get_resize_output_image_size内new_height // size_divisor * size_divisor,见 image_processing_bridgetower.py),保证后续 patch 化与注意力掩码对齐;model_input_names = ["pixel_values", "pixel_mask"]。
五、实战:三类下游任务的完整代码
本实现支持三类常见视觉-语言任务,每个任务对应一个头部模型。下面示例均采用本地生成图像的方式,避免依赖外部图片地址,可在离线环境直接运行(语义与官方文档示例保持一致)。若需真实图片,可用任意本地图文件替代 Image.fromarray 生成的张量。
5.1 图文对比学习(BridgeTowerForContrastiveLearning)
BridgeTowerForContrastiveLearning 返回 BridgeTowerContrastiveOutput(modeling_bridgetower.py),其中 text_embeds、image_embeds、cross_embeds 分别对应投影层处理后的文本、图像与跨模态嵌入,可选返回 loss(图文对比损失)。
import numpy as np
from PIL import Image
import torch
from transformers import BridgeTowerForContrastiveLearning, BridgeTowerProcessor
# 生成一张占位测试图(实际使用时替换为真实图片路径即可)
image = Image.fromarray(np.random.randint(0, 255, (400, 600, 3), dtype=np.uint8))
texts = ["An image of two cats chilling on a couch", "A football player scoring a goal"]
processor = BridgeTowerProcessor.from_pretrained("BridgeTower/bridgetower-large-itm-mlm-itc")
model = BridgeTowerForContrastiveLearning.from_pretrained("BridgeTower/bridgetower-large-itm-mlm-itc", device_map="auto")
for text in texts:
encoding = processor(image, text, return_tensors="pt").to(model.device)
outputs = model(**encoding)
print(text, outputs.keys())
5.2 图文检索(BridgeTowerForImageAndTextRetrieval)
BridgeTowerForImageAndTextRetrieval 本质是一个二分类打分头(输出 SequenceClassifierOutput),logits[0, 1] 越高代表图文越匹配。
import numpy as np
from PIL import Image
from transformers import BridgeTowerForImageAndTextRetrieval, BridgeTowerProcessor
image = Image.fromarray(np.random.randint(0, 255, (400, 600, 3), dtype=np.uint8))
texts = ["An image of two cats chilling on a couch", "A football player scoring a goal"]
processor = BridgeTowerProcessor.from_pretrained("BridgeTower/bridgetower-base-itm-mlm")
model = BridgeTowerForImageAndTextRetrieval.from_pretrained("BridgeTower/bridgetower-base-itm-mlm", device_map="auto")
scores = dict()
for text in texts:
encoding = processor(image, text, return_tensors="pt").to(model.device)
outputs = model(**encoding)
scores[text] = outputs.logits[0, 1].item()
print(scores) # 取 logits 索引 1 作为“图文匹配”分数
5.3 掩码语言建模(BridgeTowerForMaskedLM)
BridgeTowerForMaskedLM 在文本塔词表(50265)上做 MLM 预测,可借助 processor.decode 把 token id 还原成字符串。官方文档示例输出为 .a cat looking out of the window.(注意 decode 会保留 <mask> 位置被预测词补全后的原始空格与标点)。
import numpy as np
from PIL import Image
from transformers import BridgeTowerProcessor, BridgeTowerForMaskedLM
image = Image.fromarray(np.random.randint(0, 255, (400, 600, 3), dtype=np.uint8))
text = "a <mask> looking out of the window"
processor = BridgeTowerProcessor.from_pretrained("BridgeTower/bridgetower-base-itm-mlm")
model = BridgeTowerForMaskedLM.from_pretrained("BridgeTower/bridgetower-base-itm-mlm", device_map="auto")
encoding = processor(image, text, return_tensors="pt").to(model.device)
outputs = model(**encoding)
results = processor.decode(outputs.logits.argmax(dim=-1).squeeze(0).tolist())
print(results)
六、可用预训练检查点与使用提示
官方在文档中给出的 Tips 值得注意:
- 实现组合是固定的:本实现用 [
RobertaTokenizer] 生成文本嵌入,用 OpenAI CLIP/ViT 模型计算视觉嵌入(视觉侧权重在BridgeTowerVisionModel内加载),因此文本塔词表、BPE 切分均与 RoBERTa 保持一致; - 已发布检查点包括:
BridgeTower/bridgetower-base:基础预训练权重;BridgeTower/bridgetower-base-itm-mlm:在基础版之上用掩码语言建模(MLM)+ 图文匹配(ITM)继续训练的权重,适用于图文检索与 MLM;- 对比学习示例中使用的
BridgeTower/bridgetower-large-itm-mlm-itc:进一步加入图文对比(ITC)目标的 large 版。
- 性能参考:Image Retrieval 等下游任务上的详细对比请参阅论文正文(文档指向原论文 Table 5)。
从架构的通用性角度看,文档强调:原则上桥接机制可以套用任意视觉/文本/跨模态编码器,当前组合只是默认实例。
七、API 一览与延伸阅读路径
以下类均已随 transformers 顶层导出,使用 from transformers import ... 即可导入:
- 配置:
BridgeTowerConfig、BridgeTowerTextConfig、BridgeTowerVisionConfig - 预处理:
BridgeTowerImageProcessor(preprocess)、BridgeTowerImageProcessorPil(preprocess)、BridgeTowerProcessor(__call__) - 模型:
BridgeTowerModel(forward)、BridgeTowerForContrastiveLearning(forward)、BridgeTowerForMaskedLM(forward)、BridgeTowerForImageAndTextRetrieval(forward)
想继续深入本仓库,建议按以下路径阅读:
- 架构与前向逻辑:modeling_bridgetower.py(重点看
BridgeTowerLinkTower、BridgeTowerVisionTransformer、BridgeTowerModel.forward的桥接循环); - 配置默认值:configuration_bridgetower.py;
- 图像预处理:image_processing_bridgetower.py 与 image_processing_pil_bridgetower.py;
- Processor 组装:processing_bridgetower.py;
- 模型行为验证与数值断言:tests/models/bridgetower/test_modeling_bridgetower.py。
八、适用边界与注意事项
- 任务范围:当前文档与代码主要面向预训练/推理阶段的图文理解型任务(对比学习、图文匹配、图文检索、图文 MLM),不包含生成式 captioning 或检测类头;
BridgeTowerModel.forward明确不接受inputs_embeds(会抛出NotImplementedError),必须传input_ids; - 图像规格:图像预处理以 288 短边为默认基准,并把宽高对齐到 32 的倍数;更换输入分辨率时可通过配置
image_size与 processor 的size参数配合调整; - 指标口径:正文引用的 VQAv2 78.73%、81.15% 等数字出自论文在 4M 图上的实验报告,属于论文自述结果,不代表本仓库在当前硬件/数据下的复现结论;如需自行验证,请以本仓库测试与本地跑分为准。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0629
MiniCPM5-2BMiniCPM5-2B 是一款面向端侧、本地部署和资源受限场景的 2B 稠密 Transformer,能够达到同尺寸开源模型 SOTA 水平。Markdown00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
HivisionIDPhotos⚡️HivisionIDPhotos: a lightweight and efficient AI ID photos tools. 一个轻量级的AI证件照制作算法。Python07
DragonOSDragonOS is an operating system developed from scratch using Rust, with Linux compatibility. It is designed for **Serverless** scenarios. 使用Rust从0自研内核,具有Linux兼容性的操作系统,面向云计算Serverless场景而设计。Rust00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00