Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
37 commits
Select commit Hold shift + click to select a range
d75cab8
Add Qwen3.5 decode kernel scaffolding
Jun 5, 2026
7830fad
Wire Qwen3.5 decode runtime and tests
Jun 5, 2026
de00a8b
Fix Qwen3.5 decode extension build
XiaomingFun233 Jun 5, 2026
8c3cb63
Align Qwen3.5 decode with FLA semantics
XiaomingFun233 Jun 6, 2026
8c40083
Fuse Qwen3.5 layout into scalar KDA decode
XiaomingFun233 Jun 10, 2026
c475de5
test qwen35 fused layout kda decode
XiaomingFun233 Jun 10, 2026
5d1b495
benchmark qwen35 decode paths
XiaomingFun233 Jun 10, 2026
0a2579b
Add Qwen3.5 prefill CUDA path
XiaomingFun233 Jun 10, 2026
3f62c26
Add Qwen3.5 fused prefill benchmark path
XiaomingFun233 Jun 11, 2026
9c26912
Support Qwen3.5 local TP head configs
XiaomingFun233 Jun 11, 2026
a7d4733
Tune Qwen3.5 TP decode policies
XiaomingFun233 Jun 12, 2026
de9d1f9
Merge branch 'inclusionAI:main' into main
XiaomingFun233 Jun 12, 2026
b38e9a2
Optimize qwen35 decode mainloop tiling
XiaomingFun233 Jun 12, 2026
68e8fa0
Revert "Optimize qwen35 decode mainloop tiling"
XiaomingFun233 Jun 12, 2026
a9c2fa8
Optimize qwen35 long decode kernel
XiaomingFun233 Jun 13, 2026
7d60b3b
Pipeline qwen35 long decode update
XiaomingFun233 Jun 13, 2026
e99eda5
Optimize qwen35 long decode ILP
XiaomingFun233 Jun 15, 2026
a3c313b
Optimize qwen35 long decode V32 staging
XiaomingFun233 Jul 11, 2026
3bd7c80
Optimize qwen35 long decode state staging
XiaomingFun233 Jul 11, 2026
0c2450c
feat(kda): support native GVA prefill benchmarks
XiaomingFun233 Jul 12, 2026
aa3d933
perf(qwen35): optimize native GVA fused prefill
XiaomingFun233 Aug 2, 2026
3bdb09a
test(qwen35): cover native GVA fused prefill
XiaomingFun233 Aug 2, 2026
cf382eb
bench(qwen35): compare actual GDN prefill paths
XiaomingFun233 Aug 2, 2026
020172b
bench(qwen35): compare native GVA decode inputs
XiaomingFun233 Aug 2, 2026
39e2602
bench(qwen35): compare SGLang packed decode
XiaomingFun233 Aug 2, 2026
f7ee00a
perf(qwen35): overlap long decode state load
XiaomingFun233 Aug 2, 2026
6f8daff
Merge upstream main into Qwen3.5 tuning branch
XiaomingFun233 Aug 2, 2026
82b6cd6
perf(qwen35): optimize native-GVA scalar prefill on H200
XiaomingFun233 Aug 3, 2026
875537a
bench(qwen35): make scalar prefill results auditable
XiaomingFun233 Aug 3, 2026
996fbbb
bench(qwen35): add scalar prefill NCU target
XiaomingFun233 Aug 3, 2026
ad23c1b
fix(kda): preserve SM90 context-parallel launcher ABI
XiaomingFun233 Aug 3, 2026
50898ab
fix(qwen35): add scalar prefill core parameter definitions
XiaomingFun233 Aug 3, 2026
db5aaea
feat(qwen35): expose scalar prefill core ABI
XiaomingFun233 Aug 3, 2026
a08651b
feat(qwen35): connect native-GVA scalar prefill adapter
XiaomingFun233 Aug 3, 2026
dc327a3
test(qwen35): import scalar prefill core adapter
XiaomingFun233 Aug 3, 2026
83c199a
feat(qwen35): add SM100 scalar prefill state-output kernels
XiaomingFun233 Aug 3, 2026
51d6ab8
perf(qwen35): specialize compact scalar gates on SM100
XiaomingFun233 Aug 3, 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
34 changes: 26 additions & 8 deletions benchmarks/bench_kda_fused_fwd.py
Original file line number Diff line number Diff line change
Expand Up @@ -176,13 +176,22 @@ def bench_fixed(configs):
cu_seqlens=cu_seqlens,
lower_bound=lower_bound,
)
common_cula = dict(common)
if init_state is not None:
common_cula["init_state"] = init_state.transpose(-1, -2).contiguous()

