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


面向视觉 Transformer 研究的 PyTorch 实现集合,覆盖 ViT、DeiT、CaiT、CrossViT、NaViT 等大量架构与训练变体,并提供可直接组合的注意力、正则化和特征学习模块。
`sample_frac` 应缩短每个样本的序列维;层和批次只负责定位独立样本,必须在采样前后保持原位。
sample_frac = 0.5rand(tokens.shape[:2])batch reordered / duplicated
argsort(..., dim=-1)[..., :k]tokens.gather(-2, indices)actual loss = explicit token subset
`DecorrelationLoss(sample_frac < 1)` 会先把输入打包,再仅依据前两个维度生成随机索引。面对 ViT 产生的 `[layer, batch, token, dim]` 张量时,这些索引实际落在 batch 轴:序列长度完全没有减少,batch 却可能被重新排序或重复,最终损失并未计算预期的 token 子集。
采样比例描述的是每个样本内部的 token 数量,因此无论前面存在多少层维或批次维,都必须保持这些 leading dimensions 原位,只沿倒数第二个 token 轴缩短 `N → k`。索引还要为每组前导坐标独立生成,并在特征维展开后交给同一条 gather 操作。
直接从 `tokens.shape[:-1]` 生成随机分数,在最后一维排序并截取 `num_sampled` 个 token 索引;将索引扩展到 embedding 维后执行 `tokens.gather(-2, indices)`,删除 pack / unpack、batch arange 与高级索引。固定随机种子的回归将结果与显式选出的同一 token 子集逐项比较。
`sample_frac=0.5` 现在会把每个 layer / batch 的 token 维从 4 正确缩到 2,而不会改写或复制 batch;任意前导维均保持不变,损失值与显式子集计算一致。完整 ViT 前向与反向传播也通过验证,采样路径重新符合参数语义。