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

资讯详情

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

LTX-2 多GPU序列并行实战:token 切分与 all2all 内核全解析

LTX-2 多GPU序列并行实战:token 切分与 all2all 内核全解析 LTX-2 多GPU序列并行实战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跑一段 1024×1536×121 的 denoising单卡要磨很久把卡堆起来却不是为了分显存——LTX-2 的多 GPU 推理里transformer 工作副本在每张卡上都是完整复刻。那多卡到底怎么提速答案就是序列并行Sequence ParallelismSP把视频 token 序列均分给各卡每卡只算自己那片再用 all2all 内核保住全局注意力。本文从真实痛点出发拆透 LTX-2 这套方案的原理、内核实现与接入方式。一、切了 token 还能逐字节一致先弄清 SP 的承诺把视频 token 想象成一条长面包SP 把它按world_size均分每个 rank 只持有自己那一截在本地跑 transformer 前向结束后把输出 all-gather 回所有 rank见 sequence_parallel.py。你可能会问自注意力要求每个 token 看到所有 token切开了怎么算SP 的答案是换头而不是切注意力——Q/K/V 沿 head 维跨 rank 交换交换后每个 rank 持有全部 token 的一部分 head本地算完再洗牌回去。所以 token 间的交互一个都没少。这也是官方反复强调的忠实性见 sequence-parallel.md注意力保持全局all2all 只是 token×head 二维数据的一次分布变换唯一的数值差异来自浮点归约顺序all2all 内核只搬运字节往返gather(send(x)) x是逐字节精确的。结论很直接当你需要和单卡结果一致、只是更快时SP 就是正解——它不改变模型行为只改变计算在硬件上的分布方式。二、前向四步走pad 齐、切片、算、拼回每个 denoising step 中SequenceParallelModelWrapper.forward干四件事pad 对齐 → 切出本 rank 切片 → 跑模型 → all-gather 还原。第一步把 token 数 pad 到整除compute_sequence_partition要求总 token 数能被world_size整除否则直接抛ValueError。这个死板的要求其实很聪明均匀 sharding 让 All2All 自定义算子的 fake-impl 能从输入 shape 符号化推导输出 shapex.shape[1] * world_size或// world_size而不用依赖 Python int 参数——这是后面torch.compile能无 graph break 的前提。补齐由pad_modality_for_uniform_sharding负责mask 处理有两个讲究原本没有 mask 时构造key-only padding mask[0,1]形式有效 key 为 1、pad key 为 0shape 仅(1,1,T_padded)沿 batch 和 query 广播内存O(T)而不是物化稠密的(B,T,T)矩阵用户传了(B,T,T)mask 时pad 行/列会被扩展且pad 的 query 行允许 attend 所有有效 key——输出反正会被切掉但全 masked 的行会产生 NaN。第二步到第四步切片、跑模型、还原tile_modality_for_rank把 latent / timesteps / positions 按 rank 切片keyframes mask 也随行切走因为 embedding 在 all2all 之前逐 rank 施加。随后模型前向此时视频自注意力attn1与 video→audio 交叉注意力都已被 patch 成 all2all 版本。最后gather_output_tokens先把本地输出 pad 到最大长度做均匀all_gather再按各 rank 真实 token 数裁回拼接并切掉 pad 行恢复调用方的原始长度。三、all2all 内核解剖IPC 直写 SM 轮转自定义算子是ltx_kernels.All2Allltx-kernels包。内核实现all2all_heads.cu头部注释把算法讲得很透直写策略每张 GPU 通过 CUDA-IPC peer buffer 把数据直接写进目标 GPU 的内存没有中间拷贝带宽接近峰值SM 轮转分配SM i负责写 ranki % world_size天然处理 SM 数除不尽的情况——例如 132 个 SM、8 张卡时rank 0–3 各 17 个 SMrank 4–7 各 16 个每个 SM 组用 strided 方式覆盖分配给它的那个 rank 的全部 tokenbarrier 同步数据搬完后各 SM 原子递增目标 rank 的 barrier 计数器SM 0 等齐所有 rank 的信号再重置计数器供下一轮使用allgather 路径是barrier_wait_and_reset_roundrobin带timeout_cycles超时。⚡ 更关键的是它对编译友好的设计见 all_to_all.py算子用torch.library.custom_op注册为ltx_kernels::send_recv_heads与gather_headstorch.compile含reduce-overhead下的 CUDA Graph 捕获能无 graph break 地 trace 过去。world_size以 int 常量进入 traced graphDynamo 的 guard 只按 GPU 数量键控编译缓存——同一张图绝不会被换个卡数重放而每步变化的 per-rank token 数则留在 C 运行时set_rank_tokens不经过算子。另外两个细节值得一提copy_outFalse时send_recv_heads返回IPC buffer 的零拷贝视图因为该 buffer 是cudaMalloc分配不在静态图池内cudagraph_trees 下也安全实例通过weakref.finalize在销毁时自动释放 CUDA/IPC 资源。四、落地代码AttentionManager 与 SequenceParallelBuilder 各管什么AttentionManager管 buffer、管每步计数attention.py 中的AttentionManager持有 all2all buffer每个 rank 大小为ceil(max_tokens / world_size)个 tokenfrom ltx_core.multigpu.transformer.attention import AttentionManager attn_mgr AttentionManager( max_tokens32768, # 视频总 token 数上界 num_headsmodel_cfg[num_attention_heads], head_dimmodel_cfg[attention_head_dim], tensor_dtypepipeline.dtype, groupself.groups.transformer_group, )几个要点num_heads必须能被world_size整除_All2AllRedistribute.redistribute里会显式校验构造时惰性 importltx_kernels让 multigpu 模块在内核未安装时如 CPU CI 收集测试仍可导入内部创建 q/k/v/heads 共 4 个All2All实例copy_out_True时 k、v 与 q 共用。还有一个容易踩的坑all2all_timeout_seconds默认 10s对应内核configs.cuh的DEFAULT_BARRIER_TIMEOUT_SECONDS是 barrier 死锁检测超时。torch.compile首次前向时某个 rank 重编译会让它的内核启动晚于稳态超时、误触 barrier——setter 注释给出的策略是首次 compile 前向临时调大之后再复位。每步的 token 计数则由set_seqlen_all2all同步到 4 个实例的 C 运行时随后torch.distributed.barrier保证所有 rank 对齐再进模型。SequenceParallelBuilder一个包裹型buildersp_builder.py 的SequenceParallelBuilder只接受SingleGPUModelBuilder否则抛TypeError做三件事把 registry 与 LoRA 加载设备注入 inner builder用create_video_self_attention_module_ops生成 SP module-ops 追加到列表build()时经TransformerWeightTracker构建模型再包一层SequenceParallelModelWrapper。# inside runner.setup(), per stage pipeline.stage_1._transformer_builder SequenceParallelBuilder( innerpipeline.stage_1._transformer_builder, # 原单 GPU builder attn_mgrattn_mgr, registryregistry, # 进程内共享的 ModelRegistry trackertracker, # transformer_group 的 Tracker )这正是 MGPU 管线的统一模式单 GPU 管线 替换各 block 的 builder。因为SequenceParallelBuilder是包裹inner所以完整继承 checkpoint 路径、量化、编译和 LoRA 配置只叠加并行。registry让 checkpoint 每进程只从磁盘加载一次。module-ops 的注入细节create_video_self_attention_module_ops对每个BasicAVTransformerBlockattn1的attention_function换成All2AllAttention、masked_attention_function换成MaskedAll2AllAttentionvideo_to_audio_attn换成AudioAll2AllAttention/MaskedAudioAll2AllAttention。目前无人给该路径传 mask但两个槽位都换未来有人加 mask 时 SP 管道已就位不会静默绕过 All2All。视频自注意力与音频交叉注意力洗牌方式不一样视频自注意力All2AllAttentionQ/K/V 全走send_recv_heads——每个 rank 收齐全部 token 的一个 head 子集本地算heads // world_size个头再gather_heads换回音频交叉注意力AudioAll2AllAttentionQ 只在本地按 rank 切片不做跨 rank 洗牌音频序列短复制一份即可只有 K/V 走send_recv_heads输出沿 head 维用all_gather_into_tensor收集_AudioAll2AllRedistribute注释写得很明白。五、max_tokens 设多少哪些场景别选 SPmax_tokens必须覆盖最大的那个 step参考量级场景分辨率/帧数视频 token 量级备注stage 1512×768×121~6144代码注释精确值见 ti2vid_two_stages_mgpu.pydistilled shared stagefull-res1024×1536×121~24576见 distilled_mgpu.py两者默认—32768_DEFAULT_SP_MAX_TOKENS超出上界时forward会抛出明确的ValueError附带Use a smaller resolution or fewer frames.。需要更大上界时三个 MGPU runner 的setup()都接受sp_max_tokens参数显式调大。运行前提清单仅LinuxNCCL 与 CUDA-IPC peer buffer 都是 Linux-only单节点 ≥2 张支持 P2PNVLink/PCIe的 CUDA GPU不支持多节点带 CUDA 的 PyTorch且ltx-kernels已构建uv sync --group kernels需要 nvcc 与 gcc/clangSP builder 会硬性 import 它。什么时候该选 SP、什么时候不该MGPU 的定位要先摆正它是延迟工具不是显存工具——transformer 工作副本每卡都是完整复刻SP 额外分摊的是激活内存它救不了模型塞不进单卡的场景那是 FP8 量化和权重重载的活。再看与 TDP 的分工见 multigpu/README.md 能力矩阵维度SP 序列并行TDP 分块数据并行切分维度token序列维均匀切分每卡一个空间高×宽tile数值行为忠实与单 GPU 数值等价面向训练分布外分辨率upscale-only典型位置stage 1 与 shared stage 的默认方案stage 2 全分辨率如ti2vid_two_stages_mgpu下一步可以做什么直接跑python -m ltx_pipelines.ti2vid_two_stages_mgpu或ltx_pipelines.distilled_mgpu端到端体验 SP——前者 stage 1 用 SP、stage 2 切 TDP后者一个 SP 包裹同时覆盖 half-res 与 full-res。如果你要写自定义 MGPU runner照第四节的runner.setup()接入模式把目标 stage 的单 GPU builder 换成SequenceParallelBuilder即可若你的分辨率落在训练分布外把该 stage 换成TiledDataParallelBuilderSP 与 TDP 按 stage 各取所需。【免费下载链接】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),仅供参考
返回列表