[diffusion] fix: fix FA3 varlen out argument handling (#24688)

This commit is contained in:
Mick
2026-05-08 19:01:49 +08:00
committed by GitHub
parent 17888fa92a
commit 73b8eda103
2 changed files with 7 additions and 6 deletions
@@ -16,15 +16,15 @@ SGL_FA3_KERNEL_REVISION = "v1"
DEFAULT_FA3_KERNEL_LOCKFILE = "kernels.lock"
def _call_fa3_kernel(kernel, *args, out=None):
def _call_fa3_kernel(kernel, *args, out=None, **kwargs):
if out is None:
return kernel(*args)
return kernel(*args, **kwargs)
try:
return kernel(*args, out=out)
return kernel(*args, **kwargs, out=out)
except TypeError as exc:
if "unexpected keyword argument 'out'" not in str(exc):
raise
return kernel(*args)
return kernel(*args, **kwargs)
@cache_once
@@ -212,7 +212,8 @@ def flash_attn_varlen_func(
"flash_attn at sgl-kernel is only supported on sm90 and above"
)
return _load_fa3_kernels()["flash_attn_varlen_func"](
return _call_fa3_kernel(
_load_fa3_kernels()["flash_attn_varlen_func"],
q=q,
k=k,
v=v,