# Accuracy
o_fla, _ = run_fla(**common)
o_cula, _ = run_cula(**common)
o_fla, ht_fla = run_fla(**common)
o_cula, ht_cula_vk = run_cula(**common_cula)
torch.cuda.synchronize()

relative_rms_error, rel_max, mean_diff = relative_rms_error_rel_max_mean_abs(o_fla, o_cula)
if ht_fla is not None and ht_cula_vk is not None:
ht_cula = ht_cula_vk.transpose(-1, -2)
state_rms, state_max, state_mean = relative_rms_error_rel_max_mean_abs(ht_fla, ht_cula)
relative_rms_error = max(relative_rms_error, state_rms)
rel_max = max(rel_max, state_max)
mean_diff = max(mean_diff, state_mean)

# Performance
ms_fla = benchmark_cuda_mode_fn(
Expand All @@ -193,7 +202,7 @@ def bench_fixed(configs):
sanitizer_mode=SANITIZER_MODE,
)
ms_cula = benchmark_cuda_mode_fn(
lambda: run_cula(**common),
lambda: run_cula(**common_cula),
default_warmup=WARMUP,
default_rep=N_ITERS,
ncu_mode=NCU_MODE,
Expand All @@ -216,7 +225,7 @@ def bench_fixed(configs):
}
)

del o_fla, o_cula, q, k, v, g, beta, A_log, dt_bias, inputs
del o_fla, o_cula, ht_fla, ht_cula_vk, q, k, v, g, beta, A_log, dt_bias, inputs
torch.cuda.empty_cache()

return results
Expand Down Expand Up @@ -267,13 +276,22 @@ def bench_varlen(configs):
cu_seqlens=cu_seqlens,
lower_bound=lower_bound,
)
common_cula = dict(common)
if init_state is not None:
common_cula["init_state"] = init_state.transpose(-1, -2).contiguous()

# Accuracy
o_fla, _ = run_fla(**common)
o_cula, _ = run_cula(**common)
o_fla, ht_fla = run_fla(**common)
o_cula, ht_cula_vk = run_cula(**common_cula)
torch.cuda.synchronize()

relative_rms_error, rel_max, mean_diff = relative_rms_error_rel_max_mean_abs(o_fla, o_cula)
if ht_fla is not None and ht_cula_vk is not None:
ht_cula = ht_cula_vk.transpose(-1, -2)
state_rms, state_max, state_mean = relative_rms_error_rel_max_mean_abs(ht_fla, ht_cula)
relative_rms_error = max(relative_rms_error, state_rms)
rel_max = max(rel_max, state_max)
mean_diff = max(mean_diff, state_mean)

# Performance
ms_fla = benchmark_cuda_mode_fn(
Expand All @@ -284,7 +302,7 @@ def bench_varlen(configs):
sanitizer_mode=SANITIZER_MODE,
)
ms_cula = benchmark_cuda_mode_fn(
lambda: run_cula(**common),
lambda: run_cula(**common_cula),
default_warmup=WARMUP,
default_rep=N_ITERS,
ncu_mode=NCU_MODE,
Expand Down Expand Up @@ -314,7 +332,7 @@ def bench_varlen(configs):
}
)

del o_fla, o_cula, q, k, v, g, beta, A_log, dt_bias, inputs
del o_fla, o_cula, ht_fla, ht_cula_vk, q, k, v, g, beta, A_log, dt_bias, inputs
torch.cuda.empty_cache()

