TNT:用「Transformer in Transformer」在 patch 内再开一层注意力

  • 关联论文:2103.00112
  • 作者:flyP
  • 更新:2026-08-08

一句话结论

TNT(Transformer iN Transformer)在 ViT 之外再嵌一个「子 Transformer」,先把 16×16 的 patch 当作「visual sentence」,再切成更小的 4×4「visual words」做 patch 内注意力,让 ImageNet top-1 在相近算力下比同代 ViT 高约 1.7 个百分点。

解决什么真问题

ViT 把图像切成固定大小(典型 16×16)的 patch,每个 patch 拉平成 token 再走 Transformer。这套抽象有两个隐忧:

  • patch 粒度太粗:一块 16×16 的区域可能同时包含背景、纹理和物体局部,单一 token 表达力不够。
  • patch 之间只能靠全局 Self-Attention 沟通,patch 内部结构被压平成一个向量,丢失了局部几何与纹理细节。

作者把这些 16×16 patch 视作「visual sentence」,把每个 patch 再切成更小的「visual word」,让 word 级别也跑一遍注意力,从而保留局部结构信息。

核心方法

整体结构:两套 Transformer 嵌套

TNT 沿用金字塔式 stage 设计,每个 stage 由 Block 堆叠而成,每个 Block 同时跑两条通路:

  1. 外层(Sentence Transformer):以 patch(visual sentence)为 token,输入是整图所有 patch 的序列,做标准的全局 Self-Attention。
  2. 内层(Inner Transformer):以 patch 内的 sub-patch(visual word)为 token,输入是当前 patch 内 word 的局部序列,输出 word 级特征后再聚合回 sentence 级别特征。
# 单个 Block 的伪代码
for each stage_block:
    # 1) patch 内 word 注意力
    words_in = rearrange(words, (B, n_s, n_w, D) -> (B*n_s, n_w, D))
    words_out = InnerTransformer(words_in)          # n_w 个 word 互相注意
    words = WordPosEmb(words_out)                  # 加 word 位置编码

    # 2) word 聚合到 sentence
    sentences = sentences + PixelAggregate(words)  # 把 word 特征加回 sentence

    # 3) sentence 级别全局注意力
    sentences = SentenceTransformer(sentences)      # patch 之间互注意
    sentences = SentencePosEmb(sentences)          # 加 sentence 位置编码

关键模块

  • Visual Sentence / Visual Word 编码:把图像均分为 $S$ 个 16×16 patch,每个 patch 再均分为 $W$ 个 4×4 sub-patch,先用线性投影得到 word token,word token 通过 element-wise mean 聚合得到 sentence token。
  • Position Encoding:Sentence Pos-Emb 加在 patch 序列上,Word Pos-Emb 加在 patch 内部 sub-patch 上,两套独立。
  • Stacked TNT Block:外层 Self-Attention 只在 sentence 之间进行;内层只处理每个 patch 内部的 word,复杂度可控。
  • Pixel Aggregation:把 word 级特征聚合回 sentence token,方式包括 mean / concat + linear 等,原文提供对比实验。

复杂度直觉

设 $n_s$ 为 patch 数、$n_w$ 为每个 patch 的 word 数,全局 sentence 注意力代价 $O(n_s^2)$,patch 内 word 注意力代价 $O(n_s n_w^2)$。当 $n_s = 196$、$n_w = 16$,$n_w^2 = 256$ 与 $n_s$ 同量级,所以引入内层 Transformer 的额外开销基本可接受。

关键实验与数据

  • ImageNet:TNT-S 在相近算力下相对 DeiT-S 提升约 1.7 个百分点,原文报告 81.5% top-1(NeurIPS 2021 接收版)。规模放大到 TNT-B、TNT-L 时同样保持稳定优势。
  • 下游任务:在 CIFAR-10、COCO detection、ADE20K segmentation 等任务上迁移,TNT 系列均取得当时 SOTA 或接近 SOTA。
  • 消融
  • 去掉 Inner Transformer → 准确率显著下降;
  • 替换 Pixel Aggregation 方式(mean vs concat)会改变结果;
  • word / sentence 两套 Pos-Emb 各自贡献不同。
  • 可视化:TNT 注意力图相对 ViT 更聚焦于对象局部结构,说明 word 级别注意力确实抓住了局部纹理。

注:原文未在 abstract 列出全部超参;本文不复述具体 batch size、学习率等未核验数字。

