[diffusion] fix: fix FA3 varlen out argument handling (#24688)
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user