Skip to content
Draft
Show file tree
Hide file tree
Changes from 16 commits
Commits
Show all changes
34 commits
Select commit Hold shift + click to select a range
347c0ed
generalized low-latency router GEMM via cuteDSL
LopezCastroRoberto Apr 15, 2026
7b9fa1f
generalized low-latency router + A GEMM via cuteDSL (bf16/fp8)
LopezCastroRoberto Apr 17, 2026
df92ff2
add a_gemm API
LopezCastroRoberto Apr 17, 2026
ef0182b
update
LopezCastroRoberto Apr 17, 2026
6eda5a5
add TinyGEMM to the benchmark
LopezCastroRoberto Apr 22, 2026
18bc6b6
update
LopezCastroRoberto Apr 22, 2026
964dac4
add tma version
LopezCastroRoberto Apr 22, 2026
7c1546d
update benchmakrs
LopezCastroRoberto Apr 22, 2026
1d3a2e8
update
LopezCastroRoberto Apr 22, 2026
88f444b
update
LopezCastroRoberto Apr 22, 2026
eb7bffb
update
LopezCastroRoberto Apr 22, 2026
737eef7
update
LopezCastroRoberto Apr 22, 2026
4c68740
Merge branch 'main' into feature/ll_gemm_pdl
LopezCastroRoberto Apr 28, 2026
707d385
update
LopezCastroRoberto Apr 28, 2026
cd802fe
add a2a flashinfer
LopezCastroRoberto Apr 28, 2026
b3a512e
update
LopezCastroRoberto Apr 28, 2026
14728c8
update linear
LopezCastroRoberto May 4, 2026
b7d839b
update linear
LopezCastroRoberto May 5, 2026
efd8055
update
LopezCastroRoberto May 6, 2026
5111c65
update
LopezCastroRoberto May 7, 2026
d3d5316
update
LopezCastroRoberto May 9, 2026
4069ed1
update
LopezCastroRoberto May 9, 2026
d6ed40c
add benchmark + cleanup
LopezCastroRoberto May 11, 2026
27a7311
benchmark cleanup
LopezCastroRoberto May 11, 2026
236c314
remove pdl overlap bench
LopezCastroRoberto May 11, 2026
d60f987
improve test coverage
LopezCastroRoberto May 11, 2026
47fc229
cleanup fp8 path
LopezCastroRoberto May 11, 2026
f2e9e61
cleanup bf16 path
LopezCastroRoberto May 11, 2026
fefad72
remove space
LopezCastroRoberto May 11, 2026
7910943
simplify bf16 qkv_a_proj_impl integration+
LopezCastroRoberto May 11, 2026
49da9c2
cleanup up router
LopezCastroRoberto May 12, 2026
00426fd
rm tma ll gemm
LopezCastroRoberto May 12, 2026
29d3cd6
rm tma ll gemm
LopezCastroRoberto May 12, 2026
6d32d42
cleanup ll_a_gemm
LopezCastroRoberto May 12, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
129 changes: 129 additions & 0 deletions benchmarks/kernels/bench_ll_a_gemm.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,129 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

import torch
from triton.testing import do_bench_cudagraph

from vllm import _custom_ops as ops
from vllm.model_executor.layers.fused_moe.router.ll_a_gemm import ll_a_gemm
from vllm.model_executor.layers.fused_moe.router.ll_a_gemm_tma import ll_a_gemm_tma

q = [0.5, 0.2, 0.8]
_HAS_DSV3 = hasattr(ops, 'dsv3_fused_a_gemm')

try:
from flashinfer.gemm import tgv_gemm_sm100
from flashinfer import autotune
_HAS_TGV = True
except ImportError:
_HAS_TGV = False

try:
from flashinfer.gemm import tinygemm_bf16
_HAS_TINY = True
except ImportError:
_HAS_TINY = False

print(f'Device: {torch.cuda.get_device_name()}')
print(f'DSV3-A: {_HAS_DSV3} | TGV: {_HAS_TGV} | tinygemm: {_HAS_TINY}')
print()

