首页
/ BridgeTower 完整指南:借助桥接层打通视觉与文本双塔编码器(Transformers 实现解读)

BridgeTower 完整指南:借助桥接层打通视觉与文本双塔编码器(Transformers 实现解读)

2026-09-07 11:38:53作者:董灵辛Dennis

本指南以 docs/source/en/model_doc/bridgetower.md 为核心主体,结合 modeling_bridgetower.pyconfiguration_bridgetower.pyprocessing_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(双塔)架构,并在其上演化出两类设计:

  1. 轻量单模态编码器 + 深层跨模态编码器:编码器本身很浅,对齐与融合完全依赖深层跨模态编码器完成,受限于单模态特征的表达能力;
  2. 预训练好的深层单模态编码器 + 顶层跨模态编码器:只用单模态编码器的最后一层表示喂给顶层的跨模态编码器,中间层的丰富语义被丢弃。

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.pyBridgeTowerModel.__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):

  1. 前 7 层:文本与图像各自独立走单模态编码器(layer[:split_index]),不做任何跨模态交互;
  2. 第一层跨模态层特例:由于此时还没有任何跨模态输出可供桥接,第一层直接对单模态输出做线性变换(cross_modal_text/image_transform)、加上可学习的 token type embedding、过 LayerNorm 后进入跨模态层;
  3. 后 5 层桥接循环:从 i = split_index 开始,每推进一层单模态编码器,就调用一次 text_link_tower / image_link_tower,把"该层单模态输出"与"上一层跨模态输出"融合,再送入下一层跨模态编码器,如此逐层把跨模态编码器与双塔顶层"缝"在一起。

最后 get_cls_featuresmodeling_bridgetower.py)分别用文本、图像两个 pooler 取出 [CLS] 表征并经 Tanh 激活后 在最后一维拼接(维度为 2 * hidden_size),作为联合表示。

2.3 输出数据结构

BridgeTowerModel.forward 返回 BridgeTowerModelOutputmodeling_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_statesattentions 供分析使用。

三、配置体系:三个 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 = 50265hidden_size = 768num_hidden_layers = 12num_attention_heads = 12intermediate_size = 3072hidden_act = "gelu"max_position_embeddings = 514pad/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_configvision_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_idsattention_mask
  • BridgeTowerImageProcessor——负责把图像转为 pixel_values(以及 pixel_mask)。

BridgeTowerProcessor 还带有一套默认 kwargs(源码 BridgeTowerProcessorKwargs):文本侧默认 add_special_tokens=Truepadding=Falsestride=0return_overflowing_tokens=Falsereturn_special_tokens_mask=Falsereturn_offsets_mapping=Falsereturn_length=Falseverbose=True;图像侧默认 do_normalize=Truedo_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_MEANimage_std = OPENAI_CLIP_STD复用 OpenAI CLIP 的归一化统计量,与 CLIP/ViT 视觉塔的设计保持一致;
  • size_divisor = 32:resize 时把宽高规整为 32 的整数倍(get_resize_output_image_sizenew_height // size_divisor * size_divisor,见 image_processing_bridgetower.py),保证后续 patch 化与注意力掩码对齐;
  • model_input_names = ["pixel_values", "pixel_mask"]

五、实战:三类下游任务的完整代码

本实现支持三类常见视觉-语言任务,每个任务对应一个头部模型。下面示例均采用本地生成图像的方式,避免依赖外部图片地址,可在离线环境直接运行(语义与官方文档示例保持一致)。若需真实图片,可用任意本地图文件替代 Image.fromarray 生成的张量。

5.1 图文对比学习(BridgeTowerForContrastiveLearning)

BridgeTowerForContrastiveLearning 返回 BridgeTowerContrastiveOutputmodeling_bridgetower.py),其中 text_embedsimage_embedscross_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 值得注意:

  1. 实现组合是固定的:本实现用 [RobertaTokenizer] 生成文本嵌入,用 OpenAI CLIP/ViT 模型计算视觉嵌入(视觉侧权重在 BridgeTowerVisionModel 内加载),因此文本塔词表、BPE 切分均与 RoBERTa 保持一致;
  2. 已发布检查点包括:
    • BridgeTower/bridgetower-base:基础预训练权重;
    • BridgeTower/bridgetower-base-itm-mlm:在基础版之上用掩码语言建模(MLM)+ 图文匹配(ITM)继续训练的权重,适用于图文检索与 MLM;
    • 对比学习示例中使用的 BridgeTower/bridgetower-large-itm-mlm-itc:进一步加入图文对比(ITC)目标的 large 版。
  3. 性能参考:Image Retrieval 等下游任务上的详细对比请参阅论文正文(文档指向原论文 Table 5)。

从架构的通用性角度看,文档强调:原则上桥接机制可以套用任意视觉/文本/跨模态编码器,当前组合只是默认实例。

七、API 一览与延伸阅读路径

以下类均已随 transformers 顶层导出,使用 from transformers import ... 即可导入:

  • 配置:BridgeTowerConfigBridgeTowerTextConfigBridgeTowerVisionConfig
  • 预处理:BridgeTowerImageProcessorpreprocess)、BridgeTowerImageProcessorPilpreprocess)、BridgeTowerProcessor__call__
  • 模型:BridgeTowerModelforward)、BridgeTowerForContrastiveLearningforward)、BridgeTowerForMaskedLMforward)、BridgeTowerForImageAndTextRetrievalforward

想继续深入本仓库,建议按以下路径阅读:

八、适用边界与注意事项

  • 任务范围:当前文档与代码主要面向预训练/推理阶段的图文理解型任务(对比学习、图文匹配、图文检索、图文 MLM),不包含生成式 captioning 或检测类头;BridgeTowerModel.forward 明确不接受 inputs_embeds(会抛出 NotImplementedError),必须传 input_ids
  • 图像规格:图像预处理以 288 短边为默认基准,并把宽高对齐到 32 的倍数;更换输入分辨率时可通过配置 image_size 与 processor 的 size 参数配合调整;
  • 指标口径:正文引用的 VQAv2 78.73%、81.15% 等数字出自论文在 4M 图上的实验报告,属于论文自述结果,不代表本仓库在当前硬件/数据下的复现结论;如需自行验证,请以本仓库测试与本地跑分为准。
登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.74 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
857
1.35 K
docsdocs
暂无描述
Markdown
897
5.81 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
531
595
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
920
1.84 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.63 K
1.02 K
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.36 K
1.46 K
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.02 K
518
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
547
389