return results
Expand Down
200 changes: 200 additions & 0 deletions benchmarks/bench_qwen35_decode.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,200 @@
#!/usr/bin/env python3
"""Benchmark actual cuLA Qwen GDN decode against SGLang's packed inference path.

Only config.json is read. State reset is outside both CUDA event windows.
"""

from __future__ import annotations

import argparse
import csv
import json
import pathlib
import statistics
import sys
from collections.abc import Callable

import torch

ROOT = pathlib.Path(__file__).resolve().parent.parent
sys.path.insert(0, str(ROOT))

import cula.cudac as cula_cuda


def load_shape(config_path: pathlib.Path, tp_size: int) -> dict[str, int | str]:
with config_path.open(encoding="utf-8") as f:
root = json.load(f)
config = root.get("text_config", root)
global_h = int(config["linear_num_key_heads"])
global_hv = int(config["linear_num_value_heads"])
if global_h % tp_size or global_hv % tp_size:
raise ValueError(f"TP={tp_size} must divide H={global_h} and HV={global_hv}")
h, hv = global_h // tp_size, global_hv // tp_size
k = int(config["linear_key_head_dim"])
v = int(config["linear_value_head_dim"])
if hv % h or k != 128 or v != 128:
raise ValueError(f"unsupported local GVA shape H={h} HV={hv} K={k} V={v}")
return {
"model": config_path.parent.name,
"global_h": global_h,
"global_hv": global_hv,
"h": h,
"hv": hv,
"k": k,
"v": v,
}


def load_sglang(sglang_path: pathlib.Path):
for candidate in (sglang_path, sglang_path / "python"):
if candidate.exists():
sys.path.insert(0, str(candidate))
from sglang.srt.layers.attention.linear.kernels.gdn_triton import TritonGDNKernel

kernel = TritonGDNKernel()
if not kernel.supports_packed_decode:
raise RuntimeError("SGLang Triton packed GDN decode is unavailable")
return kernel


def benchmark_cuda(
fn: Callable[[], object],
*,
setup: Callable[[], None],
warmup: int,
rep: int,
) -> float:
for _ in range(warmup):
setup()
fn()
torch.cuda.synchronize()