SHAPES = [
(7168, 2112, "a_proj combined"),
(7168, 576, "kv_a_proj"),
(7168, 1536, "q_a_proj"),
(1536, 24576, "q_b_proj TP1"),
(1536, 3072, "q_b_proj TP8"),
(512, 32768, "kv_b_proj TP1"),
(512, 4096, "kv_b_proj TP8"),
]


def _bench(fn, q=q):
return do_bench_cudagraph(fn, rep=200, quantiles=q)[0] * 1000


def bench_one(M, K, N):
a = torch.randn(M, K, dtype=torch.bfloat16, device='cuda')
b = torch.randn(N, K, dtype=torch.bfloat16, device='cuda')
a8 = a.to(torch.float8_e4m3fn).view(torch.bfloat16)
b8 = b.to(torch.float8_e4m3fn).view(torch.bfloat16)

r = {}

# Peeled cp.async
r['p-bf16'] = _bench(lambda: ll_a_gemm(a, b))
r['p-fp8'] = _bench(lambda: ll_a_gemm(a8, b8, is_fp8=True))

# TMA pipeline
try:
ll_a_gemm_tma(a, b); torch.cuda.synchronize()

Check failure on line 60 in benchmarks/kernels/bench_ll_a_gemm.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (E702)

benchmarks/kernels/bench_ll_a_gemm.py:60:28: E702 Multiple statements on one line (semicolon)
r['t-bf16'] = _bench(lambda: ll_a_gemm_tma(a, b))
except Exception:
r['t-bf16'] = float('nan')

try:
ll_a_gemm_tma(a8, b8, is_fp8=True); torch.cuda.synchronize()

Check failure on line 66 in benchmarks/kernels/bench_ll_a_gemm.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (E702)

benchmarks/kernels/bench_ll_a_gemm.py:66:43: E702 Multiple statements on one line (semicolon)
r['t-fp8'] = _bench(lambda: ll_a_gemm_tma(a8, b8, is_fp8=True))
except Exception:
r['t-fp8'] = float('nan')

# DSV3 fused A GEMM (C++) — only K=7168, N=2112, M<=16
if _HAS_DSV3 and K == 7168 and N == 2112 and M <= 16:
o = torch.empty(M, N, dtype=torch.bfloat16, device='cuda')
r['DSV3'] = _bench(lambda: ops.dsv3_fused_a_gemm(o, a, b.T))
else:
r['DSV3'] = float('nan')

# TGV-sm100 (FlashInfer) — requires N%16==0
if _HAS_TGV and N % 16 == 0:
bias = torch.zeros(N, dtype=torch.bfloat16, device='cuda')
out = torch.empty(M, N, dtype=torch.bfloat16, device='cuda')
with autotune(True):
tgv_gemm_sm100(a, b.T, bias, out=out)
torch.cuda.synchronize()
r['TGV'] = _bench(lambda: tgv_gemm_sm100(a, b.T, bias, out=out))
else:
r['TGV'] = float('nan')

# tinygemm_bf16 (FlashInfer tinygemm2) — requires N%16==0
if _HAS_TINY and N % 16 == 0:
bias_t = torch.zeros(N, dtype=torch.bfloat16, device='cuda')
out_t = torch.empty(M, N, dtype=torch.bfloat16, device='cuda')
try:
with autotune(True):
tinygemm_bf16(a, b, out_t, bias=bias_t)
torch.cuda.synchronize()
r['tiny2'] = _bench(lambda: tinygemm_bf16(a, b, out_t, bias=bias_t))
except Exception:
r['tiny2'] = float('nan')
else:
r['tiny2'] = float('nan')

# cuBLAS
omm = torch.empty(M, N, dtype=torch.bfloat16, device='cuda')
r['cuBLAS'] = _bench(lambda: torch.mm(a, b.T, out=omm))

return r


cols = ['p-bf16', 'p-fp8', 't-bf16', 't-fp8', 'DSV3', 'TGV', 'tiny2', 'cuBLAS']

