尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

视频生成模型中的Token Radius Attention:局部稀疏注意力原理与PyTorch实现

视频生成模型中的Token Radius Attention:局部稀疏注意力原理与PyTorch实现 视频生成模型对算力的消耗很大一部分来自 Transformer 层处理海量 token。一个短视频在 patch 化之后会产生数万甚至数十万 token如果每一层都做全局时空注意力计算量和显存都会迅速失控。Token Radius Attention 是一种把注意力参与范围限制在 query token 附近的机制每个位置只和自己时空半径内的 token 计算注意力。它利用视频相邻帧、相邻区域强相关的先验在保持生成质量的同时减少无意义的全局交互。下面从原理讲起再用 PyTorch 写一个可运行的最小实现验证正确性、对比开销并说明真实视频生成工程中如何落地。1. 先理解视频生成里的 Attention 为什么卡在 token 规模上1.1 视频是怎么变成 token 的在视频生成任务里模型通常不会直接处理原始像素。常见流程是先用 3D VAE 把视频压缩到隐空间再把压缩后的隐特征图切分成 patch最后由线性映射变成一组 token。假设输入视频是 16 帧、分辨率 64×64patch 大小为 2×2那么每一帧会得到 32×32 个 patch。如果时间维不压缩单条视频的 token 数量是T * H * W 16 * 32 * 32 16384这里的 H 和 W 是 patch 化之后的空间网格尺寸不是原始像素分辨率。如果视频分辨率更高、帧数更多token 数量会进一步膨胀。对于视频生成模型一个 batch 里有多条视频每个 denoising step 或自回归 step 都要把这些 token 送入 Transformer 层计算量很快就压不住。这里的 token 可以简单理解成“视频块”的嵌入向量。它和 NLP 里的 token 语义类似但底层对应的是视频隐特征张量中的一个小区域。1.2 全局注意力为什么在视频场景里更贵Transformer 的核心自注意力计算需要生成一个 N×N 的注意力矩阵代表每个 query token 和每个 key token 之间的关系。N 是 token 数量。对于长度为 N 的序列标准 attention 分数矩阵大小 B * H * N * N其中 B 是 batch sizeH 是注意力头数。当 N 达到 16384即使 H 只有 8单个样本的 attention 分数矩阵大小也是8 * 16384 * 16384 2147483648也就是约 21 亿个元素。如果按 float16 计算单是保存这份分数矩阵就需要约 4GB 显存。这还没有计算多头线性变换、MLP、梯度等开销。所以在视频生成模型里直接套用图像或文本 Transformer 的全局注意力设计显存和算力会迅速成为瓶颈。计算量也同理。自注意力的矩阵乘部分大致是 O(N²) 的。N 每翻一倍注意力计算量翻四倍。视频 token 数量比图像多一个时间维度这个缺陷会被明显放大。1.3 Token Radius Attention 的出发点视频有一个很强的先验相邻帧之间的内容通常连续相邻空间区域之间的关联也远大于远处区域。一个 query token 在理解当前局部内容时不一定需要把整段视频所有 token 都看一遍。Token Radius Attention 的核心思想是给每个 token 定义一个以自己为中心的时空半径只允许 query token 和半径范围内的 key token 计算注意力。半径之外的信息被直接忽略。这个设计把注意力从“全局稠密”变成“局部稀疏”在理论上可以把每个 query 的交互范围从 N 个 token 缩减到 M 个 token其中 M 由半径决定和视频整体长度没有直接关系。2. Token Radius Attention 的设计思路和工作机制2.1 用半径定义注意力范围在视频 Token 网格中每个 token 都有三个坐标(t, h, w)分别表示时间帧索引、空间高度位置、空间宽度位置。Token Radius Attention 要求 query token 和 key token 的坐标距离小于等于某个半径 R才允许两者计算注意力。判断条件可以写成sqrt((t_q - t_k)^2 (h_q - h_k)^2 (w_q - w_k)^2) R这个公式描述的是一个三维空间中的球体邻域。每个 query token 是球心半径 R 决定了它能看到多大范围的视频内容。更通用的做法是分别设置时间半径和空间半径radius_t控制跨帧范围。radius_s控制空间邻域范围。例如|t_q - t_k| radius_t max(|h_q - h_k|, |w_q - w_k|) radius_s这种设置在实际代码里更灵活因为视频生成中时间维和空间维的分辨率差异往往很大。2.2 半径不是越大越好半径越大每个 query token 能参与的 key token 越多信息越丰富但计算量也越大。当半径大到可以覆盖所有 token 时Token Radius Attention 就会退化成普通全局注意力。半径越小计算越省但模型可能看不到跨区域、跨时间的上下文。实际使用时要平衡两点视频局部运动的连续性。模型接收全局信息的需要。对于形状变化剧烈的物体、镜头切换、大面积运动过小的半径会导致生成结果缺少全局一致性。2.3 和 window attention、token pruning 的关系Token Radius Attention 并不算一种完全独立的新类型。它和视觉 Transformer 中常见的 window attention 同属于“局部注意力”家族但关注点有一个细微差别window attention 通常固定一个方形窗口例如 7×7 或 8×8。Token Radius Attention 强调以 token 坐标为中心的“半径邻域”窗口形状可以是球形也可以是矩形甚至可以在不同层使用不同半径。它还经常和 token pruning、token merging 配合使用。先通过半径限制注意力范围再在半径范围内评估 token 的重要性把不重要的 token 合并或丢弃从而进一步减少计算量。下面用一个简化表格对比三种方案方案交互范围典型复杂度优势风险全局 attention所有 tokenO(N²)信息完整视频长序列显存压力大window attention固定窗口内O(N × W)规律清晰硬件友好窗口边缘信息割裂Token Radius Attention半径范围内O(N × M)利用时空先验范围可调半径设置不当会丢失重要信息这里的 M 是半径范围内的平均 token 数量。实际使用中M 通常远小于 N。3. 环境准备和最小实验设计3.1 实验环境依赖为了跑通下面的最小实现需要一个带 PyTorch 的 Python 环境。推荐使用 Python 3.10 或以上版本并安装 PyTorch 和 einops。创建虚拟环境并安装依赖python -m venv .venv source .venv/bin/activate pip install torch2.0 einops如果使用 GPU需要根据 CUDA 版本安装对应 PyTorch 预编译包。下面示例代码在 CPU 上也可以运行只是速度慢一些。3.2 生成最小视频数据因为重点是理解注意力机制不涉及真实视频解码这里用随机张量模拟一段视频输入。import torch B 1 # batch size C 3 # 视频通道数RGB T 4 # 帧数 H 8 # 原始视频高度 W 8 # 原始视频宽度 x torch.randn(B, C, T, H, W) print(x.shape)输出结果为torch.Size([1, 3, 4, 8, 8])这个张量可以理解成一段 4 帧、每帧 8×8 的随机视频。3.3 项目文件结构实验代码放在一个普通 Python 项目中即可建议按下面结构组织token_radius_attention/ ├── radius_attention.py ├── train_test.py └── requirements.txt这里radius_attention.py放核心模块代码train_test.py放验证脚本。4. 用 PyTorch 实现 Token Radius Attention4.1 先将视频 patch 化并生成 token模拟实现里只证明机制不追求完整视频生成模型。先用一个 3D 卷积把视频切成 patch 并映射成 token。import torch import torch.nn as nn from einops import rearrange class PatchEmbed3D(nn.Module): def __init__(self, in_channels3, embed_dim64, patch_size(1, 2, 2)): super().__init__() self.proj nn.Conv3d( in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size ) def forward(self, x): # x: (B, C, T, H, W) x self.proj(x) # x: (B, embed_dim, T, H, W) x rearrange(x, b d t h w - b (t h w) d) return x这里 patch_size 的时间维度是 1表示不压缩帧数只压缩空间。示例里输入是 8×8patch 是 2×2输出空间网格是 4×4。调用方式patch_embed PatchEmbed3D(in_channels3, embed_dim64, patch_size(1, 2, 2)) tokens patch_embed(x) print(tokens.shape)输出结果为torch.Size([1, 64, 64])这里的 token 序列长度是T * H * W 4 * 4 * 4 64每个 token 的维度是 64。4.2 生成 token 坐标网格要为每个 token 计算半径邻域必须先知道每个 token 对应的(t, h, w)坐标。import torch def make_grid_positions(T, H, W): ts torch.arange(T) hs torch.arange(H) ws torch.arange(W) grid torch.stack(torch.meshgrid(ts, hs, ws, indexingij), dim-1) return grid.reshape(-1, 3).float()调用pos make_grid_positions(T4, H4, W4) print(pos.shape)输出结果为torch.Size([64, 3])这个坐标矩阵的顺序必须和 token 展开顺序一致。这里 token 的顺序是“时间优先然后高度然后宽度”即第一个 token 是第 0 帧第 0 行第 0 列。4.3 构造 radius mask最直观的 mask 构造方式是计算所有 token 两两之间的距离然后判断距离是否小于等于半径。def make_radius_mask(T, H, W, radius): pos make_grid_positions(T, H, W) dist torch.cdist(pos, pos) return dist radiustorch.cdist计算的是三维坐标之间的欧氏距离。返回的 mask 形状是(N, N)其中mask[i, j] True表示 token j 在 token i 的半径范围内。这个版本非常直观但存在一个明显问题它本身仍然计算了 N×N 的距离矩阵显存和计算量都是 O(N²)。它适合用来验证正确性不适合直接用在真实的大视频生成模型里。真实工程中需要用局部窗口展开或稀疏注意力实现来避免生成完整的 N×N mask。4.4 完整的 Token Radius Attention 模块下面是核心模块的完整实现。import torch import torch.nn as nn from einops import rearrange class TokenRadiusAttention(nn.Module): def __init__( self, dim, num_heads, radius1.5, max_temporal_radiusNone, attn_drop0.0 ): super().__init__() assert dim % num_heads 0 self.dim dim self.num_heads num_heads self.head_dim dim // num_heads self.scale self.head_dim ** -0.5 self.radius radius self.max_temporal_radius max_temporal_radius self.qkv nn.Linear(dim, dim * 3, biasFalse) self.attn_drop nn.Dropout(attn_drop) self.proj nn.Linear(dim, dim) def forward(self, x, grid_positions): B, N, D x.shape qkv self.qkv(x) q, k, v rearrange( qkv, b n (three h d) - three b h n d, three3, hself.num_heads ) q q * self.scale attn torch.einsum(b h q d, b h k d - b h q k, q, k) dist torch.cdist(grid_positions, grid_positions) mask dist self.radius if self.max_temporal_radius is not None: dt grid_positions[:, :1] - grid_positions[:, :1].T mask dt.abs() self.max_temporal_radius attn attn.masked_fill(~mask[None, None], float(-inf)) attn attn.softmax(dim-1) attn self.attn_drop(attn) out torch.einsum(b h q k, b h k d - b h q d, attn, v) out rearrange(out, b h n d - b n (h d)) return self.proj(out)这里有几个关键点需要解释dist self.radius是半径 mask 的核心。只有半径内的 key 会参与 softmax。softmax 之前要把不允许的位置设为-inf不能设为 0。如果设为 0softmax 后这些位置仍然会有很小的非零权重。max_temporal_radius是对时间维度的额外限制。如果视频中跨帧注意力范围太远可以用这个参数单独收紧。多头计算通过rearrange一次完成避免手动循环。4.5 参数说明参数含义示例值调大影响调小影响dimtoken 向量维度64表达能力强计算量增大表达能力下降num_heads注意力头数4多头多样性更好显存增加多样性减弱radius时空半径1.5每个 query 看到更多 token计算量增大更省计算但可能丢失上下文max_temporal_radius时间维度最大跨帧距离1允许更长跨帧依赖当前帧只看到近邻帧实际项目中radius 通常需要根据 patch 大小和输入分辨率调整。patch 是 16×16 时半径 1 和 patch 是 2×2 时半径 1 代表的实际空间范围完全不同。5. 跑通验证正确性、效率和可视化5.1 当半径覆盖所有 token 时应该等价于全局注意力这是一个重要的正确性测试。如果一个 radius 大得足以覆盖所有 tokenToken Radius Attention 等价于普通全局注意力。用同一个权重初始化两个模块分别用大半径和全局注意力做对比。import torch torch.manual_seed(0) B, C, T, H, W 1, 3, 4, 8, 8 x torch.randn(B, C, T, H, W) patch_embed PatchEmbed3D(in_channels3, embed_dim64, patch_size(1, 2, 2)) tokens patch_embed(x) pos make_grid_positions(T4, H4, W4) global_attn TokenRadiusAttention(dim64, num_heads4, radius999.0) radius_attn TokenRadiusAttention(dim64, num_heads4, radius1.5) radius_attn.load_state_dict(global_attn.state_dict()) with torch.no_grad(): y_global global_attn(tokens, pos) y_radius radius_attn(tokens, pos) print(torch.allclose(y_global, y_radius, atol1e-5))因为两个模块使用同一套权重大半径模块的 mask 全为 True最终结果应该和全局注意力一致。如果输出是True说明实现里的 mask 和 softmax 逻辑没有错误。但这里要注意只有当半径覆盖所有 token 时才等价。中等半径时输出原本就会不同因为每个 token 看到的信息范围不同。5.2 计算量对比标准全局注意力的核心计算量大约是global: O(B * N * N * D)Token Radius Attention 如果实现为真正的局部稀疏注意力每个 query 只和半径内 M 个 token 交互复杂度大约是radius: O(B * N * M * D)当 M 明显小于 N 时注意力部分的理论计算量会下降。注意上面实现中的 mask 版本没有这个优势因为它仍然先计算了完整的 N×N 分数矩阵。mask 版本用于验证正确性真实生产环境需要换成局部窗口展开或稀疏注意力实现。对比表格数据规模N全局注意力 score 元素数半径注意力 score 元素数4 帧 4×4644096约 600120016 帧 16×1640961677 万几十万到百万级别32 帧 32×323276810.7 亿百万到千万级别这里的“约”字是因为不同半径和边界条件会导致 M 不同。边界 token 的邻域通常更小。5.3 输出形状和检查清单运行完整实验脚本后核心输出应该满足输入视频: torch.Size([1, 3, 4, 8, 8]) token 序列: torch.Size([1, 64, 64]) 坐标网格: torch.Size([64, 3]) attention 输出: torch.Size([1, 64, 64])验证时可以按这个清单检查patch 化后的 token 长度是否等于T * H * W。坐标网格长度是否和 token 长度一致。mask 中每个 query 至少能看到自己。大半径输出是否和全局注意力一致。输出张量形状是否与输入 token 形状一致。注意不要只验证程序能启动还要验证大半径退化结果、边界 token 的 mask、全 mask 行等边界情况。6. 常见问题与排查路径6.1 mask 构造错了注意力结果显示错乱现象输出结果和大半径版本差别很大但看不出明显报错。可能原因坐标顺序和 token 顺序不一致。比如 token 是按t, h, w展开但坐标网格是按h, w, t构造的。排查方式打印前几个 token 的坐标和tokens.shape的展开顺序对照。解决方式统一t, h, w顺序并在构造坐标后打印 shape 和顺序做断言。6.2 softmax 之后出现 NaN现象训练或测试一段时间后 loss 变成 NaN。可能原因某个 query token 的 mask 全是Falsesoftmax 对全-inf行计算时得到0/0产生 NaN。排查方式对 mask 的每一行求和检查是否所有行都至少有一个True。解决方式在 mask 中强制保证对角线可见也就是每个 token 至少能注意自己。mask dist self.radius mask mask | torch.eye(mask.size(0), dtypetorch.bool, devicemask.device)这是因为视频边缘的 token 在小半径设置下周围可能没有足够多的邻居。保留自身注意力可以避免空行问题。6.3 看起来已经加入了 radius mask但显存还是爆了现象日志里显存仍然持续上涨或者直接 OOM。可能原因实现了 mask 版稀疏注意力但仍然创建了torch.cdist结果和完整的attn分数矩阵复杂度本质上还是 O(N²)。排查方式查看attn张量形状。如果形状是(B, num_heads, N, N)说明仍然保存了完整稠密矩阵。解决方式真实生成场景应使用局部窗口展开实现。例如将特征图重排成(B, T, H, W, D)在时间窗口内切片空间邻域避免构造完整 N×N 矩阵。6.4 半径参数设置了但效果和全局注意力几乎一样现象输出和全局注意力差异极小计算量却没有任何明显下降。可能原因radius 设置过大。比如视频网格只有 4×4×4radius 设置为 3 会覆盖绝大多数 token。排查方式统计 mask 中True比例。print(mask.float().mean())解决方式从 0.5 或 1.0 开始调小半径逐步观察生成质量和计算量。常见问题总结表问题现象可能原因检查方式处理建议输出和大半径版本不一致坐标顺序错误打印 pos 前几行统一 t/h/w 展开顺序loss 出现 NaNmask 整行全 False检查每行是否有 True强制保留自身 token显存 OOM仍构造了 N×N 矩阵检查 attn 形状换成局部窗口稀疏实现明显没有加速radius 设置过大统计 mask True 比例调小半径或分开设置时间/空间半径7. 在真实视频生成工程中的落地建议7.1 和 DiT、3D VAE 结构如何配合真实视频扩散模型通常在 3D VAE 压缩后的隐空间里做 denoising。Token Radius Attention 可以替换模型内部的部分全局长距离注意力层但不需要替换所有层。常见做法是前几层使用相对较大的半径因为浅层特征需要较广的空间上下文。深层使用较小半径因为深层特征已经具备较强的语义局部性。每隔若干层保留一个全局 attention 层用来建立长距离、跨区域的依赖。这种“局部为主、全局为辅”的设计比把所有层全部改成 Radius Attention 更容易保持生成质量。7.2 学习环境和生产环境要区分对待如果只是学习机制上面基于torch.cdist的 mask 实现完全够用。它可以帮你看清 radius 是怎么影响注意力权重和输出的。如果是在真实训练或推理环境使用需要注意mask 版实现不能直接用于长视频必须先改成窗口展开或稀疏注意力。训练时可以用较大的 batch但要注意 mask 的稀疏度会随 batch 变化。推理时如果做自回归生成缓存机制需要和“半径内 token”对齐。由于每个 query 只依赖半径内历史 tokenKV cache 可以按时间窗口丢弃旧 token减少内存占用。生产环境还要考虑 kernel 是否支持稀疏 mask。部分 GPU kernel 在遇到不规则 mask 时效率反而会下降。推荐做法先在普通全局注意力上训练一个小模型作为 baseline再替换成 Token Radius Attention观察生成质量和显存变化。不要一开始就在大视频上做对比。7.3 和 token 剪枝、KV cache 结合Token Radius Attention 本身只解决“看多大范围”的问题不解决“看哪些 token 更值得看”的问题。它可以和以下策略配合token pruning在半径范围内先根据空间位置和内容相关性筛选重要 token减少参与注意力的 token 数。token merging把相邻帧中相似度极高的 token 合并减少序列长度。KV cache对于自回归视频生成缓存半径内而非全局的 key/value可以显著降低长视频推理时的缓存压力。如果视频生成模型是扩散模型半径注意力主要作用于每个 denoising step。此时 token 数量和视频长度固定计算量下降直接反映在单 step 的耗时上。7.4 落地检查清单在把 Token Radius Attention 用到真实项目前可以按下面的清单过一遍[ ] patch 化后的T, H, W是否明确。[ ] token 坐标网格是否与 token 展开顺序一致。[ ] 每个 query 至少保留自身 token避免 NaN。[ ] 时间和空间半径分别设置还是使用统一半径。[ ] mask 版实现是否只用于验证生产环境是否替换为窗口展开。[ ] 是否保留了若干层全局注意力用于长距离依赖。[ ] 是否统计了 mask 中的平均可见 token 数量。[ ] 是否用大半径退化为全局注意力验证过实现正确性。[ ] 是否用同一个权重对比过不同半径的生成结果。[ ] 推理阶段的 KV cache 是否和半径范围匹配。[ ] 是否记录了显存、耗时、生成效果三个指标。最后回到本文最核心的判断Token Radius Attention 的价值不在于提出一个复杂的数学公式而在于把“视频相邻内容强相关”这个先验显式写进注意力范围里。它减少的是无意义的跨区域、跨时间注意力计算。如果你正在深入视频生成模型可以先在小型数据集上做半径实验把大半径退化为全局注意力作为正确性基准再逐步缩小半径找到效果和效率的平衡点。对于新手来说最有价值的练习不是直接复现大规模模型而是先把本文的最小模块跑通然后手工统计不同半径下的注意力可见范围观察边界 token 和中心 token 在行为上的差异。
返回列表