尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

【Bug已解决】Finalize Guidelines for Ulysses based context parallel attention 解决方案

【Bug已解决】Finalize Guidelines for Ulysses based context parallel attention 解决方案 【Bug已解决】Finalize Guidelines for Ulysses based context parallel attention 解决方案一、现象长什么样Ulysses 风格序列并行 / context parallel的注意力是把一条长序列沿序列维切到多张卡上每张卡只算自己那段 query 对「全局 key/value」的注意力再通过全收集all-gather拿回完整输出。社区在 diffusers 的 transformers 里加这类注意力后端时文档和落地代码经常出现「指南没定稿」导致的几类问题。最直观的是照着半成品文档写代码在序列并行尺寸 1 时报RuntimeError: The expanded size of the tensor (4096) must match the existing size (1024) at non-singleton dimension 1或者通信维度对不上ValueError: context_parallel_size4 but sequence length 3000 is not divisible by cp_size更隐蔽的一类是「能跑但数值错」每张卡只用本地 query 和本地 key 算注意力忘了和别的卡交换 KV结果长程依赖全丢loss 比单卡高一大截却什么都不报。二、背景Ulysses 注意力的核心切分规则是序列维S沿context_parallel_size简称cp_size均匀切分每张卡持有S / cp_size个 token 的 query但注意力的 key/value 必须覆盖全部S个 token否则就是「局部注意力」而不是 Ulysses实现上通常用 all-to-all 或 all-gather 把 K/V 在序列维做通信计算完再还原因为序列长度必须能被cp_size整除否则切片会产生非整数边界。diffusers 在支持长上下文图像/视频生成如长视频、超大分辨率 tiling时需要把这类注意力接进AttentionProcessor。「Finalize Guidelines」诉求的就是把上面这些约束从口头约定变成可落地的、带校验的指南避免每个人接一个后端就踩一遍坑。三、根因根因可以归到「指南未固化导致三处不一致」序列长度整除性未强制cp_size和S的整除检查只在某几个后端里写了别的后端没写于是出现3000 不能被 4 整除这类ValueError且报错信息没说清该怎么 pad。KV 通信语义不一致有的 processor 在forward里只做本地注意力漏掉 all-gather KV有的做对了。没有统一指南时新接的人大概率复制了「本地版」长程依赖静默丢失。输出还原维度错all-gather 之后维度是cp_size * local_S如果还原时没按dim1或dim2取决于 KV 布局正确 reshape 回[B, S, ...]就会在和别的层拼接时炸size mismatch。本质Ulysses 注意力的切分/通信/还原三步约定散落在各个 processor没有单一真源于是每接一个新后端都在重复踩坑。四、最小可运行复现用真实 PyTorch 通信原语复现「整除性」与「KV 未通信」两类问题import torch import torch.distributed as dist def ulysses_attention_naive(q, k, v, cp_size): # q/k/v: [B, S, H, D]假设已沿 S 切到本卡 local_s q.shape[1] # 错误示范直接用本地 k/v不做 all-gather scores torch.einsum(bshd,bthd-bsht, q, k) # 只看本地 key attn scores.softmax(-1) return torch.einsum(bsht,bthd-bshd, attn, v) # 长程依赖丢失 # 复现整除性报错 cp_size 4 S 3000 if S % cp_size ! 0: raise ValueError( fcontext_parallel_size{cp_size} 但序列长度 {S} 不能被 cp_size 整除 f请对序列 pad 到 { ((S cp_size - 1) // cp_size) * cp_size } )要复现「数值错但不报错」把上面ulysses_attention_naive套到一个长序列训练里对比单卡实现会看到 loss 明显偏高却无异常。五、解决方案第一层最小直接修复最小修复是给每个 Ulysses processor 的forward加三件事整除校验、KV all-gather、输出正确还原。下面给一个可直接抄的最小正确版用torch.distributed通信原语import torch import torch.distributed as dist from torch.nn.functional import scaled_dot_product_attention def ulysses_attention(q, k, v, cp_group, cp_size): # q/k/v: [B, S_local, H, D] assert q.shape[1] * cp_size _global_seq_len(q, cp_group), S 必须能被 cp_size 整除 # 1) 沿序列维 all-gather K/V凑齐全局 [B, S, H, D] k_full torch.cat(dist.all_gather(k, groupcp_group), dim1) if cp_size 1 else k v_full torch.cat(dist.all_gather(v, groupcp_group), dim1) if cp_size 1 else v # 2) 本地 query 对全局 KV 做注意力 q q.transpose(1, 2) # [B, H, S_local, D] k_full k_full.transpose(1, 2) v_full v_full.transpose(1, 2) out scaled_dot_product_attention(q, k_full, v_full) # [B, H, S_local, D] # 3) 还原成 [B, S_local, H, D]序列维仍是本地段符合 Ulysses 输出约定 return out.transpose(1, 2).contiguous()这个版本把「KV 必须全局」和「输出仍是本地段」两条规则写死复制粘贴也不会丢长程依赖。六、解决方案第二层结构性改进把 Ulysses 接入约定收敛成一个 dataclass 单一真源并提供一个统一的AttentionProcessor基类所有 Ulysses 后端都继承它避免每个后端各自实现切分逻辑from dataclasses import dataclass, field from typing import List dataclass(frozenTrue) class UlyssesContextParallelPolicy: Ulysses 序列并行注意力接入的单一真源。 backend_name: str ulysses # 序列必须能被 cp_size 整除不满足时是否自动 pad require_seq_divisible: bool True auto_pad: bool True pad_multiple: int 1 # pad 到 cp_size * pad_multiple 的倍数 # KV 必须全局是否要求 all-gather kv_must_be_global: bool True # 输出还原维度沿序列维还原回本地段布局 output_restore_dim: int 1 # 要求所有 Ulysses processor 继承的统一基类 base_processor: str UlyssesAttentionProcessor # 校验项清单 required_checks: List[str] field(default_factorylambda: [ seq_divisible, kv_global_gathered, output_restore_dim, ]) def padded_seq_len(self, seq_len: int, cp_size: int) - int: if not self.require_seq_divisible: return seq_len base cp_size * self.pad_multiple return ((seq_len base - 1) // base) * base def validate_config(self, seq_len: int, cp_size: int) - List[str]: problems [] if self.require_seq_divisible and seq_len % cp_size ! 0: problems.append(fseq_len{seq_len} 不能被 cp_size{cp_size} 整除) return problems落库时所有 Ulysses processor 放在diffusers/models/attention_processor_ulysses.py统一继承UlyssesAttentionProcessorforward模板强制调用policy.validate_config与policy.padded_seq_len。文档指南直接由这个 dataclass 生成保证「代码即文档」。七、解决方案第三层断言 / CI 守护用 pytest 把「整除校验、KV 全局、输出维度」固化成回归最好在多卡 CI如 2/4 卡里跑import pytest import torch from mylib.ulysses_policy import UlyssesContextParallelPolicy POLICY UlyssesContextParallelPolicy() def test_seq_divisibility_guard(): assert POLICY.validate_config(seq_len3000, cp_size4) ! [] # 应报错 assert POLICY.validate_config(seq_len3008, cp_size4) [] # pad 后可整除 assert POLICY.padded_seq_len(3000, 4) 3008 def test_kv_must_be_global_contract(): # 任何 Ulysses processor 的 forward 必须 all-gather KV assert POLICY.kv_must_be_global is True assert kv_global_gathered in POLICY.required_checks pytest.mark.multi_gpu(4) def test_ulysses_matches_single_card(): # 4 卡 Ulysses 输出应与单卡全序列注意力数值接近 from mylib.ulysses import ulysses_attention out_cp ulysses_attention(q_local, k_local, v_local, cp_group, cp_size4) out_ref reference_full_attention(q_full, k_full, v_full) assert torch.allclose(out_cp, out_ref局部段, atol1e-3) def test_output_restore_dim(): assert POLICY.output_restore_dim 1CI 里加一个「Ulysses 后端新增必须继承UlyssesAttentionProcessor且通过test_ulysses_matches_single_card」的门禁杜绝「复制本地版丢长程依赖」的回归。八、排查清单Ulysses 注意力异常按顺序查序列长度能否被cp_size整除不能就用policy.padded_seq_len先 pad否则ValueError。K/V 是否在forward里 all-gather 成全局漏了长程依赖静默丢失loss 偏高但无报错。输出是否还原成[B, S_local, H, D]本地段布局还原dim错会出现size mismatch。all-gather 的通信组cp_group是否正确用错进程组会让卡之间交换到错误分片。scaled_dot_product_attention的q/k/v头维是否 transpose 一致顺序不一致会得到形状对但数值乱的结果。梯度是否需要在通信后reduce_scatter训练时 Ulysses 通常要在输出反向前reduce_scatter梯度漏了会导致多卡梯度重复累加。九、小结Ulysses 序列并行注意力的「Bug」本质是指南未固化为单一真源导致整除校验、KV 全局通信、输出还原维度三处约定各后端各写一套。第一层在每个 processor 的forward里补上整除校验 KV all-gather 正确还原第二层用UlyssesContextParallelPolicydataclass 和统一基类UlyssesAttentionProcessor收敛所有约定第三层用 pytest含多卡数值对齐测试守住「不丢长程依赖、形状正确」。把这套当成长上下文注意力后端的接入规范新后端一次写对。
返回列表