vit-pytorch README 展示的 Vision Transformer 图像分块与序列建模动画

vit-pytorch

Stars

面向视觉 Transformer 研究的 PyTorch 实现集合,覆盖 ViT、DeiT、CaiT、CrossViT、NaViT 等大量架构与训练变体,并提供可直接组合的注意力、正则化和特征学习模块。

Vision TransformerPyTorchAttention · Decorrelation
VISION TRANSFORMER · TENSOR AXIS SEMANTICS

采样 Token,不能改写 Batch

`sample_frac` 应缩短每个样本的序列维;层和批次只负责定位独立样本,必须在采样前后保持原位。

VIT ACTIVATIONS[ L · B · N · D ]
LlayerBbatchNtokenDembedding
sample_frac = 0.5
BEFORE · WRONG AXISrand(tokens.shape[:2])
B₀B₀B₂
N stays unchanged

batch reordered / duplicated

AFTER · TOKEN GATHERargsort(..., dim=-1)[..., :k]tokens.gather(-2, indices)
[ L · B ·k· D ]
LEADING DIMS UNCHANGED

actual loss = explicit token subset

轴向错位

序列没有缩短,Batch 却被重排

`DecorrelationLoss(sample_frac < 1)` 会先把输入打包,再仅依据前两个维度生成随机索引。面对 ViT 产生的 `[layer, batch, token, dim]` 张量时,这些索引实际落在 batch 轴:序列长度完全没有减少,batch 却可能被重新排序或重复,最终损失并未计算预期的 token 子集。

张量契约

前导维保持,只有 Token 轴收缩

采样比例描述的是每个样本内部的 token 数量,因此无论前面存在多少层维或批次维,都必须保持这些 leading dimensions 原位,只沿倒数第二个 token 轴缩短 `N → k`。索引还要为每组前导坐标独立生成,并在特征维展开后交给同一条 gather 操作。

索引修复

逐组生成索引并沿 −2 维 Gather

直接从 `tokens.shape[:-1]` 生成随机分数,在最后一维排序并截取 `num_sampled` 个 token 索引;将索引扩展到 embedding 维后执行 `tokens.gather(-2, indices)`,删除 pack / unpack、batch arange 与高级索引。固定随机种子的回归将结果与显式选出的同一 token 子集逐项比较。

行为验证

随机采样与显式 Token 子集一致

`sample_frac=0.5` 现在会把每个 layer / batch 的 token 维从 4 正确缩到 2,而不会改写或复制 batch;任意前导维均保持不变,损失值与显式子集计算一致。完整 ViT 前向与反向传播也通过验证,采样路径重新符合参数语义。