【Bug已解决】Incompatibility between FlashAttention and ERNIE-Image 解决方案
一、现象长什么样
ERNIE-Image 在自己的注意力实现里用 FlashAttention 加速,但装上某个版本的 FlashAttention 后,跑推理直接崩:
from diffusers import ErnieImagePipeline pipe = ErnieImagePipeline.from_pretrained("baidu/ernie-image") pipe.enable_flash_attn_2() # 开启 FA2 image = pipe("a cat").images[0]报错:
RuntimeError: expected query, key, value to have the same sequence length dimension, but got q: [1, 4096, 64] and k: [1, 64, 4096] (layout 不匹配)或者:
AttributeError: flash_attn_func() got an unexpected keyword argument 'window_size'又或静默数值错(开了 FA2 但图和不开时不一样,且更糊)。
现象总结:ERNIE-Image 的注意力代码按「某一种 FlashAttention 的布局/API」写的(如 q/k/v 的维度顺序、是否接受window_size、是否要求causal),但用户环境里的 FlashAttention 版本布局/API 不同,于是维度不匹配、参数不匹配,或静默数值错。
二、背景
FlashAttention 的接口在不同版本/变体里有差异:
- FA1 / 早期 FA2:
flash_attn_func(q, k, v, causal=...),q/k/v 形状约定为[B, S, H, D](或[B, S, 3, H, D]packed); - 部分集成版:要求 q/k/v 是
(B, S, H, D),但有的实现内部把 H 和 S 顺序搞反(ERNIE 可能按[B, H, S, D]组织); - window_size / causal:有的版本
flash_attn_func接受window_size做滑动窗口,有的没有这个参数。
ERNIE-Image 的注意力模块在写的时候,假定了某种布局(比如它内部q是[B, H, S, D],直接丢给flash_attn_func),而用户装的 FA 版本期望[B, S, H, D],于是 q/k/v 的序列维和头维对不上 → RuntimeError。或者它传了window_size但装的 FA 版本没这参数 → TypeError。
三、根因
根因两点:
- 张量布局假设不一致:ERNIE 的注意力把 q/k/v 组织成
[B, H, S, D]丢给 FA,但 FA 期望[B, S, H, D](或反之),维度顺序错配 → shape mismatch。 - API 参数假设不一致:ERNIE 调用
flash_attn_func(..., window_size=...),但用户 FA 版本没有该参数 → 或反之(ERNIE 没传但 FA 需要)。
本质:ERNIE-Image 对 FlashAttention 的「张量布局 + API 参数」做了硬编码假设,而不同 FA 版本这两点不一致,导致兼容性崩。
四、最小可运行复现
用标准库复现「布局假设不一致导致维度错配」:
import torch def flash_attn_func_expected(q, k, v): # 假设 FA 期望 [B, S, H, D] assert q.dim() == 4 and q.shape[1] == q.shape[2] * 0 + q.shape[1] # S 在 dim1 return q # 简化 def ernie_attn_call(q, k, v): # ERNIE 按 [B, H, S, D] 组织,直接丢给 FA return flash_attn_func_expected(q, k, v) # ERNIE 的 q: [B, H, S, D] = [1, 8, 4096, 64] q = torch.randn(1, 8, 4096, 64) try: ernie_attn_call(q, q, q) except AssertionError as e: print("AssertionError:", e) # FA 期望 [B,S,H,D],ERNIE 给 [B,H,S,D] 顺序错复现「参数不匹配」:flash_attn_func(q,k,v,window_size=...)在没该参数的版本上TypeError: unexpected keyword。
五、解决方案(第一层:最小直接修复)
最小修复:在 ERNIE 的注意力里做一个「布局适配 + 参数适配」的包装,按检测到的 FA 版本调整:
import torch from transformers.utils import is_flash_attn_2_available def ernie_flash_attention(q, k, v, num_heads, window_size=None): # q/k/v: ERNIE 内部 [B, H, S, D],转成 FA 期望 [B, S, H, D] def to_fa_layout(t): # [B, H, S, D] -> [B, S, H, D] return t.transpose(1, 2).contiguous() q_fa = to_fa_layout(q); k_fa = to_fa_layout(k); v_fa = to_fa_layout(v) if is_flash_attn_2_available(): from flash_attn import flash_attn_func import inspect # 参数适配:只有 FA 支持 window_size 才传 kwargs = {} if "window_size" in inspect.signature(flash_attn_func).parameters and window_size is not None: kwargs["window_size"] = window_size out = flash_attn_func(q_fa, k_fa, v_fa, **kwargs) else: # 回退 SDPA out = torch.nn.functional.scaled_dot_product_attention(q_fa, k_fa, v_fa) # 转回 ERNIE 的 [B, H, S, D] return out.transpose(1, 2).contiguous()这样无论 FA 版本布局/参数如何,都先归一化到[B,S,H,D]并只传 FA 支持的参数,兼容崩溃消失。
六、解决方案(第二层:结构性改进)
把「ERNIE-Image 注意力对 FlashAttention 的布局/参数兼容规则」收敛成一个 dataclass 单一真源:
from dataclasses import dataclass, field from typing import Dict, List, Tuple @dataclass(frozen=True) class ErnieFlashAttnPolicy: """ERNIE-Image 与 FlashAttention 兼容的单一真源。""" # ERNIE 内部布局 -> FA 期望布局 ernie_layout: str = "B H S D" fa_layout: str = "B S H D" # 需要转置的维度对(ERNIE 的 dim1<->dim2) transpose_dims: Tuple[int, int] = (1, 2) # FA 各版本支持的参数(用于参数适配) supported_kwargs_by_version: Dict[str, Tuple[str, ...]] = field(default_factory=lambda: { "2.0": ("causal", "window_size", "softmax_scale"), "1.0": ("causal",), }) # 不支持时回退的后端 fallback: str = "scaled_dot_product_attention" def to_fa_layout(self, t: torch.Tensor) -> torch.Tensor: d1, d2 = self.transpose_dims return t.transpose(d1, d2).contiguous() def filter_kwargs(self, fa_version: str, **kwargs): allowed = self.supported_kwargs_by_version.get(fa_version, ()) return {k: v for k, v in kwargs.items() if k in allowed} def detect_version(self) -> str: try: import flash_attn return getattr(flash_attn, "__version__", "2.0")[:3] except Exception: return "0.0" # 回退注意力调用统一走policy.to_fa_layout+policy.filter_kwargs,版本探测失败自动回退 SDPA。
七、解决方案(第三层:断言 / CI 守护)
用 pytest 把「布局适配 + 参数适配 + 回退」固化成回归(可在多 FA 版本矩阵跑):
import torch import pytest from mylib.ernie_flashattn import ErnieFlashAttnPolicy POLICY = ErnieFlashAttnPolicy() def test_layout_transpose(): q = torch.randn(1, 8, 4096, 64) # [B,H,S,D] q_fa = POLICY.to_fa_layout(q) assert q_fa.shape == (1, 4096, 8, 64) # -> [B,S,H,D] def test_kwargs_filtered_by_version(): # FA 1.0 不支持 window_size kw = POLICY.filter_kwargs("1.0", causal=True, window_size=(0, 0)) assert "window_size" not in kw and "causal" in kw # FA 2.0 支持 kw2 = POLICY.filter_kwargs("2.0", causal=True, window_size=(0, 0)) assert "window_size" in kw2 def test_fa_vs_sdpa_same_shape(): q = torch.randn(1, 8, 16, 64); k = q; v = q out_fa = ernie_flash_attention(q, k, v, num_heads=8) out_sdpa = torch.nn.functional.scaled_dot_product_attention( POLICY.to_fa_layout(q), POLICY.to_fa_layout(k), POLICY.to_fa_layout(v)) assert out_fa.shape[1:] == out_sdpa.shape[1:] def test_no_crash_under_any_fa_version(): # 用 fake FA 版本探测,确认不传不支持的参数 for ver in ("1.0", "2.0"): kw = POLICY.filter_kwargs(ver, window_size=(0, 0)) # 不论版本都不应引发 unexpected keyword assert isinstance(kw, dict)CI 把test_layout_transpose与test_kwargs_filtered_by_version作为 ERNIE-Image + FlashAttention 的必过项,要求「布局归一化 + 参数按版本过滤」。
八、排查清单
ERNIE-Image + FlashAttention 不兼容按顺序查:
- 报 q/k/v 维度不匹配(S 和 H 位置反)?ERNIE 用
[B,H,S,D]、FA 期望[B,S,H,D],用to_fa_layout转置。 - 报
unexpected keyword argument 'window_size'?ERNIE 传了 FA 版本不支持的参数,用filter_kwargs按版本过滤。 - 开了 FA2 后图和不打开不一样(更糊)?布局/参数静默错配,数值跑偏,必须归一化。
- 是否探测 FA 版本再决定参数?没探测就硬编码参数必崩。
- 不支持时是否回退 SDPA?
scaled_dot_product_attention是稳妥兜底。 - dtype 是否匹配?FA2 通常要求 fp16/bf16,fp32 可能报错,需统一。
九、小结
「Incompatibility between FlashAttention and ERNIE-Image」本质是ERNIE-Image 的注意力对 FlashAttention 的「张量布局([B,H,S,D] vs [B,S,H,D])+ API 参数(window_size 等)」做了硬编码假设,而不同 FA 版本这两点不一致,导致维度错配、参数报错或静默数值错。第一层加「布局转置 + 按版本过滤参数 + 回退 SDPA」的适配包装;第二层把布局/参数兼容规则收敛到ErnieFlashAttnPolicy单一真源,版本探测失败自动 SDPA;第三层用 pytest 守住「布局归一、参数按版本过滤、回退可用」。通用教训:**任何调用外部加速内核(FlashAttention 等)的代码,都必须把「对方期望的布局 + 该版本支持的参数」当作可变契约,做归一化与版本探测,而非赌某一个版本的实现。