LTX-2 多 GPU 序列并行:从 token 切分到 all2all 直写,多卡数值等价的 5 个落地要点
2026/9/18 13:18:47 网站建设 项目流程

LTX-2 多 GPU 序列并行:从 token 切分到 all2all 直写,多卡数值等价的 5 个落地要点

【免费下载链接】LTX-2Official Python inference and LoRA trainer package for the LTX-2 audio–video generative model.项目地址: https://gitcode.com/GitHub_Trending/lt/LTX-2

拆解 LTX-2 多 GPU 推理里的序列并行(Sequence Parallelism,SP):token 维均匀切分、all2all 内核的 CUDA-IPC 直写、AttentionManager 与 SequenceParallelBuilder 的接入,以及 max_tokens 报错语义。适合有 PyTorch 基础、准备接入多卡 SP 或做并行选型的推理工程师。

SP 选型决策清单:三个条件选中序列并行

先把选型讲清楚。LTX-2 的多 GPU 体系给不同子问题提供了三种切法,各管一摊:

  • SP(序列并行):沿 token 维切视频,结果与单卡数值等价
  • TDP(分块数据并行):每卡一个空间 tile,仅 upscale 场景
  • 分布式 Gemma:用 Accelerate 切分文本编码器,不碰 transformer

换句话说,MGPU 是延迟工具而非显存工具:transformer 的工作副本在每个 GPU 上都是完整副本,SP 额外把激活内存摊到各 rank。所以 SP 的决策清单只有三条:

  1. 分辨率在训练分布内?是则 SP 优先,upscale 转 TDP
  2. "结果必须和单卡一致"是硬需求?SP 是默认答案
  3. 目标是单次生成更低延迟?SP 正对此优化

官方管线里 SP 是 stage 1(ti2vid_two_stages_mgpu)与 shared stage(distilled_mgpu)的默认方案,后者的一个 SP 包裹同时覆盖 half-res 与 full-res 两次调用;stage 2 全分辨率则交给 TDP。分工一句话:分布内 + 求忠实 + 追延迟 → SP

心智模型:token 怎么切、head 怎么换的四步走

一个 denoising step 里,sequence_parallel.py 的SequenceParallelModelWrapper.forward走四步(下图为文本版数据流图,world_size=4):

T=14 tokens → pad 到 T'=16(world_size 的倍数),pad key 被 mask rank0: tok[0..3] rank1: tok[4..7] rank2: tok[8..11] rank3: tok[12..15] │ │ │ │ └─▶ send_recv_heads(all2all 换头):每 rank 取到"全部 token × 本地那部分 head" ▼ 本地注意力(每 rank 只算 heads/4 个 head) └─▶ gather_heads 洗回 → all_gather 回全长 → 切掉 pad 行 → 每 rank 得完整输出

第一步 pad 对齐:seq 维补齐到world_size的整数倍,保证每 rank 拿到等量 shard。compute_sequence_partition对"总数不整除"直接抛ValueError——均匀 sharding 还能让 all2all 自定义算子的 fake-impl 从输入 shape 符号化推导输出 shape。

pad 的 mask 处理有个讲究:原本没有 attention mask 时,构造key-only padding mask,shape(1, 1, T_padded),有效 key 为 1、pad 为 0,沿 batch 与 query 广播——只占 O(T) 内存,而不是物化稠密(B, T, T)矩阵。若用户传了(B, T, T)mask,则 pad 的 query 行被允许 attend 所有有效 key,softmax 才良定义(输出反正会切掉,但全 masked 行会产生 NaN)。

第三步 all2all 换头是自注意力"每个 token 看到所有 token"这一要求的解法:Q/K/V 的 head 跨 rank 交换,每个 rank 最终持有所有 token 的某一部分 head,本地算完再洗回。

这里点破"忠实"的来源:all2all 只改变 token×head 二维数据的分布形态,没有牺牲任何 token 间交互,与单卡的唯一数值差异是浮点归约顺序;内核只搬运字节,往返gather(send(x)) == x逐字节精确。说白了,SP 不改模型行为,只改硬件分布方式——要"更低延迟下的单卡同结果",它就是正确选择。

all2all 直写策略与 SM 轮转分配

通信原语ltx_kernels.All2All的 CUDA 实现在 all2all_heads.cu,干了三件脏活:

  • 直写:经 CUDA-IPC peer buffer 直写目标 GPU 的内存 buffer,免中间拷贝,接近峰值内存带宽
  • SM 轮转SM i 写 rank (i % world_size)。132 个 SM、8 张卡时,rank 0–3 各 17 个、rank 4–7 各 16 个;每组 SM 覆盖其目标 rank 的全部 token
  • barrier 同步:搬完后各 SM 原子递增目标 rank 的 barrier 计数器,SM 0 等齐所有 rank 的信号再复位计数器供下一轮;等待超过超时周期(默认 10 秒)即触发死锁检测

Python 侧用torch.library.custom_op注册send_recv_headsgather_headstorch.compile(含mode="reduce-overhead"的 CUDA Graph 捕获)能无 graph break 地 trace 过去:

@custom_op("ltx_kernels::send_recv_heads", mutates_args=(), device_types="cuda") def _send_recv_heads_op(x, comm_id, world_size, copy_out): return All2All._runtime_registry[comm_id].send_recv_heads(x, copy_out)