for K, N, label in SHAPES:
print(f'=== {label}: K={K}, N={N} ===')
hdr = f"{'M':>3} |" + "".join(f" {c:>8}" for c in cols)
print(hdr)
print('-' * len(hdr))

for M in [1, 4, 16]:
r = bench_one(M, K, N)
bf16_cols = ['p-bf16', 't-bf16', 'DSV3', 'TGV', 'tiny2', 'cuBLAS']
fp8_cols = ['p-fp8', 't-fp8']
best_bf16 = min((k for k in bf16_cols if r[k] == r[k]),
key=lambda k: r[k])
best_fp8 = min((k for k in fp8_cols if r[k] == r[k]),
key=lambda k: r[k])
vals = "".join(f" {r[c]:7.2f}us" if r[c] == r[c] else " N/A"
for c in cols)
print(f' {M:2d} |{vals} bf16:{best_bf16} fp8:{best_fp8}')
print()
105 changes: 105 additions & 0 deletions benchmarks/kernels/bench_ll_router_gemm.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,105 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

import argparse
import os

import torch

from vllm import _custom_ops as ops
from vllm.model_executor.layers.fused_moe.router.ll_router_gemm import (
ll_router_gemm,
)
from vllm.triton_utils import triton

_HAS_DSV3 = hasattr(ops, "dsv3_router_gemm")

_providers = ["ll-router-bf16", "ll-router-fp8", "cublas-bf16"]
if _HAS_DSV3:
_providers.append("dsv3-trtllm")


@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["M"],
x_vals=[1, 2, 4, 8, 16],
x_log=False,
line_arg="provider",
line_vals=_providers,
line_names=_providers,
ylabel="Latency (us, lower is better)",
plot_name="LL Router GEMM",
args={},
)
)
def benchmark(M, provider, N, K):
device = "cuda"
quantiles = [0.5, 0.2, 0.8]

if provider == "ll-router-bf16":
a = torch.randn(M, K, dtype=torch.bfloat16, device=device)
b = torch.randn(N, K, dtype=torch.bfloat16, device=device)
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
lambda: ll_router_gemm(a, b), quantiles=quantiles
)

elif provider == "ll-router-fp8":
a = torch.randn(M, K, device=device).to(torch.float8_e4m3fn)
b = torch.randn(N, K, device=device).to(torch.float8_e4m3fn)
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
lambda: ll_router_gemm(a, b), quantiles=quantiles
)

elif provider == "cublas-bf16":
a = torch.randn(M, K, dtype=torch.bfloat16, device=device)
b = torch.randn(N, K, dtype=torch.bfloat16, device=device)
out = torch.empty(M, N, dtype=torch.bfloat16, device=device)
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
lambda: torch.mm(a, b.T, out=out), quantiles=quantiles
)

elif provider == "dsv3-trtllm":
# DSV3 only supports N∈{256,384}, K=7168
if N not in (256, 384) or K != 7168:
return float("nan"), float("nan"), float("nan")
from vllm import _custom_ops as ops

a = torch.randn(M, K, dtype=torch.bfloat16, device=device)
b = torch.randn(N, K, dtype=torch.bfloat16, device=device)
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
lambda: ops.dsv3_router_gemm(a, b, torch.float32), quantiles=quantiles
)

# Return latency in us
return ms * 1000, min_ms * 1000, max_ms * 1000


SHAPES = [
(256, 7168, "DSV3 router"),
(256, 2048, "Small K"),
(128, 5120, "DeepSeek V2"),
(8, 4096, "Mixtral-8x7B"),
(64, 2880, "Non-aligned K"),
]


if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--save-path", type=str, default=None)
args = parser.parse_args()

print(f"Device: {torch.cuda.get_device_name()}")
print()

for N, K, desc in SHAPES:
print(f"{desc}, N={N} K={K}:")
save_dir = args.save_path or f"bench_ll_router_n{N}_k{K}"
os.makedirs(save_dir, exist_ok=True)
benchmark.run(
print_data=True,
show_plots=False,
save_path=save_dir,
N=N,
K=K,
)
print()
161 changes: 161 additions & 0 deletions benchmarks/kernels/bench_pdl_overlap.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,161 @@
import torch
from cutlass._mlir.dialects import llvm as _llvm
from cutlass.cutlass_dsl import dsl_user_op
import cutlass
import cutlass.cute as cute
from cuda.bindings.driver import CUstream
from cutlass.cute.runtime import from_dlpack
from torch.cuda import current_stream

