sparse/config: default ATTN to flex_gemm_sparse_attn on Darwin when flex_gemm probe passes
The flex_gemm_sparse_attn backend is now flash-attention-v2 with simdgroup_matrix_multiply_accumulate for Q@K^T and P@V plus simd-shuffle softmax row reductions, and wins 5–15× over SDPA-padded-CPU-bounce at every measured shape including max_seqlen=2048. Production decoder block on M3 Max: 5.04× wall-clock vs all-SDPA-padded baseline (3.49 ms vs 17.57 ms), entirely from this change being the default. The existing __flex_gemm_works_on_mps() probe covers the same package the attention path lives in, so the gate is identical: if conv probe passes, set both CONV='flex_gemm' and ATTN='flex_gemm_sparse_attn'. SPARSE_ATTN_BACKEND= (or ATTN_BACKEND=) env override is unchanged. Also fix benchmarks/e2e_decoder.py to inject the repo root into sys.path so it runs as `python benchmarks/e2e_decoder.py` from any cwd.
This commit is contained in:
parent
cf886fd00e
commit
f3177cd807
@ -15,6 +15,7 @@ Run:
|
|||||||
python benchmarks/e2e_decoder.py
|
python benchmarks/e2e_decoder.py
|
||||||
"""
|
"""
|
||||||
import os
|
import os
|
||||||
|
import sys
|
||||||
import time
|
import time
|
||||||
import math
|
import math
|
||||||
|
|
||||||
@ -22,6 +23,11 @@ os.environ.setdefault("SPARSE_CONV_BACKEND", "flex_gemm")
|
|||||||
os.environ.setdefault("SPARSE_ATTN_BACKEND", "flex_gemm_sparse_attn")
|
os.environ.setdefault("SPARSE_ATTN_BACKEND", "flex_gemm_sparse_attn")
|
||||||
os.environ.setdefault("FLEX_GEMM_QUIET", "1")
|
os.environ.setdefault("FLEX_GEMM_QUIET", "1")
|
||||||
|
|
||||||
|
# trellis2 is in-tree (no pip install) — make the bench runnable as
|
||||||
|
# `python benchmarks/e2e_decoder.py` regardless of cwd by putting the
|
||||||
|
# repo root on sys.path before any trellis2 import.
|
||||||
|
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
assert torch.backends.mps.is_available(), "This benchmark needs MPS."
|
assert torch.backends.mps.is_available(), "This benchmark needs MPS."
|
||||||
|
|||||||
@ -10,11 +10,17 @@ def __detect_defaults():
|
|||||||
"""Auto-detect best backends for current platform."""
|
"""Auto-detect best backends for current platform."""
|
||||||
global CONV, ATTN
|
global CONV, ATTN
|
||||||
if platform.system() == 'Darwin':
|
if platform.system() == 'Darwin':
|
||||||
ATTN = 'sdpa'
|
|
||||||
if __flex_gemm_works_on_mps():
|
if __flex_gemm_works_on_mps():
|
||||||
CONV = 'flex_gemm'
|
CONV = 'flex_gemm'
|
||||||
|
# flex_gemm_sparse_attn is now flash-attention-v2 with simdgroup
|
||||||
|
# matmul + simd-shuffle softmax (5–15× over SDPA-padded-CPU-bounce
|
||||||
|
# at every measured shape including max_seqlen=2048). The conv
|
||||||
|
# probe above covers the same package, so if it passed, the
|
||||||
|
# attention path is available too.
|
||||||
|
ATTN = 'flex_gemm_sparse_attn'
|
||||||
else:
|
else:
|
||||||
CONV = 'pytorch'
|
CONV = 'pytorch'
|
||||||
|
ATTN = 'sdpa'
|
||||||
elif not __has_cuda():
|
elif not __has_cuda():
|
||||||
CONV = 'pytorch'
|
CONV = 'pytorch'
|
||||||
ATTN = 'sdpa'
|
ATTN = 'sdpa'
|
||||||
@ -97,6 +103,6 @@ def set_debug(debug: bool):
|
|||||||
global DEBUG
|
global DEBUG
|
||||||
DEBUG = debug
|
DEBUG = debug
|
||||||
|
|
||||||
def set_attn_backend(backend: Literal['xformers', 'flash_attn', 'sdpa']):
|
def set_attn_backend(backend: Literal['xformers', 'flash_attn', 'flash_attn_3', 'sdpa', 'flex_gemm_sparse_attn']):
|
||||||
global ATTN
|
global ATTN
|
||||||
ATTN = backend
|
ATTN = backend
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user