From c869eb995a741dd8c486d5b1f434be8ad48d4ade Mon Sep 17 00:00:00 2001 From: Pedro Augusto <16762743+pedronaugusto@users.noreply.github.com> Date: Tue, 21 Apr 2026 16:01:32 +0100 Subject: [PATCH] sparse/attention: add 'flex_gemm_sparse_attn' backend option MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Wires the new mtlgemm fused sparse attention kernel through the ATTN backend selector. Dispatches to the fused Metal kernel when max(max_q, max_kv) <= 256 (where the naive per-thread-serial-KV kernel beats SDPA-padded on M3 Max), and falls through to an inline SDPA-padded path for larger max_seqlen. Opt in via ATTN_BACKEND=flex_gemm_sparse_attn or SPARSE_ATTN_BACKEND=flex_gemm_sparse_attn. Default on Darwin stays 'sdpa' — the threshold-based fallback doesn't yet prove a universal win across the pipeline's attention shape distribution. Co-Authored-By: Claude Opus 4.7 (1M context) --- .../modules/sparse/attention/full_attn.py | 51 +++++++++++++++++++ trellis2/modules/sparse/config.py | 4 +- 2 files changed, 54 insertions(+), 1 deletion(-) diff --git a/trellis2/modules/sparse/attention/full_attn.py b/trellis2/modules/sparse/attention/full_attn.py index a44afa5..dceceeb 100644 --- a/trellis2/modules/sparse/attention/full_attn.py +++ b/trellis2/modules/sparse/attention/full_attn.py @@ -211,6 +211,57 @@ def sparse_scaled_dot_product_attention(*args, **kwargs): max_q_seqlen = max(q_seqlen) max_kv_seqlen = max(kv_seqlen) out = flash_attn_3.flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_kv, max_q_seqlen, max_kv_seqlen) + elif config.ATTN == 'flex_gemm_sparse_attn': + # Fused variable-length sparse attention (mtlgemm Metal kernel). + # Keeps everything on MPS — avoids the MPS->CPU->MPS round-trip that + # SDPA-padded forces on PyTorch builds where new_zeros for MPS + # fp16/fp32 is broken. Wins measurably when max_seqlen is small-ish + # (trellis2 decoder's typical case); at larger max_seqlen the current + # naive per-thread-serial-KV kernel is slower than Accelerate-backed + # SDPA, so we fall back. Threshold is conservative; measured + # crossover was around 256-512 on M3 Max. + FUSED_ATTN_MAX_SEQLEN = 256 + if num_all_args == 1: + q, k, v = qkv.unbind(dim=1) + elif num_all_args == 2: + k, v = kv.unbind(dim=1) + if max(max(q_seqlen), max(kv_seqlen)) <= FUSED_ATTN_MAX_SEQLEN: + import flex_gemm, math + scale = 1.0 / math.sqrt(q.shape[-1]) + csq = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(q_seqlen), 0)]).int().to(device) + cskv = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(kv_seqlen), 0)]).int().to(device) + out = flex_gemm.kernels.cuda.sparse_attention_fwd( + q.contiguous(), k.contiguous(), v.contiguous(), + csq, cskv, max(q_seqlen), max(kv_seqlen), scale, + ) + else: + # Fall through to the SDPA implementation below by mirroring its + # preamble here (we already have q/k/v unbound). + import torch.nn.functional as F_attn + N_b = len(q_seqlen) + max_q = max(q_seqlen); max_kv = max(kv_seqlen) + H_b = q.shape[-2]; C_q_b = q.shape[-1]; C_v_b = v.shape[-1] + q_dense = q.new_zeros(N_b, max_q, H_b, C_q_b) + k_dense = k.new_zeros(N_b, max_kv, H_b, C_q_b) + v_dense = v.new_zeros(N_b, max_kv, H_b, C_v_b) + attn_mask = torch.zeros(N_b, max_q, max_kv, dtype=torch.bool, device=device) + q_off = 0; kv_off = 0 + for i in range(N_b): + ql = q_seqlen[i]; kvl = kv_seqlen[i] + q_dense[i, :ql] = q[q_off:q_off + ql] + k_dense[i, :kvl] = k[kv_off:kv_off + kvl] + v_dense[i, :kvl] = v[kv_off:kv_off + kvl] + attn_mask[i, :ql, :kvl] = True + q_off += ql; kv_off += kvl + q_dense = q_dense.permute(0, 2, 1, 3) + k_dense = k_dense.permute(0, 2, 1, 3) + v_dense = v_dense.permute(0, 2, 1, 3) + float_mask = torch.zeros(N_b, 1, max_q, max_kv, dtype=q_dense.dtype, device=device) + float_mask.masked_fill_(~attn_mask.unsqueeze(1), float('-inf')) + out_dense = F_attn.scaled_dot_product_attention(q_dense, k_dense, v_dense, attn_mask=float_mask) + out_dense = out_dense.permute(0, 2, 1, 3) + out_parts = [out_dense[i, :q_seqlen[i]] for i in range(N_b)] + out = torch.cat(out_parts, dim=0) elif config.ATTN == 'sdpa': import torch.nn.functional as F_attn if num_all_args == 1: diff --git a/trellis2/modules/sparse/config.py b/trellis2/modules/sparse/config.py index b2f1cc2..fc3c18d 100644 --- a/trellis2/modules/sparse/config.py +++ b/trellis2/modules/sparse/config.py @@ -78,7 +78,9 @@ def __from_env(): CONV = env_sparse_conv_backend if env_sparse_debug is not None: DEBUG = env_sparse_debug == '1' - if env_sparse_attn_backend is not None and env_sparse_attn_backend in ['xformers', 'flash_attn', 'flash_attn_3', 'sdpa']: + if env_sparse_attn_backend is not None and env_sparse_attn_backend in [ + 'xformers', 'flash_attn', 'flash_attn_3', 'sdpa', 'flex_gemm_sparse_attn', + ]: ATTN = env_sparse_attn_backend print(f"[SPARSE] Conv backend: {CONV}; Attention backend: {ATTN}")