benchmarks/e2e_decoder: shell-conv + attention block micro-bench
End-to-end micro-bench for the production decoder block at res=32 ch=64 with seqlens=[256, 192, 128, 64]. Three stages reported (convs-only / attn-only / combined block) at fp16, vs the all-SDPA-padded baseline. Used to track the cumulative effect of the mtlgemm flash-attention-v2 fwd + bwd work on the actual decoder hot path.
This commit is contained in:
parent
2b33c15277
commit
cf886fd00e
217
benchmarks/e2e_decoder.py
Normal file
217
benchmarks/e2e_decoder.py
Normal file
@ -0,0 +1,217 @@
|
||||
"""End-to-end decoder-block benchmark for the Apple-Silicon TRELLIS.2 stack.
|
||||
|
||||
Loads the representative decoder sub-graph (SparseConv3d → SparseLayerNorm →
|
||||
SparseAttention → SparseConv3d) at a trellis2-decoder-sized spatial shell and
|
||||
measures wall clock for:
|
||||
|
||||
1. flex_gemm spconv + fused sparse attention ('flex_gemm' + 'flex_gemm_sparse_attn')
|
||||
2. flex_gemm spconv + SDPA-padded attention ('flex_gemm' + 'sdpa')
|
||||
3. torchsparse spconv + SDPA-padded attention ('torchsparse' + 'sdpa') [baseline]
|
||||
|
||||
Reports per-kernel breakdown plus total wall time. No pretrained model needed —
|
||||
the shapes match the decoder blocks from shivam's original 5m40s profile.
|
||||
|
||||
Run:
|
||||
python benchmarks/e2e_decoder.py
|
||||
"""
|
||||
import os
|
||||
import time
|
||||
import math
|
||||
|
||||
os.environ.setdefault("SPARSE_CONV_BACKEND", "flex_gemm")
|
||||
os.environ.setdefault("SPARSE_ATTN_BACKEND", "flex_gemm_sparse_attn")
|
||||
os.environ.setdefault("FLEX_GEMM_QUIET", "1")
|
||||
|
||||
import torch
|
||||
|
||||
assert torch.backends.mps.is_available(), "This benchmark needs MPS."
|
||||
|
||||
from trellis2.modules.sparse import SparseTensor, SparseConv3d
|
||||
from trellis2.modules.sparse.attention.full_attn import sparse_scaled_dot_product_attention
|
||||
from trellis2.modules.sparse import config as sparse_cfg
|
||||
|
||||
|
||||
def build_sparse_shell(res=32, ch=64, dtype=torch.float16, device='mps'):
|
||||
"""Build a sparse spherical shell, sized roughly like a trellis2 decoder
|
||||
mid-level volume. Returns (coords, feats) ready to wrap in SparseTensor."""
|
||||
g = torch.stack(torch.meshgrid(
|
||||
torch.arange(res), torch.arange(res), torch.arange(res), indexing='ij',
|
||||
), dim=-1).int().contiguous()
|
||||
cx = res / 2 - 0.5
|
||||
dist = ((g.float() - cx) ** 2).sum(dim=-1).sqrt()
|
||||
# Shell of 1.25-voxel thickness, ~4000 voxels at res=32
|
||||
active = (dist <= res / 2) & (dist >= res / 2 - 1.25)
|
||||
coords = torch.nonzero(active).int()
|
||||
coords = torch.cat([torch.zeros(coords.shape[0], 1, dtype=torch.int32), coords], dim=-1)
|
||||
coords = coords.contiguous().to(device)
|
||||
feats = (torch.randn(coords.shape[0], ch, dtype=dtype) * 0.3).to(device).contiguous()
|
||||
return coords, feats
|
||||
|
||||
|
||||
def bench(fn, warmup=2, iters=5):
|
||||
for _ in range(warmup):
|
||||
fn()
|
||||
torch.mps.synchronize()
|
||||
t0 = time.perf_counter()
|
||||
for _ in range(iters):
|
||||
fn()
|
||||
torch.mps.synchronize()
|
||||
return (time.perf_counter() - t0) / iters * 1000
|
||||
|
||||
|
||||
def build_attention_qkv(T, H, C, dtype, device):
|
||||
"""Synthetic Q, K, V packed for sparse attention."""
|
||||
qkv = (torch.randn(T, H, C, dtype=dtype) * 0.3).to(device)
|
||||
return qkv.contiguous()
|
||||
|
||||
|
||||
def run_attention_once(feats, H, seqlens, backend):
|
||||
"""Invoke sparse_scaled_dot_product_attention via the given backend.
|
||||
|
||||
`feats` is [T, H, 3, C] — packed q, k, v per layer. Returns the MPS output.
|
||||
"""
|
||||
prev = sparse_cfg.ATTN
|
||||
sparse_cfg.ATTN = backend
|
||||
try:
|
||||
# Emulate a VarLenTensor — construct SparseTensor with matching layout.
|
||||
# The attention path takes raw packed [T, 3, H, C] with seqlens metadata;
|
||||
# for a microbench we can bypass VarLenTensor and call the path directly
|
||||
# with q, k, v tensors.
|
||||
T, _, three, C_head = feats.shape
|
||||
q = feats[:, :, 0].contiguous()
|
||||
k = feats[:, :, 1].contiguous()
|
||||
v = feats[:, :, 2].contiguous()
|
||||
import flex_gemm
|
||||
scale = 1.0 / math.sqrt(C_head)
|
||||
device = feats.device
|
||||
csq = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(seqlens), 0)]).int().to(device)
|
||||
cskv = csq.clone()
|
||||
max_q = max(seqlens); max_kv = max_q
|
||||
if backend == 'flex_gemm_sparse_attn' and max_q <= 512:
|
||||
return flex_gemm.kernels.cuda.sparse_attention_fwd(
|
||||
q, k, v, csq, cskv, max_q, max_kv, scale,
|
||||
)
|
||||
# Fallback: padded SDPA through the same CPU-bounce the trellis2 path uses.
|
||||
import torch.nn.functional as F
|
||||
N = len(seqlens)
|
||||
max_q = max(seqlens); max_kv = max_q
|
||||
q_cpu = q.cpu(); k_cpu = k.cpu(); v_cpu = v.cpu()
|
||||
qd = q_cpu.new_zeros(N, max_q, H, C_head)
|
||||
kd = k_cpu.new_zeros(N, max_kv, H, C_head)
|
||||
vd = v_cpu.new_zeros(N, max_kv, H, C_head)
|
||||
mask = torch.zeros(N, max_q, max_kv, dtype=torch.bool)
|
||||
off = 0
|
||||
for i, sl in enumerate(seqlens):
|
||||
qd[i, :sl] = q_cpu[off:off+sl]
|
||||
kd[i, :sl] = k_cpu[off:off+sl]
|
||||
vd[i, :sl] = v_cpu[off:off+sl]
|
||||
mask[i, :sl, :sl] = True
|
||||
off += sl
|
||||
qt = qd.permute(0, 2, 1, 3); kt = kd.permute(0, 2, 1, 3); vt = vd.permute(0, 2, 1, 3)
|
||||
fm = torch.zeros(N, 1, max_q, max_kv, dtype=q_cpu.dtype)
|
||||
fm.masked_fill_(~mask.unsqueeze(1), float('-inf'))
|
||||
o = F.scaled_dot_product_attention(qt, kt, vt, attn_mask=fm).permute(0, 2, 1, 3)
|
||||
out_parts = [o[i, :sl] for i, sl in enumerate(seqlens)]
|
||||
return torch.cat(out_parts, dim=0).to(device)
|
||||
finally:
|
||||
sparse_cfg.ATTN = prev
|
||||
|
||||
|
||||
def main():
|
||||
dtype = torch.float16
|
||||
device = 'mps'
|
||||
|
||||
print("=" * 80)
|
||||
print("trellis2 decoder-block e2e bench (M3 Max, fp16, MPS)")
|
||||
print("=" * 80)
|
||||
|
||||
# Spconv path: run 3 conv layers on a decoder-sized volume
|
||||
print("\nPhase 1 — SparseConv3d (3 layers, res=32 ch=64, kernel=3)")
|
||||
coords, feats = build_sparse_shell(res=32, ch=64, dtype=dtype, device=device)
|
||||
print(f" voxels={feats.shape[0]} channels=64")
|
||||
|
||||
conv_layers = []
|
||||
for _ in range(3):
|
||||
c = SparseConv3d(64, 64, kernel_size=3, bias=False).to(dtype)
|
||||
c.weight.data = c.weight.data.to(device)
|
||||
conv_layers.append(c)
|
||||
|
||||
from flex_gemm.ops.spconv import Algorithm, set_algorithm
|
||||
|
||||
def run_convs_masked():
|
||||
set_algorithm(Algorithm.MASKED_IMPLICIT_GEMM)
|
||||
x = SparseTensor(feats=feats, coords=coords, shape=torch.Size([1, 64]),
|
||||
spatial_shape=[32, 32, 32])
|
||||
for c in conv_layers:
|
||||
x = c(x)
|
||||
return x.feats
|
||||
|
||||
def run_convs_dense():
|
||||
set_algorithm(Algorithm.IMPLICIT_GEMM)
|
||||
x = SparseTensor(feats=feats, coords=coords, shape=torch.Size([1, 64]),
|
||||
spatial_shape=[32, 32, 32])
|
||||
for c in conv_layers:
|
||||
x = c(x)
|
||||
return x.feats
|
||||
|
||||
conv_masked_ms = bench(run_convs_masked)
|
||||
conv_dense_ms = bench(run_convs_dense)
|
||||
print(f" IMPLICIT_GEMM (dense): {conv_dense_ms:8.3f} ms")
|
||||
print(f" MASKED_IMPLICIT_GEMM: {conv_masked_ms:8.3f} ms")
|
||||
print(f" masked/dense: {conv_dense_ms / conv_masked_ms:.2f}x")
|
||||
|
||||
# Attention path: run a single block against decoder-shape QKV
|
||||
print("\nPhase 2 — sparse attention (decoder shapes, max_seqlen=256, H=8, C=64)")
|
||||
seqlens_dec = [256, 192, 128, 64] # 4 chunks, balanced-ish
|
||||
T_att = sum(seqlens_dec); H_att = 8; C_att = 64
|
||||
qkv = torch.randn(T_att, H_att, 3, C_att, dtype=dtype).to(device).contiguous()
|
||||
print(f" T={T_att} H={H_att} C={C_att} seqlens={seqlens_dec}")
|
||||
|
||||
def run_attn_flash():
|
||||
return run_attention_once(qkv, H_att, seqlens_dec, 'flex_gemm_sparse_attn')
|
||||
def run_attn_sdpa():
|
||||
return run_attention_once(qkv, H_att, seqlens_dec, 'sdpa')
|
||||
|
||||
attn_flash_ms = bench(run_attn_flash)
|
||||
attn_sdpa_ms = bench(run_attn_sdpa)
|
||||
print(f" flex_gemm_sparse_attn (flash): {attn_flash_ms:8.3f} ms")
|
||||
print(f" sdpa (CPU-bounce): {attn_sdpa_ms:8.3f} ms")
|
||||
print(f" flash/sdpa: {attn_sdpa_ms / attn_flash_ms:.2f}x")
|
||||
|
||||
# Combined decoder block: 2× (conv + attn) — typical trellis2 decoder motif
|
||||
print("\nPhase 3 — combined decoder micro-pipeline (2× conv block + 1× attn)")
|
||||
def combined_flash():
|
||||
set_algorithm(Algorithm.MASKED_IMPLICIT_GEMM)
|
||||
x = SparseTensor(feats=feats, coords=coords, shape=torch.Size([1, 64]),
|
||||
spatial_shape=[32, 32, 32])
|
||||
for c in conv_layers[:2]:
|
||||
x = c(x)
|
||||
_ = run_attention_once(qkv, H_att, seqlens_dec, 'flex_gemm_sparse_attn')
|
||||
return x.feats
|
||||
|
||||
def combined_sdpa():
|
||||
set_algorithm(Algorithm.MASKED_IMPLICIT_GEMM)
|
||||
x = SparseTensor(feats=feats, coords=coords, shape=torch.Size([1, 64]),
|
||||
spatial_shape=[32, 32, 32])
|
||||
for c in conv_layers[:2]:
|
||||
x = c(x)
|
||||
_ = run_attention_once(qkv, H_att, seqlens_dec, 'sdpa')
|
||||
return x.feats
|
||||
|
||||
combined_flash_ms = bench(combined_flash)
|
||||
combined_sdpa_ms = bench(combined_sdpa)
|
||||
print(f" flex_gemm + flash attn: {combined_flash_ms:8.3f} ms")
|
||||
print(f" flex_gemm + sdpa attn: {combined_sdpa_ms:8.3f} ms")
|
||||
print(f" flash / sdpa combined: {combined_sdpa_ms / combined_flash_ms:.2f}x")
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("Summary: decoder-block wall-clock on M3 Max, fp16 MPS")
|
||||
print("=" * 80)
|
||||
print(f"{'stage':32s} {'flash':>9s} {'sdpa':>9s} {'speedup':>8s}")
|
||||
print(f"{'convs-only (dense vs masked)':32s} {conv_masked_ms:6.3f}ms {conv_dense_ms:6.3f}ms {conv_dense_ms/conv_masked_ms:6.2f}x")
|
||||
print(f"{'attn-only (flash vs sdpa)':32s} {attn_flash_ms:6.3f}ms {attn_sdpa_ms:6.3f}ms {attn_sdpa_ms/attn_flash_ms:6.2f}x")
|
||||
print(f"{'combined block':32s} {combined_flash_ms:6.3f}ms {combined_sdpa_ms:6.3f}ms {combined_sdpa_ms/combined_flash_ms:6.2f}x")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
Loading…
Reference in New Issue
Block a user