starts = [torch.cuda.Event(enable_timing=True) for _ in range(rep)]
ends = [torch.cuda.Event(enable_timing=True) for _ in range(rep)]
for start, end in zip(starts, ends, strict=True):
setup()
start.record()
fn()
end.record()
torch.cuda.synchronize()
samples = sorted(start.elapsed_time(end) for start, end in zip(starts, ends, strict=True))
if len(samples) < 4:
return statistics.mean(samples)
return statistics.mean(samples[len(samples) // 4 : 3 * len(samples) // 4])


def relative_rms(reference: torch.Tensor, actual: torch.Tensor) -> float:
ref = reference.float()
diff = ref - actual.float()
return (diff.square().mean().sqrt() / ref.square().mean().sqrt().clamp_min(1e-8)).item()


def make_inputs(tokens: int, shape: dict[str, int | str], seed: int) -> dict[str, torch.Tensor]:
torch.manual_seed(seed)
device = torch.device("cuda")
h, hv, k, v = (int(shape[name]) for name in ("h", "hv", "k", "v"))
conv_dim = 2 * h * k + hv * v
state_kv = torch.randn(tokens, hv, k, v, device=device, dtype=torch.float32) * 0.01
return {
"mixed_qkv": torch.randn(tokens, conv_dim, device=device, dtype=torch.bfloat16),
"a": torch.randn(tokens, hv, device=device, dtype=torch.bfloat16),
"b": torch.randn(tokens, hv, device=device, dtype=torch.bfloat16),
"A_log": -torch.rand(hv, device=device, dtype=torch.float32),
"dt_bias": torch.randn(hv, device=device, dtype=torch.float32) * 0.1,
"state_kv": state_kv,
"state_vk": state_kv.transpose(-1, -2).contiguous(),
"indices": torch.arange(tokens, device=device, dtype=torch.int32),
}


@torch.inference_mode()
def run_case(tokens: int, shape, sglang_kernel, args) -> dict[str, float | int]:
x = make_inputs(tokens, shape, args.seed)
hv, k, v = (int(shape[name]) for name in ("hv", "k", "v"))
state_cula = torch.empty_like(x["state_kv"])
state_sglang = torch.empty_like(x["state_vk"])
out_cula = torch.empty(tokens, hv, v, device="cuda", dtype=torch.bfloat16)

def setup_cula():
state_cula.copy_(x["state_kv"])

def setup_sglang():
state_sglang.copy_(x["state_vk"])

def run_cula():
cula_cuda.qwen35_layout_scalar_kda_decode(
x["mixed_qkv"], x["a"], x["b"], x["A_log"], x["dt_bias"],
state_cula, x["indices"], out_cula,
)

def run_sglang():
return sglang_kernel.packed_decode(
mixed_qkv=x["mixed_qkv"], a=x["a"], b=x["b"],
A_log=x["A_log"], dt_bias=x["dt_bias"], scale=k**-0.5,
ssm_states=state_sglang, cache_indices=x["indices"],
num_v_heads=hv, head_v_dim=v,
)

setup_cula()
run_cula()
setup_sglang()
out_sglang = run_sglang().squeeze(0)
torch.cuda.synchronize()
out_rrms = relative_rms(out_sglang, out_cula)
state_rrms = relative_rms(state_sglang, state_cula.transpose(-1, -2))

sglang_ms = benchmark_cuda(run_sglang, setup=setup_sglang, warmup=args.warmup, rep=args.rep)
cula_ms = benchmark_cuda(run_cula, setup=setup_cula, warmup=args.warmup, rep=args.rep)
return {
"tokens": tokens,
"sglang_packed_ms": sglang_ms,
"cula_fused_ms": cula_ms,
"speedup": sglang_ms / cula_ms,
"out_rel_rms": out_rrms,
"state_rel_rms": state_rrms,
}


def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--config-json", type=pathlib.Path, required=True)
parser.add_argument("--sglang-path", type=pathlib.Path, default=pathlib.Path("/sgl-workspace/sglang"))
parser.add_argument("--tp-size", type=int, choices=(1, 2, 4, 8), default=1)
parser.add_argument("--tokens", type=int, nargs="+", default=(1, 2, 4, 8, 16, 32, 64, 128))
parser.add_argument("--warmup", type=int, default=20)
parser.add_argument("--rep", type=int, default=100)
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--csv", type=pathlib.Path)
args = parser.parse_args()

if not torch.cuda.is_available():
raise RuntimeError("CUDA is required")
shape = load_shape(args.config_json, args.tp_size)
sglang_kernel = load_sglang(args.sglang_path)
print(
f"model={shape['model']} device={torch.cuda.get_device_name(0)} TP={args.tp_size} "
f"global_H/HV={shape['global_h']}/{shape['global_hv']} "
f"local_H/HV={shape['h']}/{shape['hv']} K/V={shape['k']}/{shape['v']}"
)
print("state reset is outside timing; SGLang packed decode vs cuLA fused packed decode")
print("| tokens | sglang_packed_ms | cula_fused_ms | speedup | out_rrms | state_rrms |")
print("|---:|---:|---:|---:|---:|---:|")
rows = []
for tokens in args.tokens:
row = run_case(tokens, shape, sglang_kernel, args)
rows.append(row)
print(
f"| {tokens} | {row['sglang_packed_ms']:.4f} | {row['cula_fused_ms']:.4f} | "
f"{row['speedup']:.3f}x | {row['out_rel_rms']:.3e} | {row['state_rel_rms']:.3e} |"
)
if args.csv:
args.csv.parent.mkdir(parents=True, exist_ok=True)
with args.csv.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=list(rows[0]))
writer.writeheader()
writer.writerows(rows)


if __name__ == "__main__":
main()
Loading