from vllm.model_executor.layers.fused_moe.router.ll_a_gemm import ll_a_gemm
from vllm.model_executor.layers.fused_moe.router.ll_a_gemm_tma import ll_a_gemm_tma
from vllm.model_executor.layers.fused_moe.router.ll_router_gemm import ll_router_gemm


@dsl_user_op
def nanosleep(ns, *, loc=None, ip=None):
_llvm.inline_asm(res=None, operands_=[ns.ir_value(loc=loc, ip=ip)],
asm_string="nanosleep.u32 $0;", constraints="r",
has_side_effects=True, loc=loc, ip=ip)

@cute.kernel
def producer_k(gOut: cute.Tensor, tail_ns: cutlass.Int32):
tidx = cute.arch.thread_idx()[0]
v = cutlass.Float32(1.0)
v = v + cutlass.Float32(1.0)
cute.arch.griddepcontrol_launch_dependents()
nanosleep(tail_ns)
if tidx == 0:
gOut[0] = v

@cute.jit
def host_producer(gOut: cute.Tensor, tail_ns: cutlass.Int32, s: CUstream):
producer_k(gOut, tail_ns).launch(
grid=[1, 1, 1], block=[128, 1, 1], stream=s, use_pdl=True)

def bench_cg_n1(fn, n_retries=100):
with torch.cuda.stream(torch.cuda.Stream()):
fn(); torch.cuda.synchronize()

Check failure on line 38 in benchmarks/kernels/bench_pdl_overlap.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (E702)