这里有个刻意的设计:world_size作为 int 常量进算子,Dynamo 的 guard 只按 GPU 数量键控编译缓存——同一个图绝不会在另一个 GPU 数量下被重放;而每步变化的 per-rank token 数走set_rank_tokens下发到 C++ 运行时,不经过算子。另外copy_out=False时返回的是 IPC buffer 的零拷贝视图(buffer 由cudaMalloc分配、不在静态图池内,cudagraph_trees 下也安全),实例销毁时由weakref.finalize自动释放 CUDA/IPC 资源。

接入两个核心 API:AttentionManager 与 SequenceParallelBuilder

AttentionManager 构造的四个要点

attention.py 的AttentionManager持有 all2all buffer(每 rankceil(max_tokens / world_size)个 token):

from ltx_core.multigpu.transformer.attention import AttentionManager attn_mgr = AttentionManager( max_tokens=32768, # 视频总 token 数上界 num_heads=model_cfg["num_attention_heads"], head_dim=model_cfg["attention_head_dim"], tensor_dtype=pipeline.dtype, group=self.groups.transformer_group, )
  • num_heads须被world_size整除,redistribute中显式抛ValueError
  • 构造时才惰性 importltx_kernels,让 multigpu 模块在未装内核的 CPU CI 上仍可导入
  • 内部创建 4 个All2All实例(q / k / v / heads);copy_out_=True时 k、v 与 q 共用实例
  • all2all_timeout_seconds(默认 10.0s)管理 barrier 死锁检测

💡 超时属性对应一个真实坑:torch.compile首次前向时,某 rank 的重编译可能让它启动内核晚于稳态超时,从而撞爆 barrier。应对写在 setter 注释里——首次 compile 前向临时调大超时,之后再复位。

Builder 包裹单卡 builder 的三步

sp_builder.py 的SequenceParallelBuilder是包裹型 builder,只接受SingleGPUModelBuilder(否则TypeError),构造与构建分三步:

pipeline.stage_1._transformer_builder = SequenceParallelBuilder( inner=pipeline.stage_1._transformer_builder, # 该 stage 的单卡 builder attn_mgr=attn_mgr, registry=registry, # 进程内共享的 ModelRegistry tracker=tracker, # TransformerWeightTracker )
  1. 注入:把 registry 与 LoRA 加载设备(cuda:当前设备)注入 inner;
  2. module-ops 注入create_video_self_attention_module_ops匹配LTXModel,对每个BasicAVTransformerBlockattn1attention_function换为All2AllAttentionmasked_attention_function换为MaskedAll2AllAttentionvideo_to_audio_attn交叉注意力同理换成AudioAll2AllAttention/MaskedAudioAll2AllAttention。masked 槽位目前没有调用方(死代码路径),照样换掉——未来若有人加 mask,SP 管线已就位,不会静默绕过 All2All;
  3. build() 包裹:经TransformerWeightTracker构建模型,再包一层SequenceParallelModelWrapper返回。

两类洗牌有差异:视频自注意力的 Q/K/V 全走send_recv_heads,本地算heads // world_size个 head 后再gather_heads洗回;音频交叉注意力的Q 本地按 rank 切片、不跨 rank 洗牌(音频序列短,可复制),仅 K/V 走 all2all,输出沿 head 维all_gather_into_tensor收集。因为整个接入是"包裹",它继承 inner 的 checkpoint 路径、量化、编译与 LoRA 配置,只叠加并行——这正是 MGPU 体系"单卡管线 + 替换 builder"的模式。

max_tokens 上界、参考量级与报错语义

max_tokens决定 all2all buffer 尺寸,必须覆盖最大的那个 step。参考量级:

场景形状视频 token 数默认上界
stage 1512×768×121≈ 614432768
distilled full-res1024×1536×121≈ 2457632768

三个 MGPU runner 都默认_DEFAULT_SP_MAX_TOKENS = 32768。一旦超出,SequenceParallelModelWrapper.forward抛出信息明确的ValueError"Use a smaller resolution or fewer frames."需要更大上界时,显式给sp_max_tokens传更大的值(三个 MGPU runner 的setup()都接受该参数)——同时留意 buffer 显存成本的同步放大。

SP 运行前提与排错清单

  • 仅 Linux:NCCL 与 CUDA-IPC 是 Linux-only
  • 单节点 ≥2 张 P2P(NVLink/PCIe)CUDA GPU
  • 不支持多节点,每 GPU 一个进程
  • PyTorch 需带 CUDA
  • uv sync --group kernels构建 ltx-kernels(需 nvcc + gcc 或 clang)
  • 头数整除卡数:heads % world_size ≠ 0 抛 ValueError
  • token 数不整除卡数:pad 自动处理,无需干预
  • 首次 compile 撞 barrier 超时:临时调大 all2all_timeout_seconds

⚠️ 看到 "Use a smaller resolution or fewer frames." 不是 bug,是保护:降分辨率/帧数,或调大sp_max_tokens

💡 选型口诀:分布内求忠实 → SP;upscale-only → TDP;单文本编码器放不下 → 分布式 Gemma。

下一步:先跑哪个 runner

想上手,直接跑ti2vid_two_stages_mgpu的 CLI(完整示例见 multigpu 文档):SP 负责 stage 1、TDP 负责 stage 2,一条管线内能看到两种策略协作。要看 SP 用同一个包裹覆盖 half-res 与 full-res 两次调用,再看distilled_mgpu。token 切分、all2all 换头、gather 还原、AttentionManager 管 buffer、Builder 包裹接入——这 5 个要点吃透后,写自定义 MGPU runner 只是拼装工作。

【免费下载链接】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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询