From 1360848ee1e29083132998ae0084f9683af3033d Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <35585791+BBuf@users.noreply.github.com> Date: Sat, 2 May 2026 20:54:46 +0800 Subject: [PATCH] Optimize large GroupNorm SiLU apply (#23938) --- .../diffusion/bench_group_norm_silu.py | 281 ++++++++++++++++++ .../diffusion/triton/group_norm_silu.py | 96 ++++-- 2 files changed, 360 insertions(+), 17 deletions(-) create mode 100644 python/sglang/jit_kernel/benchmark/diffusion/bench_group_norm_silu.py diff --git a/python/sglang/jit_kernel/benchmark/diffusion/bench_group_norm_silu.py b/python/sglang/jit_kernel/benchmark/diffusion/bench_group_norm_silu.py new file mode 100644 index 000000000..0bd48a307 --- /dev/null +++ b/python/sglang/jit_kernel/benchmark/diffusion/bench_group_norm_silu.py @@ -0,0 +1,281 @@ +import argparse +import csv +import statistics +import sys +from dataclasses import dataclass +from pathlib import Path +from typing import Callable + +import torch +import torch.nn.functional as F +import triton.testing + +from sglang.jit_kernel.diffusion.triton.group_norm_silu import triton_group_norm_silu +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.utils import is_in_ci + +register_cuda_ci( + est_time=45, + suite="stage-b-kernel-benchmark-1-gpu-large", + disabled="standalone benchmark", +) + +DEVICE = "cuda" +EPS = 1e-5 +QUANTILES = [0.5, 0.2, 0.8] + + +@dataclass(frozen=True) +class Case: + name: str + shape: tuple[int, ...] + num_groups: int + + +CASES = [ + Case("token_2d", (4, 128), 32), + Case("image_2d", (2, 64, 32, 32), 32), + Case("video_3d_small", (1, 64, 4, 16, 16), 32), + Case("threshold_3d", (1, 128, 1, 256, 256), 32), + Case("hunyuan_video_large", (1, 128, 20, 256, 256), 32), +] +CASE_BY_NAME = {case.name: case for case in CASES} + + +def dtype_from_name(name: str) -> torch.dtype: + mapping = { + "bf16": torch.bfloat16, + "bfloat16": torch.bfloat16, + "fp16": torch.float16, + "float16": torch.float16, + "fp32": torch.float32, + "float32": torch.float32, + } + return mapping[name] + + +def dtype_name(dtype: torch.dtype) -> str: + mapping = { + torch.bfloat16: "bf16", + torch.float16: "fp16", + torch.float32: "fp32", + } + return mapping[dtype] + + +def parse_dtypes(text: str) -> list[torch.dtype]: + return [dtype_from_name(item.strip()) for item in text.split(",") if item.strip()] + + +def parse_cases(text: str) -> list[Case]: + if text == "all": + return CASES + names = [item.strip() for item in text.split(",") if item.strip()] + missing = sorted(set(names) - CASE_BY_NAME.keys()) + if missing: + raise ValueError(f"Unknown cases: {missing}") + return [CASE_BY_NAME[name] for name in names] + + +def tolerance(dtype: torch.dtype) -> tuple[float, float]: + if dtype == torch.float32: + return 1e-5, 1e-5 + if dtype == torch.bfloat16: + return 7e-2, 2e-2 + return 3e-3, 3e-3 + + +def native_group_norm_silu( + x: torch.Tensor, + weight: torch.Tensor, + bias: torch.Tensor, + num_groups: int, +) -> torch.Tensor: + return F.silu(F.group_norm(x, num_groups, weight=weight, bias=bias, eps=EPS)) + + +def make_inputs(case: Case, dtype: torch.dtype) -> tuple[torch.Tensor, ...]: + generator = torch.Generator(device=DEVICE) + generator.manual_seed(len(case.shape) * 1009 + case.shape[1] * 17 + case.num_groups) + x = torch.randn(case.shape, device=DEVICE, dtype=dtype, generator=generator) + weight = torch.randn(case.shape[1], device=DEVICE, dtype=dtype, generator=generator) + bias = torch.randn(case.shape[1], device=DEVICE, dtype=dtype, generator=generator) + return x, weight, bias + + +def do_bench_us(fn: Callable[[], object], warmup: int, rep: int) -> tuple[float, ...]: + median_ms, p20_ms, p80_ms = triton.testing.do_bench( + fn, + quantiles=QUANTILES, + warmup=warmup, + rep=rep, + ) + return median_ms * 1000.0, p20_ms * 1000.0, p80_ms * 1000.0 + + +def summarize(values: list[float]) -> float: + return statistics.median(values) + + +def run_case( + case: Case, + dtype: torch.dtype, + rounds: int, + warmup: int, + rep: int, +) -> dict[str, object]: + x, weight, bias = make_inputs(case, dtype) + + with torch.inference_mode(): + actual = triton_group_norm_silu( + x, weight, bias, num_groups=case.num_groups, eps=EPS + ) + expected = native_group_norm_silu(x, weight, bias, case.num_groups) + atol, rtol = tolerance(dtype) + torch.testing.assert_close(actual, expected, atol=atol, rtol=rtol) + + native_stats = [] + fused_stats = [] + for _ in range(rounds): + native_stats.append( + do_bench_us( + lambda: native_group_norm_silu(x, weight, bias, case.num_groups), + warmup=warmup, + rep=rep, + ) + ) + fused_stats.append( + do_bench_us( + lambda: triton_group_norm_silu( + x, weight, bias, num_groups=case.num_groups, eps=EPS + ), + warmup=warmup, + rep=rep, + ) + ) + + native_median_us = summarize([stats[0] for stats in native_stats]) + fused_median_us = summarize([stats[0] for stats in fused_stats]) + torch.cuda.empty_cache() + return { + "case": case.name, + "shape": "x".join(str(dim) for dim in case.shape), + "groups": case.num_groups, + "dtype": dtype_name(dtype), + "native_median_us": native_median_us, + "native_p20_us": summarize([stats[1] for stats in native_stats]), + "native_p80_us": summarize([stats[2] for stats in native_stats]), + "fused_median_us": fused_median_us, + "fused_p20_us": summarize([stats[1] for stats in fused_stats]), + "fused_p80_us": summarize([stats[2] for stats in fused_stats]), + "speedup": native_median_us / fused_median_us, + "rounds": rounds, + "warmup": warmup, + "rep": rep, + } + + +def run_profile(case: Case, dtype: torch.dtype, provider: str, iters: int) -> None: + x, weight, bias = make_inputs(case, dtype) + + if provider == "native": + + def fn() -> torch.Tensor: + return native_group_norm_silu(x, weight, bias, case.num_groups) + + elif provider == "fused": + + def fn() -> torch.Tensor: + return triton_group_norm_silu( + x, weight, bias, num_groups=case.num_groups, eps=EPS + ) + + else: + raise ValueError(f"Unknown provider: {provider}") + + with torch.inference_mode(): + for _ in range(5): + fn() + torch.cuda.synchronize() + for _ in range(iters): + fn() + torch.cuda.synchronize() + + +def write_csv(rows: list[dict[str, object]], output_path: Path) -> None: + output_path.parent.mkdir(parents=True, exist_ok=True) + fieldnames = list(rows[0].keys()) if rows else [] + with output_path.open("w", newline="", encoding="utf-8") as f: + writer = csv.DictWriter(f, fieldnames=fieldnames) + writer.writeheader() + writer.writerows(rows) + + +def print_rows(rows: list[dict[str, object]]) -> None: + header = ( + "case", + "dtype", + "shape", + "native_us", + "fused_us", + "speedup", + ) + print("| " + " | ".join(header) + " |") + print("|---|---|---|---:|---:|---:|") + for row in rows: + print( + "| {case} | {dtype} | {shape} | {native:.2f} | {fused:.2f} | {speedup:.3f}x |".format( + case=row["case"], + dtype=row["dtype"], + shape=row["shape"], + native=row["native_median_us"], + fused=row["fused_median_us"], + speedup=row["speedup"], + ) + ) + + +def main() -> None: + parser = argparse.ArgumentParser( + description="Benchmark fused GroupNorm+SiLU against PyTorch GroupNorm+SiLU." + ) + parser.add_argument("--cases", default="all") + parser.add_argument("--dtypes", default="bf16,fp16") + parser.add_argument("--rounds", type=int, default=3) + parser.add_argument("--warmup", type=int, default=25) + parser.add_argument("--rep", type=int, default=100) + parser.add_argument("--output-csv", default="") + parser.add_argument("--profile-provider", choices=["native", "fused"], default="") + parser.add_argument("--profile-iters", type=int, default=20) + args = parser.parse_args() + + if not torch.cuda.is_available(): + raise RuntimeError("CUDA is required for this benchmark.") + + cases = parse_cases(args.cases) + dtypes = parse_dtypes(args.dtypes) + + if args.profile_provider: + if len(cases) != 1 or len(dtypes) != 1: + raise ValueError( + "--profile-provider requires exactly one case and one dtype" + ) + run_profile(cases[0], dtypes[0], args.profile_provider, args.profile_iters) + return + + rows = [] + for case in cases: + for dtype in dtypes: + rows.append(run_case(case, dtype, args.rounds, args.warmup, args.rep)) + + print_rows(rows) + if args.output_csv: + write_csv(rows, Path(args.output_csv)) + print(f"Wrote {args.output_csv}") + + +if __name__ == "__main__": + if is_in_ci(): + print("Skipping bench_group_norm_silu.py in CI") + sys.exit(0) + main() diff --git a/python/sglang/jit_kernel/diffusion/triton/group_norm_silu.py b/python/sglang/jit_kernel/diffusion/triton/group_norm_silu.py index 635fa1425..dc614b594 100644 --- a/python/sglang/jit_kernel/diffusion/triton/group_norm_silu.py +++ b/python/sglang/jit_kernel/diffusion/triton/group_norm_silu.py @@ -179,6 +179,49 @@ def _group_norm_apply_kernel( tl.store(output_ptr + group_base + idx, y, mask=mask) +@triton.jit +def _group_norm_apply_scalar_affine_kernel( + input_ptr, + weight_ptr, + bias_ptr, + output_ptr, + stats_ptr, + channels, + spatial_size, + num_groups, + channels_per_group, + group_size, + chunks_per_row, + BLOCK_SIZE: tl.constexpr, + BLOCKS_PER_PROGRAM: tl.constexpr, +): + row = tl.program_id(0).to(tl.int64) + chunk_id = tl.program_id(1).to(tl.int64) + + batch_id = row // num_groups + group_id = row - batch_id * num_groups + chunk_start = chunk_id * BLOCK_SIZE * BLOCKS_PER_PROGRAM + group_base = batch_id * channels * spatial_size + group_id * group_size + + channel_id = chunk_start // spatial_size + affine_offset = group_id * channels_per_group + channel_id + weight = tl.load(weight_ptr + affine_offset).to(tl.float32) + bias = tl.load(bias_ptr + affine_offset).to(tl.float32) + + mean = tl.load(stats_ptr + row * 2) + rstd = tl.load(stats_ptr + row * 2 + 1) + offsets = tl.arange(0, BLOCK_SIZE) + + for block_id in range(BLOCKS_PER_PROGRAM): + idx = chunk_start + block_id * BLOCK_SIZE + offsets + mask = idx < group_size + x = tl.load(input_ptr + group_base + idx, mask=mask, other=0.0).to(tl.float32) + y = (x - mean) * rstd + y = y * weight + bias + y = y * tl.sigmoid(y) + tl.store(output_ptr + group_base + idx, y, mask=mask) + + def _group_norm_silu_native( x: torch.Tensor, weight: torch.Tensor, @@ -294,23 +337,42 @@ def _launch_chunked( num_stages=2, ) - _group_norm_apply_kernel[(rows, chunks_per_row)]( - x_flat, - weight, - bias, - y_flat, - stats, - channels, - spatial_size, - num_groups, - channels_per_group, - group_size, - chunks_per_row, - BLOCK_SIZE=_BLOCK_SIZE, - BLOCKS_PER_PROGRAM=_BLOCKS_PER_PROGRAM, - num_warps=8, - num_stages=3, - ) + if spatial_size % _CHUNK_SIZE == 0 and chunks_per_row >= 64: + _group_norm_apply_scalar_affine_kernel[(rows, chunks_per_row)]( + x_flat, + weight, + bias, + y_flat, + stats, + channels, + spatial_size, + num_groups, + channels_per_group, + group_size, + chunks_per_row, + BLOCK_SIZE=_BLOCK_SIZE, + BLOCKS_PER_PROGRAM=_BLOCKS_PER_PROGRAM, + num_warps=4, + num_stages=3, + ) + else: + _group_norm_apply_kernel[(rows, chunks_per_row)]( + x_flat, + weight, + bias, + y_flat, + stats, + channels, + spatial_size, + num_groups, + channels_per_group, + group_size, + chunks_per_row, + BLOCK_SIZE=_BLOCK_SIZE, + BLOCKS_PER_PROGRAM=_BLOCKS_PER_PROGRAM, + num_warps=8, + num_stages=3, + ) return y