[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
+1 -1
View File
@@ -148,7 +148,7 @@ jobs:
- ".github/workflows/pr-test-multimodal-gen.yml"
- "python/pyproject.toml"
- "python/sglang/multimodal_gen/**/!(*.md|*.ipynb)"
- "python/sglang/jit_kernel/diffusion/**"
- "python/sglang/jit_kernel/**"
- "python/sglang/jit_kernel/tests/diffusion/**"
- "python/sglang/jit_kernel/benchmark/diffusion/**"
- "python/sglang/cli/**"
@@ -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,