【Bug已解决】`_flash_3_varlen_hub` backend cannot handle non-contiguous mask 解决方案 【Bug已解决】_flash_3_varlen_hubbackend cannot handle non-contiguous mask 解决方案一、现象长什么样在 diffusers 里用 Flux 系模型的注意力后端_flash_3_varlen_hubFlash Attention 3 的 varlen 变长实现当传入的 attention mask 不是内存连续的non-contiguous时直接崩from diffusers.models.attention import FluxAttention # 构造一个非连续的 mask比如转置、切片、或 index_select 出来的 mask some_mask.transpose(-1, -2) # transpose 后通常 non-contiguous out attention( hidden_states, encoder_hidden_states, attn_maskmask, )报错RuntimeError: _flash_3_varlen_hub: attn_mask must be contiguous, got non-contiguous tensor或者更隐蔽的变体RuntimeError: expected contiguous tensor for mask argument (varlen backend)最迷惑的是同一个 mask如果是「直接torch.ones(...)造的」就没事如果是「经过mask[:, idx]/mask.T/mask.permute(...)处理过的」就炸。因为前者天然 contiguous后者在内存里是跨步stride存储flash 3 varlen 的 C kernel 只接受连续内存。这种 bug 在「变长序列打包varlen」场景特别常见——你要把不同长度的样本拼成一个 batch用arangeindex_select造 mask造出来的往往 non-contiguous于是喂给 flash 3 varlen 就炸。二、背景Flash Attention 3 的 varlenvariable-length变长后端用于「把多个不同长度的序列打包进一个连续张量、用 cu_seqlens 标记边界」的高效注意力。它底层是高度优化的 CUDA kernel对输入张量的内存布局有严格要求query/key/value 必须是contiguous.is_contiguous() Trueattention mask如果传也必须是 contiguous。PyTorch 里很多操作会返回non-contiguous张量却不报错因为普通 PyTorch 算子会自动处理 stridex torch.randn(4, 8) y x.transpose(0, 1) # y.is_contiguous() False z x[:, [0, 2, 3]] # 高级索引non-contiguous w x.permute(1, 0) # non-contiguous这些在普通torch.matmul里没事但 flash 3 varlen 的 kernel 为了提高性能直接读连续内存块不会为你做contiguous()转换。于是 non-contiguous 的 mask 一进去就RuntimeError。diffusers 的_flash_3_varlen_hub后端在封装时没有在入口帮用户把 maskcontiguous()于是这个「连续性假设」直接暴露给了调用方稍有不慎就炸。三、根因根因一句话_flash_3_varlen_hub这个 Flash Attention 3 varlen 后端要求 attention mask 必须是内存连续的contiguous但 diffusers 的封装没有在入口自动.contiguous()于是调用方传入经转置/切片/index_select 得到的 non-contiguous mask 时直接RuntimeError。三点展开kernel 连续性假设flash 3 varlen CUDA kernel 只读连续内存不为你转。封装未兜底_flash_3_varlen_hub入口没对 mask 做contiguous()转换把假设甩给调用方。non-contiguous 来源多transpose/permute/[:, idx]/index_select都产生 non-contiguousvarlen 打包场景极易踩。不是 mask 内容错是「内存布局不连续」被 kernel 拒绝。四、最小可运行复现不依赖真实 flash 3模拟「non-contiguous 被拒」import torch def fake_flash3_varlen_attn(q, k, v, attn_mask): # 模拟 kernel只接受 contiguous 输入 for name, t in [(q, q), (k, k), (v, v), (attn_mask, attn_mask)]: if t is not None and not t.is_contiguous(): raise RuntimeError(f_flash_3_varlen_hub: {name} must be contiguous) return q # 示意 q torch.randn(2, 4, 8) mask_contig torch.ones(2, 4, 4) mask_nc mask_contig.transpose(-1, -2) # non-contiguous print(contiguous mask OK:, fake_flash3_varlen_attn(q, q, q, mask_contig) is not None) try: fake_flash3_varlen_attn(q, q, q, mask_nc) except RuntimeError as e: print(non-contiguous 炸:, e) # 修复先 contiguous mask_fixed mask_nc.contiguous() print(修复后 OK:, fake_flash3_varlen_attn(q, q, q, mask_fixed) is not None)跑出来contiguous 正常transpose 出的 non-contiguous 直接RuntimeError.contiguous()后恢复。这就是「非连续 mask 崩」的精确复现。五、解决方案第一层最小直接修复最小修复在把 mask 传给 flash 3 varlen 后端之前先.contiguous()同理 query/key/value 若来自 non-contiguous 来源也要.contiguous()。import torch def safe_flash3_varlen_attn(attn_module, hidden_states, encoder_hidden_states, attn_maskNone): # 所有进 kernel 的张量都先保证连续 q hidden_states if not q.is_contiguous(): q q.contiguous() if encoder_hidden_states is not None and not encoder_hidden_states.is_contiguous(): encoder_hidden_states encoder_hidden_states.contiguous() mask attn_mask if mask is not None and not mask.is_contiguous(): mask mask.contiguous() # 关键mask 转连续 return attn_module( hidden_statesq, encoder_hidden_statesencoder_hidden_states, attn_maskmask, ) # 用法哪怕 mask 来自 transpose / 切片也先过这一层 mask raw_mask.transpose(-1, -2).contiguous() # 或者在这里 contiguous out safe_flash3_varlen_attn(attention, h, enc, attn_maskmask)要点attn_mask进 kernel 前.contiguous()transpose/切片/index_select 不再炸。query/key/value 同样处理避免任何 non-contiguous 进 kernel。在封装层统一做调用方无需关心来源。这一步单独就让 flash 3 varlen 不再因连续性报错。六、解决方案第二层结构性改进第一层是「在调用处加.contiguous()」。但 diffusers 里多个注意力后端flash 2/3、sdpa、varlen都可能对连续性有要求散落加容易漏。更稳的做法把「进注意力前统一保证连续性」收敛成单一守卫。from dataclasses import dataclass, field from typing import Optional import torch dataclass class MaskContiguityGuard: 注意力输入连续性保证的单一守卫。 # 是否对所有输入做 contiguous保险起见默认 True force_contiguous: bool True def ensure(self, *tensors: Optional[torch.Tensor]): out [] for t in tensors: if t is None: out.append(None) elif self.force_contiguous and not t.is_contiguous(): out.append(t.contiguous()) else: out.append(t) return out def prepare(self, hidden_states, encoder_hidden_statesNone, attn_maskNone): q, enc, mask self.ensure(hidden_states, encoder_hidden_states, attn_mask) return q, enc, mask def check(self, *tensors: Optional[torch.Tensor]) - list: # 返回哪些 non-contiguous用于调试/CI 断言 return [t is not None and not t.is_contiguous() for t in tensors] # 用法 guard MaskContiguityGuard() q, enc, mask guard.prepare(hidden_states, encoder_hidden_states, attn_mask) out attention(q, encoder_hidden_statesenc, attn_maskmask)结构收益单一守卫所有进注意力 kernel 的张量统一过ensure连续性假设集中处理。可观测check返回哪些 non-contiguousCI 可断言「进入 kernel 前已全部连续」。可扩展换后端flash3→flash3 varlen→sdpa都复用同一守卫。七、解决方案第三层断言 / CI 守护写 pytest 守三条(1) non-contiguous 输入被转连续(2) 已连续的输入不被无谓复制(3) 修复后 kernel 不再报错。import torch import pytest from your_lib import MaskContiguityGuard def test_noncontiguous_became_contiguous(): guard MaskContiguityGuard() nc torch.randn(2, 4, 4).transpose(-1, -2) # non-contiguous (out,) guard.ensure(nc) assert out.is_contiguous() def test_already_contiguous_kept(): guard MaskContiguityGuard() c torch.randn(2, 4, 4) (out,) guard.ensure(c) assert out.data_ptr() c.data_ptr() # 不应重新分配 def test_none_passthrough(): guard MaskContiguityGuard() (out,) guard.ensure(None) assert out is None def test_kernel_no_longer_errors(): # 模拟 kernel 拒绝 non-contiguous def kernel(mask): if not mask.is_contiguous(): raise RuntimeError(mask must be contiguous) return True guard MaskContiguityGuard() nc torch.randn(2, 4, 4).transpose(-1, -2) (mask,) guard.ensure(nc) assert kernel(mask) is True def test_prepare_tuple_order(): guard MaskContiguityGuard() h torch.randn(2, 4, 8)[:, [0, 1, 3, 2], :] # nc enc torch.randn(2, 4, 8) mask torch.ones(2, 4, 4).transpose(-1, -2) # nc q, e, m guard.prepare(h, enc, mask) assert q.is_contiguous() and e.is_contiguous() and m.is_contiguous()CI 常驻跑这五条后任何「又传 non-contiguous 进 kernel」的回归都会立刻爆红。八、排查清单flash 3 varlen 报「mask 不连续」时按顺序查先确认报错是否含must be contiguous/expected contiguous tensor——是的话定位连续性。打印attn_mask.is_contiguous()若为False即根因。检查 mask 来源transpose/permute/[:, idx]/index_select都产 non-contiguous。进 kernel 前对attn_mask及 q/k/v.contiguous()。把连续性保证收敛到MaskContiguityGuard别散落手写。注意.contiguous()会分配新内存高频路径可预先构造 contiguous mask 复用。升级 diffusers/flash-attn 后跑「varlen 打包 非连续 mask」冒烟断言不再RuntimeError。九、小结_flash_3_varlen_hub报「mask 不连续」根子是 Flash Attention 3 varlen 的 CUDA kernel 只读连续内存、不为你转 contiguous而 diffusers 封装没在入口自动.contiguous()调用方传入经transpose/切片/index_select 得到的 non-contiguous mask 就RuntimeError。修复三层次第一层进 kernel 前对attn_mask及 q/k/v.contiguous()第二层用MaskContiguityGuarddataclass 把连续性保证收敛为单一守卫第三层用 pytest 守「non-contiguous 转连续」「已连续不复制」「kernel 不再报错」。工程启示任何调用底层优化 kernelflash-attn、xformers、自定义 CUDA的封装都必须在入口统一保证输入张量的内存连续性和 dtype。PyTorch 普通算子容错性强但 C kernel 往往只认 contiguous 特定 dtype。把连续性/dtype 保证做成守卫层调用方永远不必担心「我这个 tensor 是怎么来的」。