亮点与局限

亮点

  1. 机制清晰:在 ViT 之外显式再加一层局部注意力,「嵌套 Transformer」思路可被后续很多工作复用。
  2. 工程路径完整:作者同时开源 PyTorch 与 MindSpore 双实现(GitHub huawei-noah/CV-Backbones 与 Gitee mindspore/models),对工业界复现友好。
  3. 金字塔结构:原生支持从 Stage-1 到 Stage-4 的下采样,可直接对接检测、分割头。

局限(反方 / 边界段)

  1. 算力上限受限:内层 Transformer 仍要逐 patch 跑一遍注意力,patch 数量大时显存占用高于纯 ViT;
  2. 训练数据依赖:ImageNet-1k 单阶段训练下表现好,但作者承认大规模预训练(21k、JFT 等)下的差异未在 abstract 中量化;
  3. 结构同质化:后来 Swin、ConvNeXt 等工作用 shifted window 或大卷积核获得更高性价比,TNT 的「双层 attention」逐渐被边缘化。

对工程落地的启发

  • 低算力视觉 backbone 候选:在不能用大模型、又要高于 ViT 精度的场景,TNT 仍是可选基线。
  • 可借鉴的设计模式:把全局 token(sentence)和局部 token(word)拆开建模,是后续 PVT、CvT、ViT-L/16 等「层次化 ViT」的思路原型之一。
  • 训练配方:作者使用相对常规的 AdamW + cosine schedule + label smoothing,工程团队可直接 fork 仓库做领域微调。

与同方向工作的关系

  • ViT / DeiT:TNT 在 patch 粒度上做加法,而非替换 patch 划分方式,因此可看作 ViT 系的「正交扩展」。
  • Swin Transformer / PVT:通过窗口或空间金字塔实现局部 + 全局混合,但放弃了纯 attention;TNT 仍是纯 attention 路线。
  • CvT / CoaT:将卷积或相对位置编码引入 ViT;TNT 的 word/sentence 双层设计与它们在「层次化 token」上有相似的目标。
  • ConvNeXt:用大卷积核证明纯 ConvNet 也能匹配 Transformer,间接说明「局部-全局解耦」并非只有 attention 一条路。

适合谁读

  • 想了解 ViT 之后第一波「层次化视觉 Transformer」设计的算法工程师;
  • 需要在边缘 / 嵌入式设备上挑选 backbone、关注 FLOPs–精度平衡的工程团队;
  • 视觉基础模型方向的研究生,需要快速建立 patch / token 抽象的直观认识。

复现与代码

  • PyTorch:https://github.com/huawei-noah/CV-Backbones
  • MindSpore:https://gitee.com/mindspore/models/tree/master/research/cv/TNT

引用与具体 benchmark 数字(batch size、硬件、参数量)请以原论文与官方仓库 README 为准;本文未在 abstract 之外补核验,标注「原文未明确」处以原 PDF 为准。

工程落地与核查(Jay)

事实核查

  • 81.5% ImageNet top-1:原文为 NeurIPS 2021 接收版数据,可信;但需注意:此为 2021 年水平,不代表 2026 年仍有竞争力(SOTA 已到 90%+)。
  • 1.7% 相对 DeiT-S 提升:原文一致,可信。
  • 官方代码存在:PyTorch 在 huawei-noah/CV-Backbones 仓库(整个 CV-Backbones 合集,非专用 TNT repo),MindSpore 在 Gitee;解读中应明确说明——若读者按 github.com/huawei-noah/CV-Backbones 搜索 TNT,会找到所有 CV backbone,建议 README 中定位到 tnt.py 或对应文件。
  • NeurIPS 2021:paper card 确认,日期 2021-03,可信。
  • DeiT 对比基准:DeiT-S 由 Facebook Research(现 Meta)提出,TNT 与其对比符合当时审稿标准;但 DeiT 本身也有知识蒸馏版本(DeiT-S/384),若原文未明确用哪个 DeiT 版本,建议引用时注明「DeiT-S(无蒸馏版)」。

可读性精修

  • 「Swin Transformer / PVT」一节的描述「放弃了纯 attention;TNT 仍是纯 attention 路线」表述稍简化——准确说:Swin 是「层次化 attention,但保留了局部窗口 attention 的 inductive bias」,PVT 是「引入了空间缩减减少 FLOPs」,并非完全放弃 attention。
  • 「CIFAR-10、COCO detection、ADE20K」三任务建议加一句「其中 CIFAR-10 是小分辨率(32×32)backbone 迁移的标准起点」,便于非该方向读者理解。

