8 卡怎么分 token:LTX-2 序列并行的 all2all 到底搬了什么
8 卡怎么分 tokenLTX-2 序列并行的 all2all 到底搬了什么【免费下载链接】LTX-2Official Python inference and LoRA trainer package for the LTX-2 audio–video generative model.项目地址: https://gitcode.com/GitHub_Trending/lt/LTX-2Meta Description讲清 LTX-2 多 GPU 序列并行Sequence Parallelism机制——token 怎么均匀切到各卡、all2all 内核搬的是什么数据、以及在自己的 runner 里用 3 行代码启用 SP 的最小路径。场景钩子一条 121 帧、1024×1536 的视频latent 化后约 2.4 万 token。去噪每走一步每个 token 都要和另外 2.4 万个 token 做自注意力单卡要独自算完这个平方级的大矩阵。把模型拷 8 份、每卡各跑一遍延迟几乎不动显存还白白翻倍。问题从来不是卡够不够快而是怎么把这 2.4 万个 token 拆开让每卡只算一份、又不切断任何 token 间的交互。本文回答的正是这个问题LTX-2 多 GPU 推理里干这件事的方案叫序列并行Sequence Parallelism下称 SP。核心洞察SP 把视频 token 维均匀切到各卡再用自定义 all2all 内核把每卡本地 token × 全头换过去、全 token × 本地头收回来本地注意力的结果与单卡数值等价。像把一条长视频切 8 段、8 个人并行看最后再拼回完整画面。精确表述每个 rank 只保留 1/world_size 的 token 行注意力计算前把 Q/K/V 在token × head这个二维空间做了一次分布变换——每张卡最终持有全部 token × heads/world_size 个头本地算完再洗牌回原分布。注意力始终是全局的唯一变化是浮点归约顺序往返gather(send(x)) x逐字节精确。数据流拆解token 怎么切到各卡先补齐到卡数整数倍再等长切片pad 行在注意力里被屏蔽、在 gather 后被切掉。把 seq 维补齐到world_size的整数倍 → 保证每卡分片等长all2all 算子也能从输入 shape 符号化推出输出 shape → 代价最多 world_size-1 个 pad 行8 卡时最多 7 行。实现见 sequence_parallel.py。构造 padding mask。用户没传 mask 时构造key-only maskshape(1, 1, T_padded)有效 key 为 1、pad 为 0沿 batch 和 query 广播 → 内存 O(T)不物化(B, T, T)稠密矩阵 → 用户传了(B, T, T)mask 时则扩展 pad 行列且 pad 的 query 行允许看所有有效 key——全 mask 的行 softmax 会出 NaN。把 latent / timesteps / positions / keyframes_mask 沿序列维切到本 rank → 各张量从 T 缩到 T/world_size → 后续前向的激活显存随之分摊到各卡。那 token 数不能被卡数整除怎么办必须先 pad漏掉的话 compute_sequence_partition 直接抛ValueError。⚠️ pad 后总 token 数若超过max_tokensforward抛出的报错是明确的Total video token count (...) exceeds attention_manager max_tokens (...). Use a smaller resolution or fewer frames.注意力怎么跨 rank 交换all2all 内核是纯字节搬运工输入(B, T_local, H, D)→ 输出(B, T_total, H/W, D)往返逐字节一致。算子以torch.library.custom_op注册为ltx_kernels::send_recv_heads/gather_heads→torch.compile与 CUDA Graph 捕获能无 graph break 地 trace 过去 →world_size作为 int 常量进图编译缓存按 GPU 数量键控8 卡编的图绝不会被 4 卡重放。见 all_to_all.py。每张 GPU 通过 CUDA-IPC peer buffer 把数据直写进目标 GPU 的内存 buffer → 避免中间拷贝接近峰值内存带宽 → 不引入额外的 staging 副本。SM 按轮转分给目标 rankSM i 写 ranki % world_size→ 处理 SM 数不可整除的情况132 个 SM、8 卡时 rank 0–3 各 17 个、rank 4–7 各 16 个 → 每组 SM 负责其目标 rank 的全部 token。数据搬完后每个 SM 原子递增目标 rank 的 barrier 计数器SM 0 收齐所有 rank 的信号后重置计数器 → barrier 兼作死锁检测默认超时 10 秒。协议见 all2all_heads.cu。为什么不直接走 NCCL alltoall每步注意力交换的字节量大直写 IPC 少一层暂存拷贝带宽利用率更高。⚠️num_heads不能被 world_size 整除时redistribute 阶段抛ValueError——8 卡要求头数是 8 的倍数。音频交叉注意力怎么省一次洗牌Q 在本地切片不洗牌只有 K/V 走 all2all输出沿 head 维 all_gather。Q视频 token在 token 维已被前一步切分本 rank 只取自己那段 → 无需跨 rank 交换 → 比自注意力少一轮 all2all。K/V音频侧序列够短每 rank 都持有全量但头仍经send_recv_heads重新分配 → 每卡只算heads/world_size个头 → 音频侧计算同样被分摊。本地输出(B, T_local, local_heads, D)先 permute 把 head 维提到最前用all_gather_into_tensor收齐全头 → 得到本 rank token 的完整输出形式与自注意力输出一致。逻辑在 attention.py。输出怎么拼回全长每步多一次 all_gatherpad 到最大 token 数、均匀收集、按真实长度裁回、切掉 pad 行。进模型前set_seqlen_all2all把各 rank token 数推给 C 运行时 → 内核据此计算目标 buffer 里每段 token 的写入偏移 → 随后torch.distributed.barrier保证所有 rank 对齐。gather_output_tokens 先把本地输出 pad 到各 rank 最大 token 数 → all_gather 要求各 rank 输入同尺寸 → 代价是每步一次 O(T·D) 的集合通信。gather 后按各 rank 真实 token 数裁剪、沿 dim 1 拼接 → 每个 rank 都拿到全长(B, T_padded, D)。最后把尾部 pad 行切掉恢复到调用时的原始长度 → 调用方全程不需要知道 padding 存在过。完整四步串在 SequenceParallelModelWrapper.forward。落地接入在自己 runner 里 3 行启用接入只需改这 3 行# inside runner.setup(), per stage: model_cfg pipeline.stage_1._transformer_builder.model_config().get(transformer, {}) # 读头数 attn_mgr AttentionManager( # 建 all2all bufferq/k/v/heads 共 4 个实例 max_tokens32768, num_headsmodel_cfg[num_attention_heads], head_dimmodel_cfg[attention_head_dim], tensor_dtypepipeline.dtype, groupself.groups.transformer_group, ) pipeline.stage_1._transformer_builder SequenceParallelBuilder( # 包裹单卡 builder innerpipeline.stage_1._transformer_builder, attn_mgrattn_mgr, registryregistry, trackertracker, )SequenceParallelBuilder 是包裹型 buildercheckpoint 路径、量化、编译、LoRA 全部继承自 inner只是叠加并行。它把 SP module-ops 追加进 inner——把每个BasicAVTransformerBlock的attn1和video_to_audio_attn两个槽位换成 All2All 注意力mask 版也一并换未来加 mask 不会静默绕过 all2all——build()时再包一层SequenceParallelModelWrapper返回。registry是进程内共享的ModelRegistrycheckpoint 每进程只从磁盘读一次tracker是transformer_group对应的TransformerWeightTracker该 group 由 nccl_groups.py 中dist.new_group创建。前置条件一张表先看全项要求依赖包带 CUDA 的 PyTorchltx-kernels已构建SP builder 直接 import 它硬件Linux单节点 ≥2 张支持 P2PNVLink/PCIe的 CUDA GPU不支持多节点环境约束NCCL CUDA-IPC peer buffer无 macOS/Windows每卡一个进程构建命令uv sync --group kernels需 CUDA toolkit / nvcc 与 gcc 或 clangmax_tokens 超限时报什么错max_tokens必须覆盖最大的那个 step三个 MGPU runner 默认_DEFAULT_SP_MAX_TOKENS 32768。参考量级stage 1 在 512×768×121 下约6144个视频 tokendistilled shared stage 的 full-res 调用1024×1536×121约24576个。注释分别在 ti2vid_two_stages_mgpu.py 和 distilled_mgpu.py。超限时forward抛ValueError末尾那句 Use a smaller resolution or fewer frames 就是给你的行动建议确需更大上界把sp_max_tokens参数调大即可setup()均接受该参数。选型决策什么时候选 SP、什么时候不选✅ 适合场景分辨率在训练分布内且要求与单卡结果数值一致模型权重单卡放得下想换的是每步去噪延迟stage 1 或 shared stage 场景——两个自带 runner 默认都用 SP❌ 不适合场景超出训练分布的分辨率SP 没有上采样能力那是 TDP 的活单卡根本装不下模型MGPU 是延迟工具不是省显存工具每个 rank 仍持有完整模型副本多节点或非 Linux 环境内核依赖 CUDA-IPC仅限单机一句话对照TDP 面向训练分布外分辨率、只宜做 upscale 且非数值忠实分布式解码器 并行解码 latent tile两者与 SP 互不冲突ti2vid_two_stages_mgpu就是 SP TDP 分布式 VAE 的堆叠组合。环境就绪后跑python -m ltx_pipelines.ti2vid_two_stages_mgpu --checkpoint-path checkpoint --prompt A cat --output-path out.mp4你就完整体验了这条从 token 切分、all2all 换头到输出拼回的整条链路。【免费下载链接】LTX-2Official Python inference and LoRA trainer package for the LTX-2 audio–video generative model.项目地址: https://gitcode.com/GitHub_Trending/lt/LTX-2创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考