benchmarks/kernels/bench_pdl_overlap.py:38:13: E702 Multiple statements on one line (semicolon)
g = torch.cuda.CUDAGraph()
with torch.cuda.graph(g):
fn()
torch.cuda.synchronize()
for _ in range(10):
g.replay()
torch.cuda.synchronize()
ret = []
for _ in range(n_retries):
s = torch.cuda.Event(enable_timing=True)
e = torch.cuda.Event(enable_timing=True)
s.record()
g.replay()
e.record()
torch.cuda.synchronize()
ret.append(s.elapsed_time(e) * 1000)
ret.sort()
return ret[len(ret) // 2]


# Producer
buf = torch.empty(1, dtype=torch.float32, device="cuda")
bc = from_dlpack(buf, assumed_align=16).mark_layout_dynamic()
comp_p = cute.compile(host_producer, bc, 0,
CUstream(current_stream().cuda_stream))

print("Device:", torch.cuda.get_device_name())
print("Producer: 1 block, 128 threads | n_repeat=1, n_retries=100")
print()

SHAPES = [
(7168, 256, "router gate", "router"),
(7168, 2112, "a_proj combined", "a_gemm"),
(7168, 576, "kv_a_proj", "a_gemm"),
(7168, 1536, "q_a_proj", "a_gemm"),
(1536, 3072, "q_b_proj TP8", "a_gemm"),
(512, 4096, "kv_b_proj TP8", "a_gemm"),
]

TAILS = [0, 2000, 5000, 10000, 20000, 50000]
M = 16

# Producer solo times
prod_times = {}
for tns in TAILS:
def sp(_t=tns):
comp_p(bc, _t, CUstream(current_stream().cuda_stream))
prod_times[tns] = bench_cg_n1(sp)

for K, N, label, ktype in SHAPES:
print("=" * 95)
print("M=%d K=%d N=%d — %s" % (M, K, N, label))

Check failure on line 90 in benchmarks/kernels/bench_pdl_overlap.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (UP031)

benchmarks/kernels/bench_pdl_overlap.py:90:11: UP031 Use format specifiers instead of percent format

a = torch.randn(M, K, dtype=torch.bfloat16, device="cuda")
b = torch.randn(N, K, dtype=torch.bfloat16, device="cuda")

a8 = a.to(torch.float8_e4m3fn).view(torch.bfloat16)
b8 = b.to(torch.float8_e4m3fn).view(torch.bfloat16)

kernels = {}

if ktype == "router":
ll_router_gemm(a, b); torch.cuda.synchronize()

Check failure on line 101 in benchmarks/kernels/bench_pdl_overlap.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (E702)

benchmarks/kernels/bench_pdl_overlap.py:101:29: E702 Multiple statements on one line (semicolon)
kernels['p-bf16'] = lambda: ll_router_gemm(a, b)

Check failure on line 102 in benchmarks/kernels/bench_pdl_overlap.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (B023)

benchmarks/kernels/bench_pdl_overlap.py:102:55: B023 Function definition does not bind loop variable `b`

Check failure on line 102 in benchmarks/kernels/bench_pdl_overlap.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (B023)

benchmarks/kernels/bench_pdl_overlap.py:102:52: B023 Function definition does not bind loop variable `a`
else:
ll_a_gemm(a, b); torch.cuda.synchronize()

Check failure on line 104 in benchmarks/kernels/bench_pdl_overlap.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (E702)

benchmarks/kernels/bench_pdl_overlap.py:104:24: E702 Multiple statements on one line (semicolon)
kernels['p-bf16'] = lambda: ll_a_gemm(a, b)

Check failure on line 105 in benchmarks/kernels/bench_pdl_overlap.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (B023)

benchmarks/kernels/bench_pdl_overlap.py:105:50: B023 Function definition does not bind loop variable `b`

Check failure on line 105 in benchmarks/kernels/bench_pdl_overlap.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (B023)

benchmarks/kernels/bench_pdl_overlap.py:105:47: B023 Function definition does not bind loop variable `a`

ll_a_gemm(a8, b8, is_fp8=True); torch.cuda.synchronize()
kernels['p-fp8'] = lambda: ll_a_gemm(a8, b8, is_fp8=True)

try:
ll_a_gemm_tma(a, b); torch.cuda.synchronize()
kernels['t-bf16'] = lambda: ll_a_gemm_tma(a, b)
except Exception as e:
print(" TMA bf16 error: %s" % str(e)[:60])

try:
ll_a_gemm_tma(a8, b8, is_fp8=True); torch.cuda.synchronize()
kernels['t-fp8'] = lambda: ll_a_gemm_tma(a8, b8, is_fp8=True)
except Exception as e:
print(" TMA fp8 error: %s" % str(e)[:60])

# Solo times
solos = {k: bench_cg_n1(fn) for k, fn in kernels.items()}
solo_str = " ".join("%s=%.2fus" % (k, v) for k, v in solos.items())
print(" Solo: %s" % solo_str)
print()

# Header
kcols = list(kernels.keys())
hdr = "%8s | %5s |" % ("tail", "prod")
for k in kcols:
hdr += " %8s %5s |" % (k, "ovlp")
hdr += " bf16-best fp8-best"
print(hdr)
print("-" * len(hdr))

for tns in TAILS:
pr = prod_times[tns]
pairs = {}
ovlps = {}

for k, fn in kernels.items():
def pair_fn(_t=tns, _fn=fn):
comp_p(bc, _t, CUstream(current_stream().cuda_stream))
_fn()
pairs[k] = bench_cg_n1(pair_fn)
ovlps[k] = pr + solos[k] - pairs[k]

# Winners by dtype
bf16_keys = [k for k in kcols if 'bf16' in k]
fp8_keys = [k for k in kcols if 'fp8' in k]
best_bf16 = min(bf16_keys, key=lambda k: pairs[k]) if bf16_keys else ""
best_fp8 = min(fp8_keys, key=lambda k: pairs[k]) if fp8_keys else ""

row = "%6dns | %4.1f |" % (tns, pr)
for k in kcols:
row += " %7.2fus %4.1f |" % (pairs[k], ovlps[k])
row += " %-9s %s" % (best_bf16, best_fp8)
print(row)

print()
Loading
Loading