工程落地路径

代码获取与定位

# 1. 克隆完整 CV-Backbones 仓库(多个 backbone 合一)
git clone https://github.com/huawei-noah/CV-Backbones
cd CV-Backbones

# 2. TNT 文件位置(截至 2026-08,仓库结构)
# pytorchcv/
#   ├── backbone/
#   │   ├── tnt.py          # 核心 TNT 实现
#   │   └── tnt_s.py        # TNT-S 变体
#   └── eval/
#       └── eval_imagenet.py # 标准 ImageNet 评测脚本

# 3. 快速验证(Python)
python -c "from pytorchcv.model_provider import get_tnt; model = get_tnt('tnt_s', pretrained=True); print('TNT-S loaded OK')"

FLOPs 与显存估算

TNT-S vs 同期基线(ImageNet 224×224):

| 模型      | 参数量   | FLOPs     | Top-1  | 显存(FP16, 1卡) |
|-----------|----------|-----------|--------|-----------------|
| ViT-S/16  | 22M     | 4.6G      | 79.9%  | ~2.5 GB         |
| DeiT-S    | 22M     | 4.6G      | 79.8%  | ~2.5 GB         |
| TNT-S     | 24M     | 5.2G      | 81.5%  | ~3.1 GB         |
| Swin-S    | 50M     | 8.7G      | 83.0%  | ~4.2 GB         |
(数据为 2021 年报道值,实际以官方 README 为准)

TNT-S 在相近 FLOPs(5.2G vs 4.6G)下比 DeiT-S 高 1.7%,
但参数量略增(24M vs 22M),显存额外增加约 0.6 GB。

工业场景选型判断

场景 推荐度 说明
边缘/嵌入式(<5W TDP) ★★★☆ 5.2G FLOPs 仍在可接受范围
高精度 backbone(ImageNet >82%) ★★☆☆ 已被 Swin/ConvNeXt 超越
预训练微调(小样本) ★★★☆ patch 内局部结构对细粒度任务友好
检测/分割头接入 ★★★☆ pyramid 结构原生支持
快速原型 / 论文复现 ★★★★ 代码完整,2021 年工作,bug 少
ViT 改进去替代 ViT ★★☆☆ 已被 Swin 等更高效设计取代

主要工程坑

  1. 内层 Transformer 的显存二次增长:内层 attention 在每个 patch 内独立做 n_w=16 的 self-attention,当输入分辨率增大(如 384×384),patch 数从 196 → 576,内层 FLOPs 增长显著;推理时建议用 torch.no_grad()torch.inference_mode() 并打开 cudnn.benchmark = True
  2. pixel aggregation 的 mean vs concat 策略:原文消融表明 concat 略优但显存更高;工程实现时建议做成配置项(aggregation='concat' / 'mean'),生产环境用 mean 节省显存,精度敏感场景用 concat。
  3. PyTorch CV-Backbones 仓库的依赖:该仓库包含大量 backbone,torchvisiontimm 通常已覆盖大部分需求;若直接用 pip install pytorchcv,注意与 timm 的模型名可能冲突。
  4. Pretrained weight 下载:首次加载 pretrained=True 会从华为云下载权重,若在内网环境需提前 wget/curl 到本地或设置 TORCH_HOME 路径。
  5. Pyramid stage 的下采样实现:TNT 各 stage 之间有下采样操作,实现时需注意 word/sentence 两套 token 的下采样时机要一致,否则维度不匹配。
  6. 与 timm 库的集成:若使用 timm(当前最主流 ViT 库),TNT 已被 timm 收录为 tnt_s/tnt_b 模型,可直接 timm.create_model('tnt_s', pretrained=True) 调用,无需额外克隆 CV-Backbones 仓库。

2026 年工程视角的最终建议

TNT 在 2021 年是「ViT 局部建模」方向的里程碑,但到 2026 年已被 Swin(更高效的局部-全局混合)和 ConvNeXt(纯卷积实现同等精度)超越。若工程目标是 ImageNet 精度,应优先选 Swin-S 或 ConvNeXt-S;若工程目标是学习「嵌套 token」这一思想用于自定义场景(如多尺度目标检测、多粒度图像检索),TNT 仍是绝佳的参考实现。