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 同时跑两条通路:
- 外层(Sentence Transformer):以 patch(visual sentence)为 token,输入是整图所有 patch 的序列,做标准的全局 Self-Attention。
- 内层(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、学习率等未核验数字。
亮点与局限
亮点
- 机制清晰:在 ViT 之外显式再加一层局部注意力,「嵌套 Transformer」思路可被后续很多工作复用。
- 工程路径完整:作者同时开源 PyTorch 与 MindSpore 双实现(GitHub huawei-noah/CV-Backbones 与 Gitee mindspore/models),对工业界复现友好。
- 金字塔结构:原生支持从 Stage-1 到 Stage-4 的下采样,可直接对接检测、分割头。
局限(反方 / 边界段)
- 算力上限受限:内层 Transformer 仍要逐 patch 跑一遍注意力,patch 数量大时显存占用高于纯 ViT;
- 训练数据依赖:ImageNet-1k 单阶段训练下表现好,但作者承认大规模预训练(21k、JFT 等)下的差异未在 abstract 中量化;
- 结构同质化:后来 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 等更高效设计取代 |
主要工程坑:
- 内层 Transformer 的显存二次增长:内层 attention 在每个 patch 内独立做 n_w=16 的 self-attention,当输入分辨率增大(如 384×384),patch 数从 196 → 576,内层 FLOPs 增长显著;推理时建议用
torch.no_grad()或torch.inference_mode()并打开cudnn.benchmark = True。 - pixel aggregation 的 mean vs concat 策略:原文消融表明 concat 略优但显存更高;工程实现时建议做成配置项(
aggregation='concat'/'mean'),生产环境用 mean 节省显存,精度敏感场景用 concat。 - PyTorch CV-Backbones 仓库的依赖:该仓库包含大量 backbone,
torchvision和timm通常已覆盖大部分需求;若直接用pip install pytorchcv,注意与 timm 的模型名可能冲突。 - Pretrained weight 下载:首次加载
pretrained=True会从华为云下载权重,若在内网环境需提前 wget/curl 到本地或设置TORCH_HOME路径。 - Pyramid stage 的下采样实现:TNT 各 stage 之间有下采样操作,实现时需注意 word/sentence 两套 token 的下采样时机要一致,否则维度不匹配。
- 与 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 仍是绝佳的参考实现。