chore: vendor sglang v0.5.10 snapshot
This commit is contained in:
330
third_party/sglang/benchmark/kernels/all_reduce/benchmark_aiter.py
vendored
Normal file
330
third_party/sglang/benchmark/kernels/all_reduce/benchmark_aiter.py
vendored
Normal file
@@ -0,0 +1,330 @@
|
||||
"""
|
||||
Benchmark SGLang vs Aiter custom all-reduce across message sizes.
|
||||
Usage:
|
||||
torchrun --nproc_per_node=2 benchmark_aiter.py
|
||||
torchrun --nproc_per_node=4 benchmark_aiter.py
|
||||
torchrun --nproc_per_node=8 benchmark_aiter.py
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Benchmark SGLang vs Aiter custom all-reduce across message sizes."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--backend",
|
||||
type=str,
|
||||
default="gloo",
|
||||
help="Process group backend for the custom-AR control path (must NOT be nccl).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--warmup",
|
||||
type=int,
|
||||
default=5,
|
||||
help="Warmup iterations per size per implementation.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--iters-small",
|
||||
type=int,
|
||||
default=50,
|
||||
help="Benchmark iterations for sizes <= 1MB.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--iters-large",
|
||||
type=int,
|
||||
default=20,
|
||||
help="Benchmark iterations for sizes > 1MB.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--verbose",
|
||||
action="store_true",
|
||||
help="Print per-iteration timings on rank 0 for debugging.",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def get_env_rank_world() -> Tuple[int, int, int]:
|
||||
rank = int(os.environ.get("RANK", "0"))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", "1"))
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", str(rank)))
|
||||
return rank, world_size, local_rank
|
||||
|
||||
|
||||
def init_dist(backend: str):
|
||||
rank, world_size, _ = get_env_rank_world()
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(
|
||||
backend=backend,
|
||||
init_method="env://",
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
)
|
||||
|
||||
|
||||
def get_device(local_rank: int) -> torch.device:
|
||||
torch.cuda.set_device(local_rank)
|
||||
return torch.device(f"cuda:{local_rank}")
|
||||
|
||||
|
||||
def human_size(num_bytes: int) -> str:
|
||||
units = [("B", 1), ("K", 1024), ("M", 1024 * 1024), ("G", 1024 * 1024 * 1024)]
|
||||
for suf, base in reversed(units):
|
||||
if num_bytes % base == 0 and num_bytes >= base:
|
||||
val = num_bytes // base
|
||||
return f"{val}{suf}"
|
||||
return f"{num_bytes}B"
|
||||
|
||||
|
||||
def get_message_sizes() -> List[int]:
|
||||
return [
|
||||
32 * 1024,
|
||||
64 * 1024,
|
||||
128 * 1024,
|
||||
256 * 1024,
|
||||
512 * 1024,
|
||||
1 * 1024 * 1024,
|
||||
2 * 1024 * 1024,
|
||||
4 * 1024 * 1024,
|
||||
8 * 1024 * 1024,
|
||||
16 * 1024 * 1024,
|
||||
32 * 1024 * 1024,
|
||||
64 * 1024 * 1024,
|
||||
]
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def run_once(comm, inp: torch.Tensor) -> Optional[torch.Tensor]:
|
||||
if hasattr(comm, "all_reduce_unreg"):
|
||||
return comm.all_reduce_unreg(inp)
|
||||
if hasattr(comm, "custom_all_reduce"):
|
||||
return comm.custom_all_reduce(inp)
|
||||
raise RuntimeError("No known all-reduce method found on the communicator.")
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def bench_impl(
|
||||
name: str,
|
||||
comm,
|
||||
sizes: List[int],
|
||||
device: torch.device,
|
||||
warmup: int,
|
||||
iters_small: int,
|
||||
iters_large: int,
|
||||
verbose: bool,
|
||||
pg: Optional[dist.ProcessGroup] = None,
|
||||
) -> List[Tuple[int, Optional[float]]]:
|
||||
rank = dist.get_rank()
|
||||
world_size = dist.get_world_size()
|
||||
results: List[Tuple[int, Optional[float]]] = []
|
||||
|
||||
for size_bytes in sizes:
|
||||
elems = size_bytes // 2 # float16: 2 bytes per element
|
||||
inp = torch.empty(elems, dtype=torch.float16, device=device)
|
||||
inp.uniform_(0, 1)
|
||||
|
||||
disabled = False
|
||||
dist.barrier(group=pg)
|
||||
for _ in range(warmup):
|
||||
torch.cuda.synchronize()
|
||||
out = run_once(comm, inp)
|
||||
torch.cuda.synchronize()
|
||||
if out is None:
|
||||
disabled = True
|
||||
break
|
||||
dist.barrier(group=pg)
|
||||
|
||||
if disabled:
|
||||
if rank == 0:
|
||||
print(
|
||||
f"[{name}] {human_size(size_bytes)}: custom AR disabled (skipped)"
|
||||
)
|
||||
results.append((size_bytes, None))
|
||||
continue
|
||||
|
||||
num_iters = iters_small if size_bytes <= (1 * 1024 * 1024) else iters_large
|
||||
|
||||
times_ms: List[float] = []
|
||||
for it in range(num_iters):
|
||||
dist.barrier(group=pg)
|
||||
torch.cuda.synchronize()
|
||||
t0 = time.perf_counter()
|
||||
out = run_once(comm, inp)
|
||||
torch.cuda.synchronize()
|
||||
t1 = time.perf_counter()
|
||||
dist.barrier(group=pg)
|
||||
|
||||
if out is None:
|
||||
disabled = True
|
||||
break
|
||||
|
||||
dt_ms = (t1 - t0) * 1000.0
|
||||
times_ms.append(dt_ms)
|
||||
|
||||
if verbose and rank == 0:
|
||||
print(
|
||||
f"[{name}] size={human_size(size_bytes)} iter={it} time={dt_ms:.3f} ms"
|
||||
)
|
||||
|
||||
if disabled or not times_ms:
|
||||
if rank == 0:
|
||||
print(
|
||||
f"[{name}] {human_size(size_bytes)}: custom AR disabled (no timings)"
|
||||
)
|
||||
results.append((size_bytes, None))
|
||||
continue
|
||||
|
||||
avg_ms_local = sum(times_ms) / len(times_ms)
|
||||
avg_tensor = torch.tensor([avg_ms_local], dtype=torch.float64, device=device)
|
||||
gather_list = [torch.zeros_like(avg_tensor) for _ in range(world_size)]
|
||||
dist.all_gather(gather_list, avg_tensor, group=pg)
|
||||
if rank == 0:
|
||||
avg_ms = float(torch.stack(gather_list).mean().item())
|
||||
print(
|
||||
f"[{name}] {human_size(size_bytes)}: {avg_ms:.3f} ms (avg across ranks)"
|
||||
)
|
||||
results.append((size_bytes, avg_ms))
|
||||
else:
|
||||
results.append((size_bytes, None))
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
rank, world_size, local_rank = get_env_rank_world()
|
||||
|
||||
if world_size not in (2, 4, 6, 8):
|
||||
print(
|
||||
f"[rank {rank}] WARNING: world_size={world_size} not in supported set (2,4,6,8). "
|
||||
"Custom AR may disable itself.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
|
||||
init_dist(args.backend)
|
||||
device = get_device(local_rank)
|
||||
|
||||
# Import after dist init; some libs query torch dist state on import
|
||||
sgl_comm = None
|
||||
aiter_comm = None
|
||||
HAVE_SGLANG = False
|
||||
HAVE_AITER = False
|
||||
|
||||
try:
|
||||
from sglang.srt.distributed.device_communicators.custom_all_reduce import (
|
||||
CustomAllreduce as SGLCustomAllreduce,
|
||||
)
|
||||
|
||||
HAVE_SGLANG = True
|
||||
except Exception as e:
|
||||
if rank == 0:
|
||||
print(f"SGLang CustomAllreduce import failed: {e}", file=sys.stderr)
|
||||
|
||||
try:
|
||||
from aiter.dist.device_communicators.custom_all_reduce import (
|
||||
CustomAllreduce as AiterCustomAllreduce,
|
||||
)
|
||||
|
||||
HAVE_AITER = True
|
||||
except Exception as e:
|
||||
if rank == 0:
|
||||
print(f"Aiter CustomAllreduce import failed: {e}", file=sys.stderr)
|
||||
|
||||
if rank == 0:
|
||||
print(f"Initialized PG backend={args.backend} world_size={world_size}")
|
||||
print(f"Device: {device.type}:{device.index}")
|
||||
print(f"SGLang available: {HAVE_SGLANG}, Aiter available: {HAVE_AITER}")
|
||||
|
||||
pg = dist.group.WORLD
|
||||
sizes = get_message_sizes()
|
||||
max_size = max(sizes) if sizes else (64 * 1024 * 1024)
|
||||
|
||||
if HAVE_SGLANG:
|
||||
try:
|
||||
sgl_comm = SGLCustomAllreduce(group=pg, device=device, max_size=max_size)
|
||||
except Exception as e:
|
||||
if rank == 0:
|
||||
print(
|
||||
f"Failed to construct SGLang CustomAllreduce: {e}", file=sys.stderr
|
||||
)
|
||||
sgl_comm = None
|
||||
|
||||
if HAVE_AITER:
|
||||
try:
|
||||
aiter_comm = AiterCustomAllreduce(
|
||||
group=pg, device=device, max_size=max_size
|
||||
)
|
||||
except Exception as e:
|
||||
if rank == 0:
|
||||
print(
|
||||
f"Failed to construct Aiter CustomAllreduce: {e}", file=sys.stderr
|
||||
)
|
||||
aiter_comm = None
|
||||
|
||||
sgl_results: List[Tuple[int, Optional[float]]] = []
|
||||
aiter_results: List[Tuple[int, Optional[float]]] = []
|
||||
|
||||
if sgl_comm is not None:
|
||||
sgl_results = bench_impl(
|
||||
name="SGLang",
|
||||
comm=sgl_comm,
|
||||
sizes=sizes,
|
||||
device=device,
|
||||
warmup=args.warmup,
|
||||
iters_small=args.iters_small,
|
||||
iters_large=args.iters_large,
|
||||
verbose=args.verbose,
|
||||
pg=pg,
|
||||
)
|
||||
|
||||
if aiter_comm is not None:
|
||||
aiter_results = bench_impl(
|
||||
name="Aiter",
|
||||
comm=aiter_comm,
|
||||
sizes=sizes,
|
||||
device=device,
|
||||
warmup=args.warmup,
|
||||
iters_small=args.iters_small,
|
||||
iters_large=args.iters_large,
|
||||
verbose=args.verbose,
|
||||
pg=pg,
|
||||
)
|
||||
|
||||
for comm in (sgl_comm, aiter_comm):
|
||||
if comm is not None and hasattr(comm, "close"):
|
||||
try:
|
||||
comm.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if dist.get_rank() == 0:
|
||||
print("\nResults (avg ms across ranks; None = disabled/unavailable):")
|
||||
header = f"{'Size':>8} {'SGLang(ms)':>12} {'Aiter(ms)':>11}"
|
||||
print(header)
|
||||
print("-" * len(header))
|
||||
|
||||
sgl_map = {s: v for s, v in sgl_results if v is not None}
|
||||
aiter_map = {s: v for s, v in aiter_results if v is not None}
|
||||
|
||||
for s in sizes:
|
||||
sgl_ms = sgl_map.get(s, None)
|
||||
aiter_ms = aiter_map.get(s, None)
|
||||
print(
|
||||
f"{human_size(s):>8} {('%.3f' % sgl_ms) if sgl_ms is not None else 'None':>12} "
|
||||
f"{('%.3f' % aiter_ms) if aiter_ms is not None else 'None':>11}"
|
||||
)
|
||||
|
||||
dist.barrier()
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
351
third_party/sglang/benchmark/kernels/all_reduce/benchmark_all_reduce.py
vendored
Normal file
351
third_party/sglang/benchmark/kernels/all_reduce/benchmark_all_reduce.py
vendored
Normal file
@@ -0,0 +1,351 @@
|
||||
"""
|
||||
Benchmark SGLang custom all-reduce vs Torch symm-mem all-reduce across message sizes.
|
||||
Usage:
|
||||
torchrun --nproc_per_node=2 benchmark_all_reduce.py
|
||||
torchrun --nproc_per_node=4 benchmark_all_reduce.py
|
||||
torchrun --nproc_per_node=8 benchmark_all_reduce.py
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from sglang.srt.distributed.parallel_state import (
|
||||
destroy_distributed_environment,
|
||||
destroy_model_parallel,
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Benchmark SGLang custom all-reduce vs Torch symm-mem all-reduce across message sizes."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--backend",
|
||||
type=str,
|
||||
default="gloo",
|
||||
help="Process group backend for the custom-AR control path (must NOT be nccl).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--warmup",
|
||||
type=int,
|
||||
default=5,
|
||||
help="Warmup iterations per size per implementation.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--iters-small",
|
||||
type=int,
|
||||
default=50,
|
||||
help="Benchmark iterations for sizes <= 1MB.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--iters-large",
|
||||
type=int,
|
||||
default=20,
|
||||
help="Benchmark iterations for sizes > 1MB.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--verbose",
|
||||
action="store_true",
|
||||
help="Print per-iteration timings on rank 0 for debugging.",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def get_env_rank_world() -> Tuple[int, int, int]:
|
||||
rank = int(os.environ.get("RANK", "0"))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", "1"))
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", str(rank)))
|
||||
return rank, world_size, local_rank
|
||||
|
||||
|
||||
def init_dist(backend: str):
|
||||
rank, world_size, _ = get_env_rank_world()
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(
|
||||
backend=backend,
|
||||
init_method="env://",
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
)
|
||||
|
||||
device = torch.device(f"cuda:{rank}")
|
||||
torch.cuda.set_device(device)
|
||||
distributed_init_method = f"tcp://localhost:23456"
|
||||
init_distributed_environment(
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
distributed_init_method=distributed_init_method,
|
||||
local_rank=rank,
|
||||
)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
return dist.group.WORLD
|
||||
|
||||
|
||||
def get_device(local_rank: int) -> torch.device:
|
||||
torch.cuda.set_device(local_rank)
|
||||
return torch.device(f"cuda:{local_rank}")
|
||||
|
||||
|
||||
def human_size(num_bytes: int) -> str:
|
||||
units = [("B", 1), ("K", 1024), ("M", 1024 * 1024), ("G", 1024 * 1024 * 1024)]
|
||||
for suf, base in reversed(units):
|
||||
if num_bytes % base == 0 and num_bytes >= base:
|
||||
val = num_bytes // base
|
||||
return f"{val}{suf}"
|
||||
return f"{num_bytes}B"
|
||||
|
||||
|
||||
def get_message_sizes() -> List[int]:
|
||||
return [
|
||||
32 * 1024,
|
||||
64 * 1024,
|
||||
128 * 1024,
|
||||
256 * 1024,
|
||||
512 * 1024,
|
||||
1 * 1024 * 1024,
|
||||
2 * 1024 * 1024,
|
||||
4 * 1024 * 1024,
|
||||
8 * 1024 * 1024,
|
||||
16 * 1024 * 1024,
|
||||
32 * 1024 * 1024,
|
||||
64 * 1024 * 1024,
|
||||
]
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def run_once(comm, inp: torch.Tensor) -> Optional[torch.Tensor]:
|
||||
if hasattr(comm, "custom_all_reduce"):
|
||||
return comm.custom_all_reduce(inp)
|
||||
if hasattr(comm, "all_reduce"):
|
||||
return comm.all_reduce(inp)
|
||||
raise RuntimeError("No known all-reduce method found on the communicator.")
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def bench_impl(
|
||||
name: str,
|
||||
comm,
|
||||
sizes: List[int],
|
||||
device: torch.device,
|
||||
warmup: int,
|
||||
iters_small: int,
|
||||
iters_large: int,
|
||||
verbose: bool,
|
||||
pg: Optional[dist.ProcessGroup] = None,
|
||||
) -> List[Tuple[int, Optional[float]]]:
|
||||
rank = dist.get_rank()
|
||||
world_size = dist.get_world_size()
|
||||
results: List[Tuple[int, Optional[float]]] = []
|
||||
|
||||
for size_bytes in sizes:
|
||||
elems = size_bytes // 2 # float16: 2 bytes per element
|
||||
inp = torch.empty(elems, dtype=torch.float16, device=device)
|
||||
inp.uniform_(0, 1)
|
||||
|
||||
disabled = False
|
||||
dist.barrier(group=pg)
|
||||
for _ in range(warmup):
|
||||
torch.cuda.synchronize()
|
||||
out = run_once(comm, inp)
|
||||
torch.cuda.synchronize()
|
||||
if out is None:
|
||||
disabled = True
|
||||
break
|
||||
dist.barrier(group=pg)
|
||||
|
||||
if disabled:
|
||||
if rank == 0:
|
||||
print(
|
||||
f"[{name}] {human_size(size_bytes)}: custom AR disabled (skipped)"
|
||||
)
|
||||
results.append((size_bytes, None))
|
||||
continue
|
||||
|
||||
num_iters = iters_small if size_bytes <= (1 * 1024 * 1024) else iters_large
|
||||
|
||||
times_ms: List[float] = []
|
||||
for it in range(num_iters):
|
||||
dist.barrier(group=pg)
|
||||
torch.cuda.synchronize()
|
||||
t0 = time.perf_counter()
|
||||
out = run_once(comm, inp)
|
||||
torch.cuda.synchronize()
|
||||
t1 = time.perf_counter()
|
||||
dist.barrier(group=pg)
|
||||
|
||||
if out is None:
|
||||
disabled = True
|
||||
break
|
||||
|
||||
dt_ms = (t1 - t0) * 1000.0
|
||||
times_ms.append(dt_ms)
|
||||
|
||||
if verbose and rank == 0:
|
||||
print(
|
||||
f"[{name}] size={human_size(size_bytes)} iter={it} time={dt_ms:.3f} ms"
|
||||
)
|
||||
|
||||
if disabled or not times_ms:
|
||||
if rank == 0:
|
||||
print(
|
||||
f"[{name}] {human_size(size_bytes)}: custom AR disabled (no timings)"
|
||||
)
|
||||
results.append((size_bytes, None))
|
||||
continue
|
||||
|
||||
avg_ms_local = sum(times_ms) / len(times_ms)
|
||||
avg_tensor = torch.tensor([avg_ms_local], dtype=torch.float64, device=device)
|
||||
gather_list = [torch.zeros_like(avg_tensor) for _ in range(world_size)]
|
||||
dist.all_gather(gather_list, avg_tensor, group=pg)
|
||||
if rank == 0:
|
||||
avg_ms = float(torch.stack(gather_list).mean().item())
|
||||
print(
|
||||
f"[{name}] {human_size(size_bytes)}: {avg_ms:.3f} ms (avg across ranks)"
|
||||
)
|
||||
results.append((size_bytes, avg_ms))
|
||||
else:
|
||||
results.append((size_bytes, None))
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
rank, world_size, local_rank = get_env_rank_world()
|
||||
|
||||
if world_size not in (2, 4, 6, 8):
|
||||
print(
|
||||
f"[rank {rank}] WARNING: world_size={world_size} not in supported set (2,4,6,8). "
|
||||
"Custom AR may disable itself.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
|
||||
group = init_dist(args.backend)
|
||||
device = get_device(local_rank)
|
||||
|
||||
# Import after dist init; some libs query torch dist state on import
|
||||
torch_symm_mem_comm = None
|
||||
HAVE_SGLANG_CUSTOM = False
|
||||
HAVE_TORCH_SYMM_MEM = False
|
||||
|
||||
try:
|
||||
from sglang.srt.distributed.device_communicators.custom_all_reduce import (
|
||||
CustomAllreduce as SGLCustomAllreduce,
|
||||
)
|
||||
|
||||
HAVE_SGLANG_CUSTOM = True
|
||||
except Exception as e:
|
||||
if rank == 0:
|
||||
print(f"SGLang CustomAllreduce import failed: {e}", file=sys.stderr)
|
||||
|
||||
try:
|
||||
from sglang.srt.distributed.device_communicators.torch_symm_mem import (
|
||||
TorchSymmMemCommunicator as TorchSymmMemAllreduce,
|
||||
)
|
||||
|
||||
HAVE_TORCH_SYMM_MEM = True
|
||||
except Exception as e:
|
||||
if rank == 0:
|
||||
print(f"TorchSymmMemAllreduce import failed: {e}", file=sys.stderr)
|
||||
|
||||
if rank == 0:
|
||||
print(f"Initialized PG backend={args.backend} world_size={world_size}")
|
||||
print(f"Device: {device.type}:{device.index}")
|
||||
print(
|
||||
f"SGLang Custom available: {HAVE_SGLANG_CUSTOM}, Torch Symm-Mem available: {HAVE_TORCH_SYMM_MEM}"
|
||||
)
|
||||
|
||||
sizes = get_message_sizes()
|
||||
max_size = max(sizes) if sizes else (128 * 1024 * 1024)
|
||||
|
||||
if HAVE_SGLANG_CUSTOM:
|
||||
try:
|
||||
sgl_custom_comm = SGLCustomAllreduce(
|
||||
group=group, device=device, max_size=max_size
|
||||
)
|
||||
except Exception as e:
|
||||
if rank == 0:
|
||||
print(
|
||||
f"Failed to construct SGLangCustomAllreduce: {e}", file=sys.stderr
|
||||
)
|
||||
sgl_custom_comm = None
|
||||
|
||||
if HAVE_TORCH_SYMM_MEM:
|
||||
try:
|
||||
torch_symm_mem_comm = TorchSymmMemAllreduce(group=group, device=device)
|
||||
except Exception as e:
|
||||
if rank == 0:
|
||||
print(
|
||||
f"Failed to construct TorchSymmMemAllreduce: {e}", file=sys.stderr
|
||||
)
|
||||
torch_symm_mem_comm = None
|
||||
|
||||
sgl_custom_results: List[Tuple[int, Optional[float]]] = []
|
||||
symm_mem_results: List[Tuple[int, Optional[float]]] = []
|
||||
|
||||
if sgl_custom_comm is not None:
|
||||
sgl_custom_results = bench_impl(
|
||||
name="SGLangCustom",
|
||||
comm=sgl_custom_comm,
|
||||
sizes=sizes,
|
||||
device=device,
|
||||
warmup=args.warmup,
|
||||
iters_small=args.iters_small,
|
||||
iters_large=args.iters_large,
|
||||
verbose=args.verbose,
|
||||
pg=group,
|
||||
)
|
||||
|
||||
if torch_symm_mem_comm is not None:
|
||||
symm_mem_results = bench_impl(
|
||||
name="TorchSymmMem",
|
||||
comm=torch_symm_mem_comm,
|
||||
sizes=sizes,
|
||||
device=device,
|
||||
warmup=args.warmup,
|
||||
iters_small=args.iters_small,
|
||||
iters_large=args.iters_large,
|
||||
verbose=args.verbose,
|
||||
pg=group,
|
||||
)
|
||||
|
||||
for comm in (sgl_custom_comm, torch_symm_mem_comm):
|
||||
if comm is not None and hasattr(comm, "close"):
|
||||
try:
|
||||
comm.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if dist.get_rank() == 0:
|
||||
print(
|
||||
f"\nResults (avg ms across {world_size} ranks; None = disabled/unavailable):"
|
||||
)
|
||||
header = f"{'Size':>8} {'CustomAR(ms)':>12} {'TorchSymmMem(ms)':>11}"
|
||||
print(header)
|
||||
print("-" * len(header))
|
||||
|
||||
sgl_custom_map = {s: v for s, v in sgl_custom_results if v is not None}
|
||||
symm_mem_map = {s: v for s, v in symm_mem_results if v is not None}
|
||||
|
||||
for s in sizes:
|
||||
sgl_ms = sgl_custom_map.get(s, None)
|
||||
symm_mem_ms = symm_mem_map.get(s, None)
|
||||
print(
|
||||
f"{human_size(s):>8} {('%.3f' % sgl_ms) if sgl_ms is not None else 'None':>12} "
|
||||
f"{('%.3f' % symm_mem_ms) if symm_mem_ms is not None else 'None':>11}"
|
||||
)
|
||||
torch.distributed.barrier(group=group)
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
536
third_party/sglang/benchmark/kernels/all_reduce/benchmark_fused_ar_rms_amd.py
vendored
Normal file
536
third_party/sglang/benchmark/kernels/all_reduce/benchmark_fused_ar_rms_amd.py
vendored
Normal file
@@ -0,0 +1,536 @@
|
||||
"""
|
||||
Benchmark fused allreduce+rmsnorm on AMD with correctness checks.
|
||||
|
||||
This script targets the same fused op used by SGLang:
|
||||
`tensor_model_parallel_fused_allreduce_rmsnorm`.
|
||||
|
||||
It reports:
|
||||
- eager mode latency (prefill-like)
|
||||
- graph mode latency (decode-like)
|
||||
- fused availability (whether fused path returns non-None)
|
||||
- correctness (fused output matches split allreduce + rmsnorm reference)
|
||||
|
||||
Usage example:
|
||||
torchrun --nproc_per_node=8 \
|
||||
benchmark/kernels/all_reduce/benchmark_fused_ar_rms_amd.py \
|
||||
--dtype bfloat16 \
|
||||
--prefill-shapes 2048x8192,8192x8192 \
|
||||
--decode-shapes 1x8192,4x8192,16x8192 \
|
||||
--warmup 10 --iters 30 --repeats 5
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import os
|
||||
import statistics
|
||||
from typing import Dict, List, Optional, Sequence, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn.functional as F
|
||||
|
||||
from sglang.srt.distributed.communication_op import (
|
||||
tensor_model_parallel_all_reduce,
|
||||
tensor_model_parallel_fused_allreduce_rmsnorm,
|
||||
)
|
||||
from sglang.srt.distributed.parallel_state import (
|
||||
destroy_distributed_environment,
|
||||
destroy_model_parallel,
|
||||
graph_capture,
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
set_custom_all_reduce,
|
||||
)
|
||||
|
||||
Shape = Tuple[int, int]
|
||||
|
||||
|
||||
def parse_shapes(raw: str) -> List[Shape]:
|
||||
shapes: List[Shape] = []
|
||||
for item in [x.strip() for x in raw.split(",") if x.strip()]:
|
||||
if "x" not in item:
|
||||
raise ValueError(f"Invalid shape '{item}', expected MxN format.")
|
||||
m_str, n_str = item.split("x", 1)
|
||||
m = int(m_str)
|
||||
n = int(n_str)
|
||||
if m <= 0 or n <= 0:
|
||||
raise ValueError(f"Invalid shape '{item}', both dims must be positive.")
|
||||
shapes.append((m, n))
|
||||
if not shapes:
|
||||
raise ValueError("Empty shape list is not allowed.")
|
||||
return shapes
|
||||
|
||||
|
||||
def dtype_from_name(name: str) -> torch.dtype:
|
||||
mapping = {
|
||||
"float16": torch.float16,
|
||||
"fp16": torch.float16,
|
||||
"bfloat16": torch.bfloat16,
|
||||
"bf16": torch.bfloat16,
|
||||
}
|
||||
if name not in mapping:
|
||||
raise ValueError(f"Unsupported dtype: {name}")
|
||||
return mapping[name]
|
||||
|
||||
|
||||
def check_close(
|
||||
a: torch.Tensor, b: torch.Tensor, dtype: torch.dtype
|
||||
) -> Tuple[bool, str]:
|
||||
if dtype == torch.bfloat16:
|
||||
rtol, atol = 2e-2, 1.25e-1
|
||||
else:
|
||||
rtol, atol = 1e-2, 2e-2
|
||||
try:
|
||||
torch.testing.assert_close(a, b, rtol=rtol, atol=atol)
|
||||
return True, "PASS"
|
||||
except AssertionError:
|
||||
max_diff = torch.max(torch.abs(a - b)).item()
|
||||
mean_diff = torch.mean(torch.abs(a - b)).item()
|
||||
return False, f"FAIL(max={max_diff:.6f},mean={mean_diff:.6f})"
|
||||
|
||||
|
||||
def _measure_us(
|
||||
fn,
|
||||
warmup: int,
|
||||
iters: int,
|
||||
repeats: int,
|
||||
device: torch.device,
|
||||
) -> Tuple[float, Dict[str, float]]:
|
||||
for _ in range(warmup):
|
||||
fn()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
start_event = torch.cuda.Event(enable_timing=True)
|
||||
end_event = torch.cuda.Event(enable_timing=True)
|
||||
samples_us: List[float] = []
|
||||
|
||||
for _ in range(max(1, repeats)):
|
||||
_barrier(device)
|
||||
torch.cuda.synchronize()
|
||||
start_event.record()
|
||||
for _ in range(iters):
|
||||
fn()
|
||||
end_event.record()
|
||||
end_event.synchronize()
|
||||
samples_us.append(start_event.elapsed_time(end_event) * 1000.0 / iters)
|
||||
|
||||
sorted_samples = sorted(samples_us)
|
||||
p50 = float(statistics.median(sorted_samples))
|
||||
p95 = float(sorted_samples[int((len(sorted_samples) - 1) * 0.95)])
|
||||
return p50, {
|
||||
"p50_us": p50,
|
||||
"p95_us": p95,
|
||||
"min_us": float(sorted_samples[0]),
|
||||
"max_us": float(sorted_samples[-1]),
|
||||
}
|
||||
|
||||
|
||||
def _barrier(device: torch.device):
|
||||
try:
|
||||
dist.barrier(device_ids=[device.index])
|
||||
except TypeError:
|
||||
dist.barrier()
|
||||
|
||||
|
||||
def _mean_across_ranks(value: float, device: torch.device) -> float:
|
||||
t = torch.tensor([value], dtype=torch.float64, device=device)
|
||||
dist.all_reduce(t, op=dist.ReduceOp.SUM)
|
||||
t /= dist.get_world_size()
|
||||
return float(t.item())
|
||||
|
||||
|
||||
def _all_true_across_ranks(value: bool, device: torch.device) -> bool:
|
||||
t = torch.tensor([1 if value else 0], dtype=torch.int32, device=device)
|
||||
dist.all_reduce(t, op=dist.ReduceOp.MIN)
|
||||
return bool(int(t.item()))
|
||||
|
||||
|
||||
def _make_inputs(
|
||||
shape: Shape,
|
||||
dtype: torch.dtype,
|
||||
seed: int,
|
||||
residual_mode: str,
|
||||
rank: int,
|
||||
device: torch.device,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
m, n = shape
|
||||
torch.manual_seed(seed + rank * 17)
|
||||
x = torch.randn((m, n), dtype=torch.float32, device=device).to(dtype)
|
||||
if residual_mode == "self":
|
||||
residual = x.clone()
|
||||
elif residual_mode == "random":
|
||||
residual = torch.randn((m, n), dtype=torch.float32, device=device).to(dtype)
|
||||
elif residual_mode == "zero":
|
||||
residual = torch.zeros((m, n), dtype=dtype, device=device)
|
||||
else:
|
||||
raise ValueError(f"Unknown residual_mode: {residual_mode}")
|
||||
weight = torch.randn((n,), dtype=torch.float32, device=device).to(dtype)
|
||||
return x, residual, weight
|
||||
|
||||
|
||||
def _split_reference(
|
||||
x: torch.Tensor, residual: torch.Tensor, weight: torch.Tensor, eps: float
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
ar_out = tensor_model_parallel_all_reduce(x.clone())
|
||||
residual_out = ar_out + residual
|
||||
out = F.rms_norm(
|
||||
input=residual_out,
|
||||
normalized_shape=(residual_out.shape[-1],),
|
||||
weight=weight,
|
||||
eps=eps,
|
||||
)
|
||||
return out, residual_out
|
||||
|
||||
|
||||
def bench_eager(
|
||||
x: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
eps: float,
|
||||
warmup: int,
|
||||
iters: int,
|
||||
repeats: int,
|
||||
) -> Dict[str, object]:
|
||||
split_fn = lambda: _split_reference(x, residual, weight, eps)
|
||||
split_us, split_stats = _measure_us(split_fn, warmup, iters, repeats, x.device)
|
||||
|
||||
fused_probe = tensor_model_parallel_fused_allreduce_rmsnorm(
|
||||
x.clone(), residual.clone(), weight, eps
|
||||
)
|
||||
fused_available = fused_probe is not None
|
||||
|
||||
fused_us: Optional[float] = None
|
||||
fused_stats: Optional[Dict[str, float]] = None
|
||||
if fused_available:
|
||||
fused_fn = lambda: tensor_model_parallel_fused_allreduce_rmsnorm(
|
||||
x, residual, weight, eps
|
||||
)
|
||||
fused_us, fused_stats = _measure_us(fused_fn, warmup, iters, repeats, x.device)
|
||||
|
||||
ref_out, ref_residual = _split_reference(x, residual, weight, eps)
|
||||
if fused_available:
|
||||
fused_out, fused_residual = tensor_model_parallel_fused_allreduce_rmsnorm(
|
||||
x.clone(), residual.clone(), weight, eps
|
||||
)
|
||||
out_ok, out_detail = check_close(fused_out, ref_out, x.dtype)
|
||||
res_ok, res_detail = check_close(fused_residual, ref_residual, x.dtype)
|
||||
correctness_ok = out_ok and res_ok
|
||||
correctness_detail = f"out={out_detail}, residual={res_detail}"
|
||||
else:
|
||||
correctness_ok = True
|
||||
correctness_detail = "SKIP(fused_unavailable)"
|
||||
|
||||
return {
|
||||
"split_us": split_us,
|
||||
"split_stats": split_stats,
|
||||
"fused_available": fused_available,
|
||||
"fused_us": fused_us,
|
||||
"fused_stats": fused_stats,
|
||||
"correctness_ok": correctness_ok,
|
||||
"correctness_detail": correctness_detail,
|
||||
}
|
||||
|
||||
|
||||
def bench_graph(
|
||||
x: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
eps: float,
|
||||
warmup: int,
|
||||
iters: int,
|
||||
repeats: int,
|
||||
) -> Dict[str, object]:
|
||||
split_x = x.clone()
|
||||
split_res = residual.clone()
|
||||
split_graph_out: Optional[torch.Tensor] = None
|
||||
|
||||
with graph_capture() as gc:
|
||||
split_graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(split_graph, stream=gc.stream):
|
||||
split_graph_out, _ = _split_reference(split_x, split_res, weight, eps)
|
||||
|
||||
def split_replay():
|
||||
split_graph.replay()
|
||||
|
||||
split_us, split_stats = _measure_us(split_replay, warmup, iters, repeats, x.device)
|
||||
|
||||
fused_probe = tensor_model_parallel_fused_allreduce_rmsnorm(
|
||||
x.clone(), residual.clone(), weight, eps
|
||||
)
|
||||
fused_available = fused_probe is not None
|
||||
|
||||
fused_us: Optional[float] = None
|
||||
fused_stats: Optional[Dict[str, float]] = None
|
||||
fused_graph_out: Optional[torch.Tensor] = None
|
||||
fused_graph_residual: Optional[torch.Tensor] = None
|
||||
|
||||
if fused_available:
|
||||
fused_x = x.clone()
|
||||
fused_res = residual.clone()
|
||||
with graph_capture() as gc:
|
||||
fused_graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(fused_graph, stream=gc.stream):
|
||||
fused_graph_out, fused_graph_residual = (
|
||||
tensor_model_parallel_fused_allreduce_rmsnorm(
|
||||
fused_x, fused_res, weight, eps
|
||||
)
|
||||
)
|
||||
|
||||
def fused_replay():
|
||||
fused_graph.replay()
|
||||
|
||||
fused_us, fused_stats = _measure_us(
|
||||
fused_replay, warmup, iters, repeats, x.device
|
||||
)
|
||||
|
||||
ref_out, ref_residual = _split_reference(x, residual, weight, eps)
|
||||
if (
|
||||
fused_available
|
||||
and fused_graph_out is not None
|
||||
and fused_graph_residual is not None
|
||||
):
|
||||
fused_graph.replay()
|
||||
torch.cuda.synchronize()
|
||||
out_ok, out_detail = check_close(fused_graph_out, ref_out, x.dtype)
|
||||
res_ok, res_detail = check_close(fused_graph_residual, ref_residual, x.dtype)
|
||||
correctness_ok = out_ok and res_ok
|
||||
correctness_detail = f"out={out_detail}, residual={res_detail}"
|
||||
else:
|
||||
correctness_ok = True
|
||||
correctness_detail = "SKIP(fused_unavailable)"
|
||||
|
||||
return {
|
||||
"split_us": split_us,
|
||||
"split_stats": split_stats,
|
||||
"fused_available": fused_available,
|
||||
"fused_us": fused_us,
|
||||
"fused_stats": fused_stats,
|
||||
"correctness_ok": correctness_ok,
|
||||
"correctness_detail": correctness_detail,
|
||||
}
|
||||
|
||||
|
||||
def _shape_bytes(shape: Shape, dtype: torch.dtype) -> int:
|
||||
m, n = shape
|
||||
return m * n * torch.tensor([], dtype=dtype).element_size()
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Benchmark fused allreduce+rmsnorm (prefill eager + decode graph)."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dtype",
|
||||
type=str,
|
||||
default="bf16",
|
||||
choices=["fp16", "bf16", "float16", "bfloat16"],
|
||||
)
|
||||
parser.add_argument("--eps", type=float, default=1e-6)
|
||||
parser.add_argument("--seed", type=int, default=1234)
|
||||
parser.add_argument(
|
||||
"--residual-mode",
|
||||
type=str,
|
||||
default="self",
|
||||
choices=["self", "random", "zero"],
|
||||
help="Use residual=x (self) to match aiter test behavior by default.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prefill-shapes",
|
||||
type=str,
|
||||
default="2048x8192,8192x8192,16384x8192",
|
||||
help="Comma-separated MxN shapes for eager mode.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--decode-shapes",
|
||||
type=str,
|
||||
default="1x8192,2x8192,4x8192,8x8192,16x8192",
|
||||
help="Comma-separated MxN shapes for graph mode.",
|
||||
)
|
||||
parser.add_argument("--warmup", type=int, default=10)
|
||||
parser.add_argument("--iters", type=int, default=30)
|
||||
parser.add_argument("--repeats", type=int, default=5)
|
||||
parser.add_argument(
|
||||
"--mode",
|
||||
type=str,
|
||||
default="both",
|
||||
choices=["eager", "graph", "both"],
|
||||
)
|
||||
parser.add_argument(
|
||||
"--csv-out",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Optional output CSV path (written on rank 0 only).",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
dtype = dtype_from_name(args.dtype)
|
||||
rank = int(os.environ.get("RANK", "0"))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", "1"))
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", str(rank)))
|
||||
torch.cuda.set_device(local_rank % torch.cuda.device_count())
|
||||
device = torch.device(f"cuda:{local_rank % torch.cuda.device_count()}")
|
||||
|
||||
set_custom_all_reduce(True)
|
||||
init_distributed_environment(
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=local_rank,
|
||||
distributed_init_method="env://",
|
||||
backend="nccl",
|
||||
)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
|
||||
prefill_shapes = parse_shapes(args.prefill_shapes)
|
||||
decode_shapes = parse_shapes(args.decode_shapes)
|
||||
|
||||
if rank == 0:
|
||||
print(
|
||||
"Config: "
|
||||
f"world_size={world_size}, dtype={dtype}, residual_mode={args.residual_mode}, "
|
||||
f"warmup={args.warmup}, iters={args.iters}, repeats={args.repeats}"
|
||||
)
|
||||
|
||||
run_modes: Sequence[str]
|
||||
if args.mode == "both":
|
||||
run_modes = ("eager", "graph")
|
||||
else:
|
||||
run_modes = (args.mode,)
|
||||
csv_rows: List[Dict[str, object]] = []
|
||||
|
||||
for mode in run_modes:
|
||||
shapes = prefill_shapes if mode == "eager" else decode_shapes
|
||||
if rank == 0:
|
||||
phase_name = "prefill(eager)" if mode == "eager" else "decode(graph)"
|
||||
print("\n" + "=" * 120)
|
||||
print(f"Mode: {phase_name}")
|
||||
print(
|
||||
"| Shape | Input bytes/rank | Split p50 (us) | Fused p50 (us) | Speedup | Fused available | Correctness |"
|
||||
)
|
||||
print(
|
||||
"|:------|-----------------:|---------------:|---------------:|--------:|:----------------|:------------|"
|
||||
)
|
||||
|
||||
for shape in shapes:
|
||||
x, residual, weight = _make_inputs(
|
||||
shape=shape,
|
||||
dtype=dtype,
|
||||
seed=args.seed,
|
||||
residual_mode=args.residual_mode,
|
||||
rank=rank,
|
||||
device=device,
|
||||
)
|
||||
|
||||
if mode == "eager":
|
||||
metrics = bench_eager(
|
||||
x=x,
|
||||
residual=residual,
|
||||
weight=weight,
|
||||
eps=args.eps,
|
||||
warmup=args.warmup,
|
||||
iters=args.iters,
|
||||
repeats=args.repeats,
|
||||
)
|
||||
else:
|
||||
metrics = bench_graph(
|
||||
x=x,
|
||||
residual=residual,
|
||||
weight=weight,
|
||||
eps=args.eps,
|
||||
warmup=args.warmup,
|
||||
iters=args.iters,
|
||||
repeats=args.repeats,
|
||||
)
|
||||
|
||||
split_us = _mean_across_ranks(float(metrics["split_us"]), device)
|
||||
fused_available = _all_true_across_ranks(
|
||||
bool(metrics["fused_available"]), device
|
||||
)
|
||||
correctness_ok = _all_true_across_ranks(
|
||||
bool(metrics["correctness_ok"]), device
|
||||
)
|
||||
|
||||
fused_us: Optional[float] = None
|
||||
if fused_available and metrics["fused_us"] is not None:
|
||||
fused_us = _mean_across_ranks(float(metrics["fused_us"]), device)
|
||||
|
||||
if rank == 0:
|
||||
m, n = shape
|
||||
shape_str = f"{m}x{n}"
|
||||
bytes_per_rank = _shape_bytes(shape, dtype)
|
||||
if fused_us is not None and fused_us > 0:
|
||||
speedup = split_us / fused_us
|
||||
speedup_str = f"{speedup:.3f}x"
|
||||
fused_str = f"{fused_us:.1f}"
|
||||
else:
|
||||
speedup_str = "N/A"
|
||||
fused_str = "N/A"
|
||||
correctness_text = (
|
||||
"PASS" if correctness_ok else str(metrics["correctness_detail"])
|
||||
)
|
||||
print(
|
||||
f"| {shape_str} | {bytes_per_rank} | {split_us:.1f} | {fused_str} | "
|
||||
f"{speedup_str} | {str(fused_available)} | {correctness_text} |"
|
||||
)
|
||||
csv_rows.append(
|
||||
{
|
||||
"mode": mode,
|
||||
"shape": shape_str,
|
||||
"m": m,
|
||||
"n": n,
|
||||
"bytes_per_rank": bytes_per_rank,
|
||||
"split_p50_us": split_us,
|
||||
"fused_p50_us": fused_us if fused_us is not None else "",
|
||||
"speedup_split_over_fused": (
|
||||
split_us / fused_us
|
||||
if fused_us is not None and fused_us > 0
|
||||
else ""
|
||||
),
|
||||
"fused_available": fused_available,
|
||||
"correctness_ok": correctness_ok,
|
||||
"correctness_detail": correctness_text,
|
||||
"dtype": str(dtype),
|
||||
"world_size": world_size,
|
||||
"residual_mode": args.residual_mode,
|
||||
"warmup": args.warmup,
|
||||
"iters": args.iters,
|
||||
"repeats": args.repeats,
|
||||
}
|
||||
)
|
||||
|
||||
if rank == 0 and args.csv_out:
|
||||
os.makedirs(os.path.dirname(args.csv_out) or ".", exist_ok=True)
|
||||
fieldnames = [
|
||||
"mode",
|
||||
"shape",
|
||||
"m",
|
||||
"n",
|
||||
"bytes_per_rank",
|
||||
"split_p50_us",
|
||||
"fused_p50_us",
|
||||
"speedup_split_over_fused",
|
||||
"fused_available",
|
||||
"correctness_ok",
|
||||
"correctness_detail",
|
||||
"dtype",
|
||||
"world_size",
|
||||
"residual_mode",
|
||||
"warmup",
|
||||
"iters",
|
||||
"repeats",
|
||||
]
|
||||
with open(args.csv_out, "w", newline="", encoding="utf-8") as f:
|
||||
writer = csv.DictWriter(f, fieldnames=fieldnames)
|
||||
writer.writeheader()
|
||||
writer.writerows(csv_rows)
|
||||
print(f"\nSaved CSV to: {args.csv_out}")
|
||||
|
||||
_barrier(device)
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
224
third_party/sglang/benchmark/kernels/all_reduce/benchmark_mscclpp.py
vendored
Normal file
224
third_party/sglang/benchmark/kernels/all_reduce/benchmark_mscclpp.py
vendored
Normal file
@@ -0,0 +1,224 @@
|
||||
"""For Now, MSCCL is only supported on TP16 and TP8 case
|
||||
|
||||
export WORLD_SIZE=1
|
||||
export RANK=0
|
||||
export MASTER_ADDR=127.0.0.1
|
||||
export MASTER_PORT=12345
|
||||
|
||||
torchrun --nproc_per_node gpu \
|
||||
--nnodes $WORLD_SIZE \
|
||||
--node_rank $RANK \
|
||||
--master_addr $MASTER_ADDR \
|
||||
--master_port $MASTER_PORT benchmark/kernels/all_reduce/benchmark_mscclpp.py
|
||||
"""
|
||||
|
||||
import os
|
||||
from contextlib import nullcontext
|
||||
from typing import List
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import ProcessGroup
|
||||
|
||||
from sglang.srt.distributed import init_distributed_environment
|
||||
from sglang.srt.distributed.device_communicators.pymscclpp import PyMscclppCommunicator
|
||||
from sglang.srt.distributed.device_communicators.pynccl import PyNcclCommunicator
|
||||
from sglang.srt.distributed.parallel_state import (
|
||||
get_tensor_model_parallel_group,
|
||||
graph_capture,
|
||||
initialize_model_parallel,
|
||||
set_mscclpp_all_reduce,
|
||||
)
|
||||
|
||||
|
||||
def torch_allreduce(torch_input: torch.Tensor, group: ProcessGroup) -> torch.Tensor:
|
||||
dist.all_reduce(torch_input, group=group)
|
||||
return torch_input
|
||||
|
||||
|
||||
def msccl_allreduce(
|
||||
msccl_input: torch.Tensor, msccl_comm: PyMscclppCommunicator
|
||||
) -> torch.Tensor:
|
||||
return msccl_comm.all_reduce(msccl_input)
|
||||
|
||||
|
||||
def pynccl_allreduce(
|
||||
msccl_input: torch.Tensor, pynccl_comm: PyNcclCommunicator
|
||||
) -> torch.Tensor:
|
||||
pynccl_comm.all_reduce(msccl_input)
|
||||
return msccl_input
|
||||
|
||||
|
||||
def _bench_graph_time(func, inp_randn, warmup_loop=2, graph_loop=10, test_loop=10):
|
||||
graph_input = inp_randn.clone()
|
||||
with graph_capture() as graph_capture_context:
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph, stream=graph_capture_context.stream):
|
||||
for _ in range(graph_loop):
|
||||
graph_out = func(graph_input)
|
||||
|
||||
graph.replay()
|
||||
func_output = graph_out.clone()
|
||||
|
||||
for _ in range(warmup_loop):
|
||||
graph.replay()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
start_event = torch.cuda.Event(enable_timing=True)
|
||||
end_event = torch.cuda.Event(enable_timing=True)
|
||||
|
||||
latencies: List[float] = []
|
||||
for _ in range(test_loop):
|
||||
torch.cuda.synchronize()
|
||||
dist.barrier()
|
||||
start_event.record()
|
||||
graph.replay()
|
||||
end_event.record()
|
||||
end_event.synchronize()
|
||||
latencies.append(start_event.elapsed_time(end_event))
|
||||
func_cost_us = sum(latencies) / len(latencies) / graph_loop * 1000
|
||||
graph.reset()
|
||||
return func_output, func_cost_us
|
||||
|
||||
|
||||
def _bench_eager_time(func, inp_randn, warmup_loop=2, test_loop=10):
|
||||
eager_input = inp_randn.clone()
|
||||
eager_output = func(eager_input)
|
||||
func_output = eager_output.clone()
|
||||
|
||||
for _ in range(warmup_loop):
|
||||
func(eager_input)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
start_event = torch.cuda.Event(enable_timing=True)
|
||||
end_event = torch.cuda.Event(enable_timing=True)
|
||||
torch.cuda.synchronize()
|
||||
start_event.record()
|
||||
for _ in range(test_loop):
|
||||
func(eager_input)
|
||||
end_event.record()
|
||||
torch.cuda.synchronize()
|
||||
func_cost_us = start_event.elapsed_time(end_event) / test_loop * 1000
|
||||
|
||||
return func_output, func_cost_us
|
||||
|
||||
|
||||
def get_torch_prof_ctx(do_prof: bool):
|
||||
ctx = (
|
||||
torch.profiler.profile(
|
||||
activities=[
|
||||
torch.profiler.ProfilerActivity.CPU,
|
||||
torch.profiler.ProfilerActivity.CUDA,
|
||||
],
|
||||
record_shapes=True,
|
||||
with_stack=True,
|
||||
)
|
||||
if do_prof
|
||||
else nullcontext()
|
||||
)
|
||||
return ctx
|
||||
|
||||
|
||||
def human_readable_size(size, decimal_places=1):
|
||||
for unit in ["B", "KiB", "MiB", "GiB", "TiB", "PiB"]:
|
||||
if size < 1024.0 or unit == "PiB":
|
||||
break
|
||||
size /= 1024.0
|
||||
return f"{size:.{decimal_places}f} {unit}"
|
||||
|
||||
|
||||
try:
|
||||
from tabulate import tabulate
|
||||
except ImportError:
|
||||
print("tabulate not installed, skipping table printing")
|
||||
tabulate = None
|
||||
|
||||
|
||||
def print_markdown_table(data):
|
||||
if tabulate is not None:
|
||||
print(tabulate(data, headers="keys", tablefmt="github"))
|
||||
return
|
||||
headers = data[0].keys()
|
||||
header_row = "| " + " | ".join(headers) + " |"
|
||||
separator = "| " + " | ".join(["---"] * len(headers)) + " |"
|
||||
rows = []
|
||||
for item in data:
|
||||
row = "| " + " | ".join(str(item[key]) for key in headers) + " |"
|
||||
rows.append(row)
|
||||
markdown_table = "\n".join([header_row, separator] + rows)
|
||||
print(markdown_table)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import logging
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s - %(levelname)s - %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
force=True,
|
||||
)
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(backend="nccl")
|
||||
world, world_size = dist.group.WORLD, dist.get_world_size()
|
||||
rank = dist.get_rank()
|
||||
torch.cuda.set_device(rank % 8)
|
||||
device = torch.cuda.current_device()
|
||||
set_mscclpp_all_reduce(True)
|
||||
init_distributed_environment(
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=rank % 8,
|
||||
)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
group = get_tensor_model_parallel_group().device_group
|
||||
cpu_group = get_tensor_model_parallel_group().cpu_group
|
||||
pynccl_comm = get_tensor_model_parallel_group().pynccl_comm
|
||||
pymscclpp_comm = get_tensor_model_parallel_group().pymscclpp_comm
|
||||
dist.barrier()
|
||||
profile = False
|
||||
dtype = torch.bfloat16
|
||||
ctx = get_torch_prof_ctx(profile)
|
||||
result = []
|
||||
|
||||
with ctx:
|
||||
for i in range(10, 20):
|
||||
sz = 2**i
|
||||
if sz * dtype.itemsize > 2**20:
|
||||
break
|
||||
inp_randn = torch.randint(1, 16, (sz,), dtype=dtype, device=device)
|
||||
|
||||
memory = torch.empty_like(inp_randn)
|
||||
memory_out = torch.empty_like(memory)
|
||||
torch_eager_output, torch_eager_time = _bench_eager_time(
|
||||
lambda inp: torch_allreduce(inp, group), inp_randn
|
||||
)
|
||||
msccl_eager_output, msccl_eager_time = _bench_eager_time(
|
||||
lambda inp: msccl_allreduce(inp, pymscclpp_comm), inp_randn
|
||||
)
|
||||
msccl_graph_output, msccl_graph_time = _bench_graph_time(
|
||||
lambda inp: msccl_allreduce(inp, pymscclpp_comm), inp_randn
|
||||
)
|
||||
# since pynccl is inplace op, this return result is not correct if graph loop > 1
|
||||
_, pynccl_graph_time = _bench_graph_time(
|
||||
lambda inp: pynccl_allreduce(inp, pynccl_comm), inp_randn
|
||||
)
|
||||
torch.testing.assert_close(torch_eager_output, msccl_graph_output)
|
||||
torch.testing.assert_close(torch_eager_output, msccl_eager_output)
|
||||
result.append(
|
||||
{
|
||||
"msg_size": human_readable_size(inp_randn.nbytes),
|
||||
"torch eager time": torch_eager_time,
|
||||
"msccl eager time": msccl_eager_time,
|
||||
"msccl graph time": msccl_graph_time,
|
||||
"pynccl graph time": pynccl_graph_time,
|
||||
}
|
||||
)
|
||||
if rank == 0:
|
||||
print(f"sz={sz}, dtype={dtype}: correctness check PASS!")
|
||||
if rank == 0:
|
||||
print_markdown_table(result)
|
||||
if profile:
|
||||
prof_dir = f"prof/msccl"
|
||||
os.makedirs(prof_dir, exist_ok=True)
|
||||
ctx.export_chrome_trace(f"{prof_dir}/trace_rank{dist.get_rank()}.json.gz")
|
||||
248
third_party/sglang/benchmark/kernels/all_reduce/benchmark_torch_symm_mem.py
vendored
Normal file
248
third_party/sglang/benchmark/kernels/all_reduce/benchmark_torch_symm_mem.py
vendored
Normal file
@@ -0,0 +1,248 @@
|
||||
"""For Now, TORCH_SYMM_MEM is only supported on following limited tp case
|
||||
|
||||
SM90: {
|
||||
2: 64 * MiB, # 64 MB
|
||||
4: 64 * MiB, # 64 MB
|
||||
6: 128 * MiB, # 128 MB
|
||||
8: 128 * MiB, # 128 MB
|
||||
},
|
||||
SM100: {
|
||||
2: 64 * MiB, # 64 MB
|
||||
4: 64 * MiB, # 64 MB
|
||||
6: 128 * MiB, # 128 MB
|
||||
8: 128 * MiB, # 128 MB
|
||||
}
|
||||
|
||||
export WORLD_SIZE=8
|
||||
export RANK=0
|
||||
export MASTER_ADDR=127.0.0.1
|
||||
export MASTER_PORT=12345
|
||||
|
||||
torchrun --nproc_per_node gpu \
|
||||
--nnodes $WORLD_SIZE \
|
||||
--node_rank $RANK \
|
||||
--master_addr $MASTER_ADDR \
|
||||
--master_port $MASTER_PORT ./benchmark/kernels/all_reduce/benchmark_torch_symm_mem.py
|
||||
"""
|
||||
|
||||
import os
|
||||
from contextlib import nullcontext
|
||||
from typing import List
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import ProcessGroup
|
||||
|
||||
from sglang.srt.distributed import init_distributed_environment
|
||||
from sglang.srt.distributed.device_communicators.pynccl import PyNcclCommunicator
|
||||
from sglang.srt.distributed.device_communicators.torch_symm_mem import (
|
||||
TorchSymmMemCommunicator,
|
||||
)
|
||||
from sglang.srt.distributed.parallel_state import (
|
||||
get_tensor_model_parallel_group,
|
||||
graph_capture,
|
||||
initialize_model_parallel,
|
||||
set_torch_symm_mem_all_reduce,
|
||||
)
|
||||
from sglang.utils import is_in_ci
|
||||
|
||||
IS_CI = is_in_ci()
|
||||
|
||||
|
||||
def torch_allreduce(torch_input: torch.Tensor, group: ProcessGroup) -> torch.Tensor:
|
||||
dist.all_reduce(torch_input, group=group)
|
||||
return torch_input
|
||||
|
||||
|
||||
def torch_symm_mem_allreduce(
|
||||
torch_symm_mem_input: torch.Tensor, torch_symm_mem_comm: TorchSymmMemCommunicator
|
||||
) -> torch.Tensor:
|
||||
return torch_symm_mem_comm.all_reduce(torch_symm_mem_input)
|
||||
|
||||
|
||||
def pynccl_allreduce(
|
||||
pynccl_input: torch.Tensor, pynccl_comm: PyNcclCommunicator
|
||||
) -> torch.Tensor:
|
||||
pynccl_comm.all_reduce(pynccl_input)
|
||||
return pynccl_input
|
||||
|
||||
|
||||
def _bench_graph_time(func, inp_randn, warmup_loop=2, graph_loop=10, test_loop=10):
|
||||
graph_input = inp_randn.clone()
|
||||
with graph_capture() as graph_capture_context:
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph, stream=graph_capture_context.stream):
|
||||
for _ in range(graph_loop):
|
||||
graph_out = func(graph_input)
|
||||
|
||||
graph.replay()
|
||||
func_output = graph_out.clone()
|
||||
|
||||
for _ in range(warmup_loop):
|
||||
graph.replay()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
start_event = torch.cuda.Event(enable_timing=True)
|
||||
end_event = torch.cuda.Event(enable_timing=True)
|
||||
|
||||
latencies: List[float] = []
|
||||
for _ in range(test_loop):
|
||||
torch.cuda.synchronize()
|
||||
dist.barrier()
|
||||
start_event.record()
|
||||
graph.replay()
|
||||
end_event.record()
|
||||
end_event.synchronize()
|
||||
latencies.append(start_event.elapsed_time(end_event))
|
||||
func_cost_us = sum(latencies) / len(latencies) / graph_loop * 1000
|
||||
graph.reset()
|
||||
return func_output, func_cost_us
|
||||
|
||||
|
||||
def _bench_eager_time(func, inp_randn, warmup_loop=2, test_loop=10):
|
||||
eager_input = inp_randn.clone()
|
||||
eager_output = func(eager_input)
|
||||
func_output = eager_output.clone()
|
||||
|
||||
for _ in range(warmup_loop):
|
||||
func(eager_input)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
start_event = torch.cuda.Event(enable_timing=True)
|
||||
end_event = torch.cuda.Event(enable_timing=True)
|
||||
torch.cuda.synchronize()
|
||||
start_event.record()
|
||||
for _ in range(test_loop):
|
||||
func(eager_input)
|
||||
end_event.record()
|
||||
torch.cuda.synchronize()
|
||||
func_cost_us = start_event.elapsed_time(end_event) / test_loop * 1000
|
||||
|
||||
return func_output, func_cost_us
|
||||
|
||||
|
||||
def get_torch_prof_ctx(do_prof: bool):
|
||||
ctx = (
|
||||
torch.profiler.profile(
|
||||
activities=[
|
||||
torch.profiler.ProfilerActivity.CPU,
|
||||
torch.profiler.ProfilerActivity.CUDA,
|
||||
],
|
||||
record_shapes=True,
|
||||
with_stack=True,
|
||||
)
|
||||
if do_prof
|
||||
else nullcontext()
|
||||
)
|
||||
return ctx
|
||||
|
||||
|
||||
def human_readable_size(size, decimal_places=1):
|
||||
for unit in ["B", "KiB", "MiB", "GiB", "TiB", "PiB"]:
|
||||
if size < 1024.0 or unit == "PiB":
|
||||
break
|
||||
size /= 1024.0
|
||||
return f"{size:.{decimal_places}f} {unit}"
|
||||
|
||||
|
||||
try:
|
||||
from tabulate import tabulate
|
||||
except ImportError:
|
||||
print("tabulate not installed, skipping table printing")
|
||||
tabulate = None
|
||||
|
||||
|
||||
def print_markdown_table(data):
|
||||
if tabulate is not None:
|
||||
print(tabulate(data, headers="keys", tablefmt="github"))
|
||||
return
|
||||
headers = data[0].keys()
|
||||
header_row = "| " + " | ".join(headers) + " |"
|
||||
separator = "| " + " | ".join(["---"] * len(headers)) + " |"
|
||||
rows = []
|
||||
for item in data:
|
||||
row = "| " + " | ".join(str(item[key]) for key in headers) + " |"
|
||||
rows.append(row)
|
||||
markdown_table = "\n".join([header_row, separator] + rows)
|
||||
print(markdown_table)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import logging
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s - %(levelname)s - %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
force=True,
|
||||
)
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(backend="nccl")
|
||||
world, world_size = dist.group.WORLD, dist.get_world_size()
|
||||
rank = dist.get_rank()
|
||||
torch.cuda.set_device(rank % 8)
|
||||
device = torch.cuda.current_device()
|
||||
set_torch_symm_mem_all_reduce(True)
|
||||
init_distributed_environment(
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=rank % 8,
|
||||
)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
group = get_tensor_model_parallel_group().device_group
|
||||
cpu_group = get_tensor_model_parallel_group().cpu_group
|
||||
pynccl_comm = get_tensor_model_parallel_group().pynccl_comm
|
||||
torch_symm_mem_comm = get_tensor_model_parallel_group().torch_symm_mem_comm
|
||||
dist.barrier()
|
||||
profile = False
|
||||
dtype = torch.bfloat16
|
||||
ctx = get_torch_prof_ctx(profile)
|
||||
result = []
|
||||
|
||||
with ctx:
|
||||
if IS_CI:
|
||||
i_range = range(10, 11)
|
||||
else:
|
||||
i_range = range(10, 20)
|
||||
for i in i_range:
|
||||
sz = 2**i
|
||||
if sz * dtype.itemsize > 2**24:
|
||||
break
|
||||
inp_randn = torch.randint(1, 16, (sz,), dtype=dtype, device=device)
|
||||
|
||||
memory = torch.empty_like(inp_randn)
|
||||
memory_out = torch.empty_like(memory)
|
||||
torch_eager_output, torch_eager_time = _bench_eager_time(
|
||||
lambda inp: torch_allreduce(inp, group), inp_randn
|
||||
)
|
||||
symm_mem_eager_output, symm_mem_eager_time = _bench_eager_time(
|
||||
lambda inp: torch_symm_mem_allreduce(inp, torch_symm_mem_comm),
|
||||
inp_randn,
|
||||
)
|
||||
symm_mem_graph_output, symm_mem_graph_time = _bench_graph_time(
|
||||
lambda inp: torch_symm_mem_allreduce(inp, torch_symm_mem_comm),
|
||||
inp_randn,
|
||||
)
|
||||
# since pynccl is inplace op, this return result is not correct if graph loop > 1
|
||||
_, pynccl_graph_time = _bench_graph_time(
|
||||
lambda inp: pynccl_allreduce(inp, pynccl_comm), inp_randn
|
||||
)
|
||||
torch.testing.assert_close(torch_eager_output, symm_mem_graph_output)
|
||||
torch.testing.assert_close(torch_eager_output, symm_mem_eager_output)
|
||||
result.append(
|
||||
{
|
||||
"msg_size": human_readable_size(inp_randn.nbytes),
|
||||
"torch eager time": torch_eager_time,
|
||||
"symm mem eager time": symm_mem_eager_time,
|
||||
"symm mem graph time": symm_mem_graph_time,
|
||||
"pynccl graph time": pynccl_graph_time,
|
||||
}
|
||||
)
|
||||
if rank == 0:
|
||||
print(f"sz={sz}, dtype={dtype}: correctness check PASS!")
|
||||
if rank == 0:
|
||||
print_markdown_table(result)
|
||||
if profile:
|
||||
prof_dir = f"prof/torch_symm_mem"
|
||||
os.makedirs(prof_dir, exist_ok=True)
|
||||
ctx.export_chrome_trace(f"{prof_dir}/trace_rank{dist.get_rank()}.json.gz")
|
||||
403
third_party/sglang/benchmark/kernels/decoding_attention_triton/triton_flashinfer_cudnn.py
vendored
Normal file
403
third_party/sglang/benchmark/kernels/decoding_attention_triton/triton_flashinfer_cudnn.py
vendored
Normal file
@@ -0,0 +1,403 @@
|
||||
import itertools
|
||||
import math
|
||||
|
||||
import cudnn
|
||||
import torch
|
||||
import torch.utils.benchmark as benchmark
|
||||
from flashinfer import BatchDecodeWithPagedKVCacheWrapper
|
||||
|
||||
from sglang.srt.layers.attention.flashinfer_backend import should_use_tensor_core
|
||||
from sglang.srt.layers.attention.triton_ops.decode_attention import decode_attention_fwd
|
||||
|
||||
|
||||
def benchmark_forward(
|
||||
fn,
|
||||
*inputs,
|
||||
repeats=10,
|
||||
amp=False,
|
||||
amp_dtype=torch.float16,
|
||||
**kwinputs,
|
||||
):
|
||||
def amp_wrapper(*inputs, **kwinputs):
|
||||
with torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=amp):
|
||||
fn(*inputs, **kwinputs)
|
||||
|
||||
t = benchmark.Timer(
|
||||
stmt="fn_amp(*inputs, **kwinputs)",
|
||||
globals={"fn_amp": amp_wrapper, "inputs": inputs, "kwinputs": kwinputs},
|
||||
num_threads=torch.get_num_threads(),
|
||||
)
|
||||
m = t.timeit(repeats)
|
||||
return t, m
|
||||
|
||||
|
||||
def time_fwd(func, *args, **kwargs):
|
||||
time_f = benchmark_forward(func, *args, **kwargs)
|
||||
return time_f[1].mean * 1e6
|
||||
|
||||
|
||||
def decode_attention_sglang(
|
||||
q,
|
||||
kv_data,
|
||||
batch_size,
|
||||
kv_len,
|
||||
head_num_q,
|
||||
head_num_kv,
|
||||
head_dim,
|
||||
num_kv_splits,
|
||||
warmup=10,
|
||||
):
|
||||
|
||||
k_buffer = kv_data[0].view(-1, head_num_kv, head_dim)
|
||||
v_buffer = kv_data[1].view(-1, head_num_kv, head_dim)
|
||||
o = torch.empty_like(q)
|
||||
total_tokens = batch_size * kv_len
|
||||
req_to_token = torch.arange(0, total_tokens).to(0).int().view(batch_size, kv_len)
|
||||
b_req_idx = torch.arange(0, batch_size).to(0).int()
|
||||
b_seq_len = torch.full((batch_size,), kv_len, dtype=torch.int32, device="cuda")
|
||||
max_len_in_batch = kv_len
|
||||
sm_scale = 1.0 / (head_dim**0.5)
|
||||
|
||||
attn_logits = torch.empty(
|
||||
(batch_size, head_num_q, num_kv_splits, head_dim + 1),
|
||||
dtype=torch.float32,
|
||||
device="cuda",
|
||||
)
|
||||
|
||||
for _ in range(warmup):
|
||||
decode_attention_fwd(
|
||||
q,
|
||||
k_buffer,
|
||||
v_buffer,
|
||||
o,
|
||||
req_to_token,
|
||||
b_req_idx,
|
||||
b_seq_len,
|
||||
attn_logits,
|
||||
num_kv_splits,
|
||||
sm_scale,
|
||||
)
|
||||
|
||||
f = time_fwd(
|
||||
decode_attention_fwd,
|
||||
q,
|
||||
k_buffer,
|
||||
v_buffer,
|
||||
o,
|
||||
req_to_token,
|
||||
b_req_idx,
|
||||
b_seq_len,
|
||||
attn_logits,
|
||||
num_kv_splits,
|
||||
sm_scale,
|
||||
)
|
||||
|
||||
return f, o
|
||||
|
||||
|
||||
def decode_attention_flashinfer(dtype, head_num_q, head_num_kv):
|
||||
workspace_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.int8, device="cuda")
|
||||
use_tensor_cores = should_use_tensor_core(
|
||||
kv_cache_dtype=dtype,
|
||||
num_attention_heads=head_num_q,
|
||||
num_kv_heads=head_num_kv,
|
||||
)
|
||||
flashinfer_decode_wrapper = BatchDecodeWithPagedKVCacheWrapper(
|
||||
workspace_buffer, "NHD", use_tensor_cores=use_tensor_cores
|
||||
)
|
||||
|
||||
class FlashinferAttention(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx,
|
||||
q,
|
||||
kv_data,
|
||||
batch_size,
|
||||
kv_len,
|
||||
head_num_q,
|
||||
head_num_kv,
|
||||
head_dim,
|
||||
dtype,
|
||||
warmup=10,
|
||||
):
|
||||
total_tokens = batch_size * kv_len
|
||||
kv_indptr = torch.arange(0, batch_size + 1).to(0).int() * kv_len
|
||||
kv_indices = torch.arange(0, total_tokens).to(0).int()
|
||||
kv_last_page_len = torch.full(
|
||||
(batch_size,), 1, dtype=torch.int32, device="cuda"
|
||||
)
|
||||
|
||||
flashinfer_decode_wrapper.end_forward()
|
||||
flashinfer_decode_wrapper.begin_forward(
|
||||
kv_indptr,
|
||||
kv_indices,
|
||||
kv_last_page_len,
|
||||
head_num_q,
|
||||
head_num_kv,
|
||||
head_dim,
|
||||
1,
|
||||
pos_encoding_mode="NONE",
|
||||
data_type=dtype,
|
||||
)
|
||||
|
||||
for _ in range(warmup):
|
||||
o = flashinfer_decode_wrapper.forward(
|
||||
q.contiguous().view(-1, head_num_q, head_dim), kv_data
|
||||
)
|
||||
|
||||
f = time_fwd(
|
||||
flashinfer_decode_wrapper.forward,
|
||||
q.contiguous().view(-1, head_num_q, head_dim),
|
||||
kv_data,
|
||||
)
|
||||
|
||||
return f, o
|
||||
|
||||
return FlashinferAttention
|
||||
|
||||
|
||||
def convert_to_cudnn_type(torch_type):
|
||||
if torch_type == torch.float16:
|
||||
return cudnn.data_type.HALF
|
||||
elif torch_type == torch.bfloat16:
|
||||
return cudnn.data_type.BFLOAT16
|
||||
elif torch_type == torch.float32:
|
||||
return cudnn.data_type.FLOAT
|
||||
elif torch_type == torch.int32:
|
||||
return cudnn.data_type.INT32
|
||||
elif torch_type == torch.int64:
|
||||
return cudnn.data_type.INT64
|
||||
else:
|
||||
raise ValueError("Unsupported tensor data type.")
|
||||
|
||||
|
||||
def decode_attention_cudnn(
|
||||
q, kv_data, batch_size, kv_len, head_num_q, head_num_kv, head_dim, dtype, warmup=10
|
||||
):
|
||||
# Prepare data: continuous q,k,v
|
||||
dims_q = (batch_size, head_num_q, 1, head_dim)
|
||||
strides_q = (head_num_q * head_dim, head_dim, head_num_q * head_dim, 1)
|
||||
q_gpu = q.as_strided(dims_q, strides_q)
|
||||
o_gpu = (
|
||||
torch.empty(batch_size * head_num_q * head_dim)
|
||||
.half()
|
||||
.cuda()
|
||||
.as_strided(dims_q, strides_q)
|
||||
)
|
||||
|
||||
dims_kv = (batch_size, head_num_kv, kv_len, head_dim)
|
||||
strides_kv = (
|
||||
kv_len * head_num_kv * head_dim,
|
||||
head_dim,
|
||||
head_num_kv * head_dim,
|
||||
1,
|
||||
)
|
||||
k_gpu = kv_data[0].as_strided(dims_kv, strides_kv)
|
||||
v_gpu = kv_data[1].as_strided(dims_kv, strides_kv)
|
||||
|
||||
seq_len_q_gpu = torch.full((batch_size, 1, 1, 1), 1, device="cuda")
|
||||
seq_len_kv_gpu = torch.full((batch_size, 1, 1, 1), kv_len, device="cuda")
|
||||
attn_scale = 1.0 / (head_dim**0.5)
|
||||
|
||||
# Prepare data: paged k,v
|
||||
block_size = 1
|
||||
blocks_per_batch = math.ceil(kv_len / block_size)
|
||||
# [num_blocks, head_num_kv, block_size, head_dim], num_blocks = batch_size * blocks_per_batch
|
||||
container_k_gpu = torch.cat(k_gpu.chunk(blocks_per_batch, dim=2), dim=0)
|
||||
container_v_gpu = torch.cat(v_gpu.chunk(blocks_per_batch, dim=2), dim=0)
|
||||
page_table_k_gpu = (
|
||||
torch.linspace(
|
||||
0,
|
||||
batch_size * blocks_per_batch - 1,
|
||||
batch_size * blocks_per_batch,
|
||||
device="cuda",
|
||||
dtype=torch.int32,
|
||||
)
|
||||
.reshape(blocks_per_batch, 1, batch_size, 1)
|
||||
.transpose(0, 2)
|
||||
)
|
||||
page_table_v_gpu = page_table_k_gpu.clone()
|
||||
|
||||
graph = cudnn.pygraph(
|
||||
io_data_type=convert_to_cudnn_type(dtype),
|
||||
intermediate_data_type=cudnn.data_type.FLOAT,
|
||||
compute_data_type=cudnn.data_type.FLOAT,
|
||||
)
|
||||
|
||||
q = graph.tensor_like(q_gpu)
|
||||
container_k = graph.tensor_like(container_k_gpu)
|
||||
container_v = graph.tensor_like(container_v_gpu)
|
||||
page_table_k = graph.tensor_like(page_table_k_gpu)
|
||||
page_table_v = graph.tensor_like(page_table_v_gpu)
|
||||
|
||||
seq_len_q = graph.tensor_like(seq_len_q_gpu)
|
||||
seq_len_kv = graph.tensor_like(seq_len_kv_gpu)
|
||||
|
||||
o, _ = graph.sdpa(
|
||||
name="sdpa",
|
||||
q=q,
|
||||
k=container_k, # Container K: non contiguous container with K blocks
|
||||
v=container_v, # Container V: non contiguous container with V blocks
|
||||
is_inference=True,
|
||||
attn_scale=attn_scale,
|
||||
use_causal_mask=False,
|
||||
use_padding_mask=True,
|
||||
seq_len_q=seq_len_q,
|
||||
seq_len_kv=seq_len_kv,
|
||||
paged_attention_k_table=page_table_k, # Page Table K: Tensor containing offsets to the container with K blocks
|
||||
paged_attention_v_table=page_table_v, # Page Table V: Tensor containing offsets to the container with V blocks
|
||||
paged_attention_max_seq_len_kv=kv_len, # The maximum sequence length for K caches (this is optional, but recommended)
|
||||
)
|
||||
|
||||
o.set_output(True).set_dim(dims_q).set_stride(strides_q)
|
||||
|
||||
graph.validate()
|
||||
graph.build_operation_graph()
|
||||
graph.create_execution_plans([cudnn.heur_mode.A])
|
||||
graph.check_support()
|
||||
graph.build_plans()
|
||||
|
||||
workspace = torch.empty(
|
||||
graph.get_workspace_size(), device="cuda", dtype=torch.uint8
|
||||
)
|
||||
|
||||
variant_pack = {
|
||||
q: q_gpu,
|
||||
container_k: container_k_gpu,
|
||||
container_v: container_v_gpu,
|
||||
page_table_k: page_table_k_gpu,
|
||||
page_table_v: page_table_v_gpu,
|
||||
seq_len_q: seq_len_q_gpu,
|
||||
seq_len_kv: seq_len_kv_gpu,
|
||||
o: o_gpu,
|
||||
}
|
||||
|
||||
for _ in range(warmup):
|
||||
graph.execute(variant_pack, workspace)
|
||||
|
||||
f = time_fwd(
|
||||
graph.execute,
|
||||
variant_pack,
|
||||
workspace,
|
||||
)
|
||||
|
||||
return f, o_gpu.squeeze(dim=2)
|
||||
|
||||
|
||||
def calculate_diff():
|
||||
|
||||
dtype = torch.float16
|
||||
batch_size = 64
|
||||
kv_len = 4096
|
||||
head_num_q = 64
|
||||
head_num_kv = 8
|
||||
head_dim = 128
|
||||
|
||||
q = torch.randn(batch_size, head_num_q, head_dim, dtype=dtype, device="cuda")
|
||||
kv_data = (
|
||||
torch.randn(
|
||||
batch_size * kv_len, head_num_kv, head_dim, dtype=dtype, device="cuda"
|
||||
),
|
||||
torch.randn(
|
||||
batch_size * kv_len, head_num_kv, head_dim, dtype=dtype, device="cuda"
|
||||
),
|
||||
)
|
||||
|
||||
_, output_sglang = decode_attention_sglang(
|
||||
q,
|
||||
kv_data,
|
||||
batch_size,
|
||||
kv_len,
|
||||
head_num_q,
|
||||
head_num_kv,
|
||||
head_dim,
|
||||
num_kv_splits=8,
|
||||
)
|
||||
|
||||
attn_flashinfer = decode_attention_flashinfer(dtype, head_num_q, head_num_kv).apply
|
||||
_, output_flashinfer = attn_flashinfer(
|
||||
q, kv_data, batch_size, kv_len, head_num_q, head_num_kv, head_dim, dtype
|
||||
)
|
||||
|
||||
_, output_cudnn = decode_attention_cudnn(
|
||||
q, kv_data, batch_size, kv_len, head_num_q, head_num_kv, head_dim, dtype
|
||||
)
|
||||
|
||||
print(f"SGLang output={output_sglang}")
|
||||
print(f"FlashInfer output={output_flashinfer}")
|
||||
print(f"cuDNN output={output_cudnn}")
|
||||
if torch.allclose(output_sglang, output_flashinfer, atol=1e-2, rtol=1e-2):
|
||||
print("✅ SGLang[Triton] and FlashInfer match")
|
||||
else:
|
||||
print("❌ SGLang[Triton] and FlashInfer differ")
|
||||
|
||||
if torch.allclose(output_sglang, output_cudnn, atol=1e-2, rtol=1e-2):
|
||||
print("✅ SGLang[Triton] and cuDNN match")
|
||||
else:
|
||||
print("❌ SGLang[Triton] and cuDNN differ")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
calculate_diff()
|
||||
|
||||
head_dim = 128
|
||||
dtype = torch.float16
|
||||
batch_size_range = [2**i for i in range(0, 8, 2)]
|
||||
kv_len_range = [2**i for i in range(6, 13, 1)]
|
||||
configs = list(itertools.product(batch_size_range, kv_len_range))
|
||||
|
||||
for head_num_q, head_num_kv in [[32, 32], [64, 8], [40, 8]]:
|
||||
attn_flashinfer = decode_attention_flashinfer(
|
||||
dtype, head_num_q, head_num_kv
|
||||
).apply
|
||||
for batch_size, kv_len in configs:
|
||||
q = torch.randn(
|
||||
batch_size, head_num_q, head_dim, dtype=dtype, device="cuda"
|
||||
)
|
||||
kv_data = (
|
||||
torch.randn(
|
||||
batch_size * kv_len,
|
||||
head_num_kv,
|
||||
head_dim,
|
||||
dtype=dtype,
|
||||
device="cuda",
|
||||
),
|
||||
torch.randn(
|
||||
batch_size * kv_len,
|
||||
head_num_kv,
|
||||
head_dim,
|
||||
dtype=dtype,
|
||||
device="cuda",
|
||||
),
|
||||
)
|
||||
us_cudnn, output_cudnn = decode_attention_cudnn(
|
||||
q, kv_data, batch_size, kv_len, head_num_q, head_num_kv, head_dim, dtype
|
||||
)
|
||||
us_sglang, output_sglang = decode_attention_sglang(
|
||||
q,
|
||||
kv_data,
|
||||
batch_size,
|
||||
kv_len,
|
||||
head_num_q,
|
||||
head_num_kv,
|
||||
head_dim,
|
||||
num_kv_splits=8,
|
||||
)
|
||||
us_flashinfer, _ = attn_flashinfer(
|
||||
q, kv_data, batch_size, kv_len, head_num_q, head_num_kv, head_dim, dtype
|
||||
)
|
||||
print(
|
||||
head_num_q,
|
||||
" ",
|
||||
head_num_kv,
|
||||
" ",
|
||||
batch_size,
|
||||
" ",
|
||||
kv_len,
|
||||
" ",
|
||||
us_cudnn,
|
||||
" ",
|
||||
us_sglang,
|
||||
" ",
|
||||
us_flashinfer,
|
||||
)
|
||||
218
third_party/sglang/benchmark/kernels/deepep/deepep_utils.py
vendored
Normal file
218
third_party/sglang/benchmark/kernels/deepep/deepep_utils.py
vendored
Normal file
@@ -0,0 +1,218 @@
|
||||
# ADAPTED FROM https://github.com/deepseek-ai/DeepEP/blob/main/tests/utils.py
|
||||
|
||||
import os
|
||||
import sys
|
||||
from typing import Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
|
||||
def init_dist(local_rank: int, num_local_ranks: int, args):
|
||||
ip = args.master_addr
|
||||
port = args.master_port
|
||||
num_nodes = args.nnodes
|
||||
node_rank = args.node_rank
|
||||
assert (num_local_ranks < 8 and num_nodes == 1) or num_local_ranks == 8
|
||||
|
||||
dist.init_process_group(
|
||||
backend="nccl",
|
||||
init_method=f"tcp://{ip}:{port}",
|
||||
world_size=num_nodes * num_local_ranks,
|
||||
rank=node_rank * num_local_ranks + local_rank,
|
||||
)
|
||||
torch.set_default_dtype(torch.bfloat16)
|
||||
torch.set_default_device("cuda")
|
||||
torch.cuda.set_device(local_rank)
|
||||
|
||||
return (
|
||||
dist.get_rank(),
|
||||
dist.get_world_size(),
|
||||
dist.new_group(list(range(num_local_ranks * num_nodes))),
|
||||
)
|
||||
|
||||
|
||||
def calc_diff(x: torch.Tensor, y: torch.Tensor):
|
||||
x, y = x.double() + 1, y.double() + 1
|
||||
denominator = (x * x + y * y).sum()
|
||||
sim = 2 * (x * y).sum() / denominator
|
||||
return (1 - sim).item()
|
||||
|
||||
|
||||
def per_token_cast_to_fp8(x: torch.Tensor):
|
||||
assert x.dim() == 2 and x.size(1) % 128 == 0
|
||||
m, n = x.shape
|
||||
x_view = x.view(m, -1, 128)
|
||||
x_amax = x_view.abs().float().amax(dim=2).view(m, -1).clamp(1e-4)
|
||||
return (x_view * (448.0 / x_amax.unsqueeze(2))).to(torch.float8_e4m3fn).view(
|
||||
m, n
|
||||
), (x_amax / 448.0).view(m, -1)
|
||||
|
||||
|
||||
def per_token_cast_back(x_fp8: torch.Tensor, x_scales: torch.Tensor):
|
||||
x_fp32 = x_fp8.to(torch.float32).view(x_fp8.size(0), -1, 128)
|
||||
x_scales = x_scales.view(x_fp8.size(0), -1, 1)
|
||||
return (x_fp32 * x_scales).view(x_fp8.shape).to(torch.bfloat16)
|
||||
|
||||
|
||||
def inplace_unique(x: torch.Tensor, num_slots: int):
|
||||
assert x.dim() == 2
|
||||
mask = x < 0
|
||||
x_padded = x.masked_fill(mask, num_slots)
|
||||
bin_count = torch.zeros((x.size(0), num_slots + 1), dtype=x.dtype, device=x.device)
|
||||
bin_count.scatter_add_(1, x_padded, torch.ones_like(x_padded))
|
||||
bin_count = bin_count[:, :num_slots]
|
||||
sorted_bin_count, sorted_bin_idx = torch.sort(bin_count, dim=-1, descending=True)
|
||||
sorted_bin_idx.masked_fill_(sorted_bin_count == 0, -1)
|
||||
sorted_bin_idx = torch.sort(sorted_bin_idx, descending=True, dim=-1).values
|
||||
x[:, :].fill_(-1)
|
||||
valid_len = min(num_slots, x.size(1))
|
||||
x[:, :valid_len] = sorted_bin_idx[:, :valid_len]
|
||||
|
||||
|
||||
def create_grouped_scores(
|
||||
scores: torch.Tensor, group_idx: torch.Tensor, num_groups: int
|
||||
):
|
||||
num_tokens, num_experts = scores.shape
|
||||
scores = scores.view(num_tokens, num_groups, -1)
|
||||
mask = torch.zeros((num_tokens, num_groups), dtype=torch.bool, device=scores.device)
|
||||
mask = mask.scatter_(1, group_idx, True).unsqueeze(-1).expand_as(scores)
|
||||
return (scores * mask).view(num_tokens, num_experts)
|
||||
|
||||
|
||||
def bench(fn, num_warmups: int = 20, num_tests: int = 30, post_fn=None):
|
||||
# Flush L2 cache with 256 MB data
|
||||
torch.cuda.synchronize()
|
||||
cache = torch.empty(int(256e6 // 4), dtype=torch.int, device="cuda")
|
||||
|
||||
# Warmup
|
||||
for _ in range(num_warmups):
|
||||
fn()
|
||||
|
||||
# Flush L2
|
||||
cache.zero_()
|
||||
|
||||
# Testing
|
||||
start_events = [torch.cuda.Event(enable_timing=True) for _ in range(num_tests)]
|
||||
end_events = [torch.cuda.Event(enable_timing=True) for _ in range(num_tests)]
|
||||
for i in range(num_tests):
|
||||
# Record
|
||||
start_events[i].record()
|
||||
fn()
|
||||
end_events[i].record()
|
||||
if post_fn is not None:
|
||||
post_fn()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
times = np.array(
|
||||
[s.elapsed_time(e) / 1e3 for s, e in zip(start_events, end_events)]
|
||||
)[1:]
|
||||
return np.average(times), np.min(times), np.max(times)
|
||||
|
||||
|
||||
class empty_suppress:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *_):
|
||||
pass
|
||||
|
||||
|
||||
class suppress_stdout_stderr:
|
||||
def __enter__(self):
|
||||
self.outnull_file = open(os.devnull, "w")
|
||||
self.errnull_file = open(os.devnull, "w")
|
||||
|
||||
self.old_stdout_fileno_undup = sys.stdout.fileno()
|
||||
self.old_stderr_fileno_undup = sys.stderr.fileno()
|
||||
|
||||
self.old_stdout_fileno = os.dup(sys.stdout.fileno())
|
||||
self.old_stderr_fileno = os.dup(sys.stderr.fileno())
|
||||
|
||||
self.old_stdout = sys.stdout
|
||||
self.old_stderr = sys.stderr
|
||||
|
||||
os.dup2(self.outnull_file.fileno(), self.old_stdout_fileno_undup)
|
||||
os.dup2(self.errnull_file.fileno(), self.old_stderr_fileno_undup)
|
||||
|
||||
sys.stdout = self.outnull_file
|
||||
sys.stderr = self.errnull_file
|
||||
return self
|
||||
|
||||
def __exit__(self, *_):
|
||||
sys.stdout = self.old_stdout
|
||||
sys.stderr = self.old_stderr
|
||||
|
||||
os.dup2(self.old_stdout_fileno, self.old_stdout_fileno_undup)
|
||||
os.dup2(self.old_stderr_fileno, self.old_stderr_fileno_undup)
|
||||
|
||||
os.close(self.old_stdout_fileno)
|
||||
os.close(self.old_stderr_fileno)
|
||||
|
||||
self.outnull_file.close()
|
||||
self.errnull_file.close()
|
||||
|
||||
|
||||
def bench_kineto(
|
||||
fn,
|
||||
kernel_names,
|
||||
num_tests: int = 30,
|
||||
suppress_kineto_output: bool = False,
|
||||
trace_path: Optional[str] = None,
|
||||
barrier_comm_profiling: bool = False,
|
||||
):
|
||||
# Profile
|
||||
suppress = suppress_stdout_stderr if suppress_kineto_output else empty_suppress
|
||||
with suppress():
|
||||
schedule = torch.profiler.schedule(wait=0, warmup=1, active=1, repeat=1)
|
||||
with torch.profiler.profile(
|
||||
activities=[torch.profiler.ProfilerActivity.CUDA], schedule=schedule
|
||||
) as prof:
|
||||
for i in range(2):
|
||||
# NOTES: use a large kernel and a barrier to eliminate the unbalanced CPU launch overhead
|
||||
if barrier_comm_profiling:
|
||||
lhs = torch.randn((8192, 8192), dtype=torch.float, device="cuda")
|
||||
rhs = torch.randn((8192, 8192), dtype=torch.float, device="cuda")
|
||||
lhs @ rhs
|
||||
dist.all_reduce(torch.ones(1, dtype=torch.float, device="cuda"))
|
||||
for _ in range(num_tests):
|
||||
fn()
|
||||
prof.step()
|
||||
|
||||
# Parse the profiling table
|
||||
assert isinstance(kernel_names, str) or isinstance(kernel_names, tuple)
|
||||
is_tupled = isinstance(kernel_names, tuple)
|
||||
prof_lines = (
|
||||
prof.key_averages()
|
||||
.table(sort_by="cuda_time_total", max_name_column_width=100)
|
||||
.split("\n")
|
||||
)
|
||||
kernel_names = (kernel_names,) if isinstance(kernel_names, str) else kernel_names
|
||||
assert all([isinstance(name, str) for name in kernel_names])
|
||||
for name in kernel_names:
|
||||
assert (
|
||||
sum([name in line for line in prof_lines]) == 1
|
||||
), f"Errors of the kernel {name} in the profiling table"
|
||||
|
||||
# Save chrome traces
|
||||
if trace_path is not None:
|
||||
prof.export_chrome_trace(trace_path)
|
||||
|
||||
# Return average kernel times
|
||||
units = {"ms": 1e3, "us": 1e6}
|
||||
kernel_times = []
|
||||
for name in kernel_names:
|
||||
for line in prof_lines:
|
||||
if name in line:
|
||||
time_str = line.split()[-2]
|
||||
for unit, scale in units.items():
|
||||
if unit in time_str:
|
||||
kernel_times.append(float(time_str.replace(unit, "")) / scale)
|
||||
break
|
||||
break
|
||||
return tuple(kernel_times) if is_tupled else kernel_times[0]
|
||||
|
||||
|
||||
def hash_tensor(t: torch.Tensor):
|
||||
return t.view(torch.int64).sum().item()
|
||||
480
third_party/sglang/benchmark/kernels/deepep/tuning_deepep.py
vendored
Normal file
480
third_party/sglang/benchmark/kernels/deepep/tuning_deepep.py
vendored
Normal file
@@ -0,0 +1,480 @@
|
||||
# MODIFIED FROM https://github.com/deepseek-ai/DeepEP/blob/main/tests/test_internode.py
|
||||
|
||||
"""
|
||||
Example usage:
|
||||
python tuning_deepep.py --nnodes 4 --node-rank $MY_NODE_RANK --master-addr 1.2.3.4
|
||||
Then check `deepep_tuned.json`
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from copy import deepcopy
|
||||
from pathlib import Path
|
||||
|
||||
# noinspection PyUnresolvedReferences
|
||||
import deep_ep
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from deepep_utils import (
|
||||
bench,
|
||||
calc_diff,
|
||||
create_grouped_scores,
|
||||
init_dist,
|
||||
inplace_unique,
|
||||
per_token_cast_back,
|
||||
per_token_cast_to_fp8,
|
||||
)
|
||||
|
||||
|
||||
def test_main(
|
||||
num_sms: int,
|
||||
local_rank: int,
|
||||
num_local_ranks: int,
|
||||
num_ranks: int,
|
||||
num_nodes: int,
|
||||
rank: int,
|
||||
buffer: deep_ep.Buffer,
|
||||
group: dist.ProcessGroup,
|
||||
args,
|
||||
):
|
||||
# Settings
|
||||
num_tokens, hidden, num_topk_groups, num_topk, num_experts = (
|
||||
args.num_tokens,
|
||||
args.hidden,
|
||||
min(num_nodes, 4),
|
||||
args.num_topk,
|
||||
(args.num_experts // num_ranks) * num_ranks,
|
||||
)
|
||||
assert num_experts % num_ranks == 0 and num_local_ranks == 8
|
||||
if local_rank == 0:
|
||||
print(
|
||||
f"[config] num_tokens={num_tokens}, hidden={hidden}, num_topk_groups={num_topk_groups}, num_topk={num_topk}",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
# Random data
|
||||
x = torch.ones((num_tokens, hidden), dtype=torch.bfloat16, device="cuda") * rank
|
||||
x_pure_rand = torch.randn((num_tokens, hidden), dtype=torch.bfloat16, device="cuda")
|
||||
x_e4m3 = per_token_cast_to_fp8(x)
|
||||
scores = (
|
||||
torch.randn((num_tokens, num_experts), dtype=torch.float32, device="cuda").abs()
|
||||
+ 1
|
||||
)
|
||||
group_scores = scores.view(num_tokens, num_nodes, -1).amax(dim=-1)
|
||||
group_idx = torch.topk(
|
||||
group_scores, k=num_topk_groups, dim=-1, sorted=False
|
||||
).indices
|
||||
masked_scores = create_grouped_scores(scores, group_idx, num_nodes)
|
||||
topk_idx = torch.topk(masked_scores, num_topk, dim=-1, largest=True, sorted=False)[
|
||||
1
|
||||
]
|
||||
topk_weights = (
|
||||
torch.ones((num_tokens, num_topk), dtype=torch.float32, device="cuda") * rank
|
||||
)
|
||||
topk_weights_pure_rand = torch.randn(
|
||||
(num_tokens, num_topk), dtype=torch.float32, device="cuda"
|
||||
)
|
||||
rank_idx = topk_idx // (num_experts // num_ranks)
|
||||
rank_idx.masked_fill_(topk_idx == -1, -1)
|
||||
inplace_unique(rank_idx, num_ranks)
|
||||
rdma_rank_idx = rank_idx // num_local_ranks
|
||||
rdma_rank_idx.masked_fill_(rank_idx == -1, -1)
|
||||
inplace_unique(rdma_rank_idx, num_nodes)
|
||||
|
||||
# RDMA dispatch counts
|
||||
rdma_idx = topk_idx // (num_experts // num_nodes)
|
||||
rdma_idx.masked_fill_(topk_idx == -1, -1)
|
||||
inplace_unique(rdma_idx, num_nodes)
|
||||
num_rdma_token_sent = rdma_idx.ne(-1).sum().item()
|
||||
|
||||
# Expert meta
|
||||
num_tokens_per_expert = torch.zeros((num_experts,), dtype=torch.int, device="cuda")
|
||||
for i in range(num_experts):
|
||||
num_tokens_per_expert[i] = (topk_idx == i).sum()
|
||||
gbl_num_tokens_per_expert = num_tokens_per_expert.clone()
|
||||
dist.all_reduce(gbl_num_tokens_per_expert, group=group)
|
||||
|
||||
# Rank layout meta
|
||||
num_tokens_per_rank = torch.empty((num_ranks,), dtype=torch.int, device="cuda")
|
||||
num_tokens_per_rdma_rank = torch.empty((num_nodes,), dtype=torch.int, device="cuda")
|
||||
token_idx_in_rank = torch.full(
|
||||
(num_ranks, num_tokens), -1, dtype=torch.long, device="cuda"
|
||||
)
|
||||
for i in range(num_ranks):
|
||||
num_tokens_per_rank[i] = (rank_idx == i).sum()
|
||||
token_sel = (rank_idx == i).max(dim=-1)[0]
|
||||
count = token_sel.sum().item()
|
||||
tokens = torch.sort(token_sel.to(torch.int), descending=True)[1]
|
||||
tokens[:count] = torch.sort(tokens[:count])[0]
|
||||
token_idx_in_rank[i][tokens[:count]] = torch.arange(
|
||||
count, dtype=torch.long, device="cuda"
|
||||
)
|
||||
for i in range(num_nodes):
|
||||
num_tokens_per_rdma_rank[i] = (rdma_rank_idx == i).sum()
|
||||
token_idx_in_rank = token_idx_in_rank.T.contiguous().to(torch.int)
|
||||
is_token_in_rank = token_idx_in_rank >= 0
|
||||
gbl_num_tokens_per_rank = num_tokens_per_rank.clone()
|
||||
dist.all_reduce(gbl_num_tokens_per_rank, group=group)
|
||||
|
||||
(
|
||||
ref_num_tokens_per_rank,
|
||||
ref_num_tokens_per_rdma_rank,
|
||||
ref_num_tokens_per_expert,
|
||||
ref_is_token_in_rank,
|
||||
_,
|
||||
) = buffer.get_dispatch_layout(topk_idx, num_experts)
|
||||
assert torch.allclose(ref_num_tokens_per_rank, num_tokens_per_rank)
|
||||
assert torch.allclose(ref_num_tokens_per_rdma_rank, num_tokens_per_rdma_rank)
|
||||
assert torch.allclose(ref_num_tokens_per_expert, num_tokens_per_expert)
|
||||
assert torch.allclose(ref_is_token_in_rank, is_token_in_rank)
|
||||
t = bench(lambda: buffer.get_dispatch_layout(topk_idx, num_experts))[0]
|
||||
if local_rank == 0:
|
||||
print(f"[layout] Kernel performance: {t * 1000:.3f} ms", flush=True)
|
||||
print("", flush=True)
|
||||
group.barrier()
|
||||
time.sleep(1)
|
||||
|
||||
# Config
|
||||
rdma_buffer_size, nvl_buffer_size = 128, (720 if num_ranks in (144, 160) else 512)
|
||||
config = deep_ep.Config(num_sms, 8, nvl_buffer_size, 16, rdma_buffer_size)
|
||||
|
||||
# Test dispatch
|
||||
# noinspection PyShadowingNames
|
||||
def check_data(check_x, recv_gbl_rank_prefix_sum):
|
||||
assert torch.allclose(check_x.amin(dim=1), check_x.amax(dim=1))
|
||||
check_start = 0
|
||||
for i in range(num_ranks):
|
||||
check_end = recv_gbl_rank_prefix_sum[i].item()
|
||||
assert (check_x[check_start:check_end, :].int() - i).sum().item() == 0
|
||||
check_start = check_end
|
||||
|
||||
for previous_mode in (False, True):
|
||||
for async_mode in (False, True):
|
||||
for current_x in (x_pure_rand, x, x_e4m3):
|
||||
for with_topk in (False, True):
|
||||
if local_rank == 0:
|
||||
print(
|
||||
f'[testing] Running with {"FP8" if isinstance(current_x, tuple) else "BF16"}, {"with" if with_topk else "without"} top-k (async={async_mode}, previous={previous_mode}) ...',
|
||||
flush=True,
|
||||
end="",
|
||||
)
|
||||
dispatch_args = {
|
||||
"x": current_x,
|
||||
"num_tokens_per_rank": num_tokens_per_rank,
|
||||
"num_tokens_per_rdma_rank": num_tokens_per_rdma_rank,
|
||||
"is_token_in_rank": is_token_in_rank,
|
||||
"num_tokens_per_expert": num_tokens_per_expert,
|
||||
"config": config,
|
||||
"async_finish": async_mode,
|
||||
}
|
||||
if with_topk:
|
||||
dispatch_args.update(
|
||||
{
|
||||
"topk_idx": topk_idx,
|
||||
"topk_weights": (
|
||||
topk_weights_pure_rand
|
||||
if current_x is x_pure_rand
|
||||
else topk_weights
|
||||
),
|
||||
}
|
||||
)
|
||||
if previous_mode:
|
||||
dispatch_args.update({"previous_event": buffer.capture()})
|
||||
(
|
||||
recv_x,
|
||||
recv_topk_idx,
|
||||
recv_topk_weights,
|
||||
recv_num_tokens_per_expert_list,
|
||||
handle,
|
||||
event,
|
||||
) = buffer.dispatch(**dispatch_args)
|
||||
event.current_stream_wait() if async_mode else ()
|
||||
recv_x = (
|
||||
per_token_cast_back(*recv_x)
|
||||
if isinstance(recv_x, tuple)
|
||||
else recv_x
|
||||
)
|
||||
|
||||
# Checks
|
||||
recv_gbl_rank_prefix_sum = handle[-4]
|
||||
assert gbl_num_tokens_per_rank[rank].item() == recv_x.size(
|
||||
0
|
||||
), f"{gbl_num_tokens_per_rank[rank].item()} != {recv_x.size(0)}"
|
||||
assert (
|
||||
gbl_num_tokens_per_expert.view(num_ranks, -1)[rank].tolist()
|
||||
== recv_num_tokens_per_expert_list
|
||||
)
|
||||
if current_x is not x_pure_rand:
|
||||
check_data(recv_x, recv_gbl_rank_prefix_sum)
|
||||
if with_topk:
|
||||
# Check `topk_idx`
|
||||
assert (
|
||||
recv_topk_idx.eq(-1)
|
||||
| (
|
||||
(recv_topk_idx >= 0)
|
||||
& (recv_topk_idx < (num_experts // num_ranks))
|
||||
)
|
||||
).sum().item() == recv_topk_idx.numel()
|
||||
for i, count in enumerate(recv_num_tokens_per_expert_list):
|
||||
assert recv_topk_idx.eq(i).sum().item() == count
|
||||
|
||||
# Check `topk_weights`
|
||||
if current_x is not x_pure_rand:
|
||||
recv_topk_weights[recv_topk_idx.eq(-1)] = (
|
||||
recv_topk_weights.amax(dim=1, keepdim=True).expand_as(
|
||||
recv_topk_weights
|
||||
)[recv_topk_idx.eq(-1)]
|
||||
)
|
||||
check_data(recv_topk_weights, recv_gbl_rank_prefix_sum)
|
||||
|
||||
# Test cached dispatch (must without top-k staffs)
|
||||
if not with_topk:
|
||||
dispatch_args = {
|
||||
"x": current_x,
|
||||
"handle": handle,
|
||||
"config": config,
|
||||
"async_finish": async_mode,
|
||||
}
|
||||
if previous_mode:
|
||||
dispatch_args.update({"previous_event": buffer.capture()})
|
||||
recv_x, _, _, _, _, event = buffer.dispatch(**dispatch_args)
|
||||
event.current_stream_wait() if async_mode else ()
|
||||
recv_x = (
|
||||
per_token_cast_back(*recv_x)
|
||||
if isinstance(recv_x, tuple)
|
||||
else recv_x
|
||||
)
|
||||
if current_x is not x_pure_rand:
|
||||
check_data(recv_x, recv_gbl_rank_prefix_sum)
|
||||
|
||||
# Test combine
|
||||
combine_args = {
|
||||
"x": recv_x,
|
||||
"handle": handle,
|
||||
"config": config,
|
||||
"async_finish": async_mode,
|
||||
}
|
||||
if with_topk:
|
||||
combine_args.update({"topk_weights": recv_topk_weights})
|
||||
if previous_mode:
|
||||
combine_args.update({"previous_event": buffer.capture()})
|
||||
combined_x, combined_topk_weights, event = buffer.combine(
|
||||
**combine_args
|
||||
)
|
||||
event.current_stream_wait() if async_mode else ()
|
||||
check_x = combined_x.float() / is_token_in_rank.sum(
|
||||
dim=1
|
||||
).unsqueeze(1)
|
||||
ref_x = x_pure_rand if current_x is x_pure_rand else x
|
||||
assert calc_diff(check_x, ref_x) < 5e-6
|
||||
if with_topk:
|
||||
check_topk_weights = (
|
||||
combined_topk_weights
|
||||
if (current_x is x_pure_rand)
|
||||
else (
|
||||
combined_topk_weights
|
||||
/ is_token_in_rank.sum(dim=1).unsqueeze(1)
|
||||
)
|
||||
)
|
||||
ref_topk_weights = (
|
||||
topk_weights_pure_rand
|
||||
if current_x is x_pure_rand
|
||||
else topk_weights
|
||||
)
|
||||
assert calc_diff(check_topk_weights, ref_topk_weights) < 1e-9
|
||||
|
||||
# For later tuning
|
||||
dispatch_bf16_rdma_send_bytes = num_rdma_token_sent * hidden * 2
|
||||
dispatch_bf16_nvl_recv_bytes = recv_x.numel() * 2
|
||||
combine_bf16_nvl_send_bytes = dispatch_bf16_nvl_recv_bytes
|
||||
combine_bf16_rdma_recv_bytes = dispatch_bf16_rdma_send_bytes
|
||||
|
||||
if local_rank == 0:
|
||||
print(" passed", flush=True)
|
||||
if local_rank == 0:
|
||||
print("", flush=True)
|
||||
|
||||
output_data = {}
|
||||
|
||||
# Tune dispatch performance
|
||||
best_dispatch_results = None
|
||||
fp8_factor = (1 + 4 / 128) / 2
|
||||
for current_x in (x_e4m3, x):
|
||||
best_time, best_results = 1e10, None
|
||||
rdma_send_bytes = (
|
||||
(dispatch_bf16_rdma_send_bytes * fp8_factor)
|
||||
if isinstance(current_x, tuple)
|
||||
else dispatch_bf16_rdma_send_bytes
|
||||
)
|
||||
nvl_recv_bytes = (
|
||||
(dispatch_bf16_nvl_recv_bytes * fp8_factor)
|
||||
if isinstance(current_x, tuple)
|
||||
else dispatch_bf16_nvl_recv_bytes
|
||||
)
|
||||
for nvl_chunk_size in range(4, 33, 4):
|
||||
for rdma_chunk_size in range(4, 33, 4):
|
||||
config_kwargs = {
|
||||
"num_sms": num_sms,
|
||||
"num_max_nvl_chunked_send_tokens": nvl_chunk_size,
|
||||
"num_max_nvl_chunked_recv_tokens": nvl_buffer_size,
|
||||
"num_max_rdma_chunked_send_tokens": rdma_chunk_size,
|
||||
"num_max_rdma_chunked_recv_tokens": rdma_buffer_size,
|
||||
}
|
||||
config = deep_ep.Config(**config_kwargs)
|
||||
tune_args = {"x": current_x, "handle": handle, "config": config}
|
||||
t = bench(lambda: buffer.dispatch(**tune_args))[0]
|
||||
if t < best_time:
|
||||
best_time, best_results = t, (
|
||||
num_sms,
|
||||
nvl_chunk_size,
|
||||
rdma_chunk_size,
|
||||
config_kwargs,
|
||||
)
|
||||
if local_rank == 0:
|
||||
print(
|
||||
f"[tuning] SMs {num_sms}, NVL chunk {nvl_chunk_size}, RDMA chunk {rdma_chunk_size}: {rdma_send_bytes / 1e9 / t:.2f} GB/s (RDMA), {nvl_recv_bytes / 1e9 / t:.2f} GB/s (NVL) ",
|
||||
flush=True,
|
||||
)
|
||||
if local_rank == 0:
|
||||
print(
|
||||
f'[tuning] Best dispatch ({"FP8" if isinstance(current_x, tuple) else "BF16"}): SMs {best_results[0]}, NVL chunk {best_results[1]}, RDMA chunk {best_results[2]}: {rdma_send_bytes / 1e9 / best_time:.2f} GB/s (RDMA), {nvl_recv_bytes / 1e9 / best_time:.2f} GB/s (NVL)',
|
||||
flush=True,
|
||||
)
|
||||
print("", flush=True)
|
||||
is_fp8 = isinstance(current_x, tuple)
|
||||
if is_fp8:
|
||||
output_data["normal_dispatch"] = deepcopy(best_results[3])
|
||||
|
||||
if isinstance(current_x, tuple):
|
||||
# Gather FP8 the best config from rank 0
|
||||
best_dispatch_results = torch.tensor(
|
||||
[best_results[0], best_results[1], best_results[2]],
|
||||
dtype=torch.int32,
|
||||
device="cuda",
|
||||
)
|
||||
all_best_fp8_results_list = [
|
||||
torch.zeros_like(best_dispatch_results)
|
||||
for _ in range(torch.distributed.get_world_size())
|
||||
]
|
||||
dist.all_gather(
|
||||
all_best_fp8_results_list, best_dispatch_results, group=group
|
||||
)
|
||||
best_dispatch_results = all_best_fp8_results_list[0].tolist()
|
||||
dispatch_config = deep_ep.Config(
|
||||
best_dispatch_results[0],
|
||||
best_dispatch_results[1],
|
||||
nvl_buffer_size,
|
||||
best_dispatch_results[2],
|
||||
rdma_buffer_size,
|
||||
)
|
||||
|
||||
dispatch_args = {
|
||||
"x": x,
|
||||
"num_tokens_per_rank": num_tokens_per_rank,
|
||||
"num_tokens_per_rdma_rank": num_tokens_per_rdma_rank,
|
||||
"is_token_in_rank": is_token_in_rank,
|
||||
"num_tokens_per_expert": num_tokens_per_expert,
|
||||
"config": dispatch_config if dispatch_config is not None else config,
|
||||
}
|
||||
recv_x, _, _, _, handle, _ = buffer.dispatch(**dispatch_args)
|
||||
|
||||
# Tune combine performance
|
||||
best_time, best_results = 1e10, None
|
||||
for nvl_chunk_size in range(1, 8, 1):
|
||||
for rdma_chunk_size in range(12 if num_nodes == 2 else 8, 33, 4):
|
||||
config_kwargs = {
|
||||
"num_sms": num_sms,
|
||||
"num_max_nvl_chunked_send_tokens": nvl_chunk_size,
|
||||
"num_max_nvl_chunked_recv_tokens": nvl_buffer_size,
|
||||
"num_max_rdma_chunked_send_tokens": rdma_chunk_size,
|
||||
"num_max_rdma_chunked_recv_tokens": rdma_buffer_size,
|
||||
}
|
||||
config = deep_ep.Config(**config_kwargs)
|
||||
tune_args = {"x": recv_x, "handle": handle, "config": config}
|
||||
t = bench(lambda: buffer.combine(**tune_args))[0]
|
||||
if local_rank == 0:
|
||||
print(
|
||||
f"[tuning] SMs {num_sms}, NVL chunk {nvl_chunk_size}, RDMA chunk {rdma_chunk_size}: {combine_bf16_rdma_recv_bytes / 1e9 / t:.2f} GB/s (RDMA), {combine_bf16_nvl_send_bytes / 1e9 / t:.2f} GB/s (NVL) ",
|
||||
flush=True,
|
||||
)
|
||||
if t < best_time:
|
||||
best_time, best_results = t, (
|
||||
num_sms,
|
||||
nvl_chunk_size,
|
||||
rdma_chunk_size,
|
||||
config_kwargs,
|
||||
)
|
||||
|
||||
if local_rank == 0:
|
||||
print(
|
||||
f"[tuning] Best combine: SMs {best_results[0]}, NVL chunk {best_results[1]}, RDMA chunk {best_results[2]}: {combine_bf16_rdma_recv_bytes / 1e9 / best_time:.2f} GB/s (RDMA), {combine_bf16_nvl_send_bytes / 1e9 / best_time:.2f} GB/s (NVL)",
|
||||
flush=True,
|
||||
)
|
||||
print("", flush=True)
|
||||
output_data["normal_combine"] = deepcopy(best_results[3])
|
||||
|
||||
if rank == 0 and local_rank == 0:
|
||||
_write_output(args, output_data)
|
||||
|
||||
|
||||
def _write_output(args, output_data):
|
||||
text = json.dumps(output_data, indent=4)
|
||||
output_path = args.output_path
|
||||
print(f"Write to {output_path} with {text}")
|
||||
Path(output_path).write_text(text)
|
||||
|
||||
|
||||
# noinspection PyUnboundLocalVariable
|
||||
def test_loop(local_rank: int, num_local_ranks: int, args):
|
||||
num_nodes = args.nnodes
|
||||
rank, num_ranks, group = init_dist(local_rank, num_local_ranks, args)
|
||||
|
||||
num_sms = args.num_sms
|
||||
num_qps_per_rank = num_sms // 2
|
||||
|
||||
buffer = deep_ep.Buffer(
|
||||
group,
|
||||
int(1e9),
|
||||
int(1e9),
|
||||
low_latency_mode=False,
|
||||
num_qps_per_rank=num_qps_per_rank,
|
||||
)
|
||||
assert num_local_ranks == 8 and num_ranks > 8
|
||||
torch.manual_seed(rank)
|
||||
|
||||
for i in (num_sms,):
|
||||
test_main(
|
||||
i,
|
||||
local_rank,
|
||||
num_local_ranks,
|
||||
num_ranks,
|
||||
num_nodes,
|
||||
rank,
|
||||
buffer,
|
||||
group,
|
||||
args,
|
||||
)
|
||||
if local_rank == 0:
|
||||
print("", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--num-sms", type=int, default=24)
|
||||
parser.add_argument("--num-tokens", type=int, default=4096)
|
||||
parser.add_argument("--hidden", type=int, default=7168)
|
||||
parser.add_argument("--num-topk", type=int, default=8)
|
||||
parser.add_argument("--num-experts", type=int, default=256)
|
||||
parser.add_argument("--output-path", type=str, default="deepep_tuned.json")
|
||||
parser.add_argument("--nnodes", type=int, default=1)
|
||||
parser.add_argument("--node-rank", type=int, default=0)
|
||||
parser.add_argument("--master-addr", type=str, default="127.0.0.1")
|
||||
parser.add_argument("--master-port", type=int, default=8361)
|
||||
args = parser.parse_args()
|
||||
print(f"Start system with {args=}")
|
||||
|
||||
num_processes = 8
|
||||
torch.multiprocessing.spawn(
|
||||
test_loop, args=(num_processes, args), nprocs=num_processes
|
||||
)
|
||||
19
third_party/sglang/benchmark/kernels/deepseek/README.md
vendored
Normal file
19
third_party/sglang/benchmark/kernels/deepseek/README.md
vendored
Normal file
@@ -0,0 +1,19 @@
|
||||
## DeepSeek kernels benchmark
|
||||
|
||||
|
||||
### Prerequisites
|
||||
- You should install [DeepGemm](https://github.com/deepseek-ai/DeepGEMM) from source before run `benchmark_deepgemm_fp8_gemm.py` and `benchmark_deepgemm_fp8_group_gemm.py`.
|
||||
|
||||
### Benchmark
|
||||
- `benchmark_deepgemm_fp8_gemm.py`
|
||||
```bash
|
||||
python benchmark_deepgemm_fp8_gemm.py --run_correctness --tp_size 1
|
||||
```
|
||||
|
||||
- `benchmark_deepgemm_fp8_group_gemm.py`
|
||||
```bash
|
||||
python benchmark_deepgemm_fp8_group_gemm.py --run_correctness --tp_size 1
|
||||
```
|
||||
|
||||
- You can use the `--run_correctness` parameter to verify all kernels results's correctness.
|
||||
- You can use the `--tp_size` parameter to benchmark all FP8 w8a8 block-wise matrix multiplications involved in DeepSeek V3/R1 under the current tensor parallelism (TP) setting. This benchmark compares DeepSeek's open-source [DeepGemm](https://github.com/deepseek-ai/DeepGEMM) implementation with SGLang's and VLLM Triton implementation.
|
||||
250
third_party/sglang/benchmark/kernels/deepseek/benchmark_deepgemm_dsv3_router_gemm_blackwell.py
vendored
Normal file
250
third_party/sglang/benchmark/kernels/deepseek/benchmark_deepgemm_dsv3_router_gemm_blackwell.py
vendored
Normal file
@@ -0,0 +1,250 @@
|
||||
import argparse
|
||||
import os
|
||||
from typing import List
|
||||
|
||||
import torch
|
||||
import triton
|
||||
from flashinfer.gemm import mm_M1_16_K7168_N256
|
||||
from sgl_kernel import dsv3_router_gemm
|
||||
|
||||
N = 256
|
||||
K = 7168
|
||||
|
||||
|
||||
def create_benchmark_configs(tp_sizes: List[int]):
|
||||
configs = []
|
||||
for tp_size in tp_sizes:
|
||||
for m in range(1, 17):
|
||||
configs.append((m, N, K, tp_size))
|
||||
return configs
|
||||
|
||||
|
||||
def dsv3_router_gemm_flashinfer(
|
||||
hidden_states: torch.Tensor,
|
||||
router_weights: torch.Tensor,
|
||||
):
|
||||
"""Flashinfer implementation of dsv3 router gemm"""
|
||||
output = torch.empty(
|
||||
hidden_states.shape[0],
|
||||
router_weights.shape[0],
|
||||
device="cuda",
|
||||
dtype=torch.float32,
|
||||
)
|
||||
mm_M1_16_K7168_N256(
|
||||
hidden_states, router_weights.t(), output, launch_with_pdl=args.use_pdl
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
def dsv3_router_gemm_sgl(
|
||||
hidden_states: torch.Tensor,
|
||||
router_weights: torch.Tensor,
|
||||
):
|
||||
"""SGLang implementation of dsv3 router gemm"""
|
||||
output = dsv3_router_gemm(
|
||||
hidden_states,
|
||||
router_weights,
|
||||
out_dtype=torch.float32,
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
def check_accuracy(a, b, atol, rtol, percent):
|
||||
"""Unified accuracy checking function with detailed error reporting."""
|
||||
if not torch.isfinite(a).all():
|
||||
print("Non-finite values in reference output")
|
||||
return False
|
||||
if not torch.isfinite(b).all():
|
||||
print("Non-finite values in actual output")
|
||||
return False
|
||||
assert a.shape == b.shape, f"Shape mismatch: {a.shape} vs {b.shape}"
|
||||
|
||||
close = torch.isclose(a, b, atol=atol, rtol=rtol)
|
||||
match_ratio = close.float().mean()
|
||||
if match_ratio >= percent:
|
||||
return True
|
||||
|
||||
mismatch_percent = 1.0 - match_ratio.item()
|
||||
if mismatch_percent > 1 - percent:
|
||||
print(
|
||||
f"Mismatch percentage is {mismatch_percent:.4f} for rtol {rtol} "
|
||||
f"(threshold: {1 - percent:.4f})"
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def calculate_diff(m: int, n: int, k: int):
|
||||
hidden_states = torch.randn((m, k), device="cuda", dtype=torch.bfloat16)
|
||||
router_weights = torch.randn((n, k), device="cuda", dtype=torch.bfloat16)
|
||||
|
||||
out_flashinfer = dsv3_router_gemm_flashinfer(
|
||||
hidden_states.clone(memory_format=torch.contiguous_format),
|
||||
router_weights.clone(memory_format=torch.contiguous_format),
|
||||
)
|
||||
|
||||
out_sgl = dsv3_router_gemm_sgl(
|
||||
hidden_states.clone(memory_format=torch.contiguous_format),
|
||||
router_weights.clone(memory_format=torch.contiguous_format),
|
||||
)
|
||||
|
||||
print(f"Shape m={m}, n={n}, k={k}:")
|
||||
print(f"Using PDL={args.use_pdl}")
|
||||
print(f"Flashinfer output: {out_flashinfer[0, 0:5]}")
|
||||
print(f"SGLang output: {out_sgl[0, 0:5]}")
|
||||
|
||||
flashinfer_sgl_match = check_accuracy(out_flashinfer, out_sgl, 0.1, 0.6, 0.95)
|
||||
print("Correctness check:")
|
||||
print(f" - Flashinfer vs SGLang: {'✅' if flashinfer_sgl_match else '❌'}")
|
||||
|
||||
|
||||
def _benchmark(m, n, k, tp_size, provider):
|
||||
print(f"Shape (m={m}, n={n}, k={k}, tp={tp_size}), Provider: {provider}")
|
||||
hidden_states = torch.randn(
|
||||
(m, k), device="cuda", dtype=torch.bfloat16
|
||||
).contiguous()
|
||||
router_weights = torch.randn(
|
||||
(n, k), device="cuda", dtype=torch.bfloat16
|
||||
).contiguous()
|
||||
|
||||
quantiles = [0.5, 0.2, 0.8]
|
||||
|
||||
if provider == "sglang":
|
||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
||||
lambda: dsv3_router_gemm_sgl(
|
||||
hidden_states.clone(memory_format=torch.contiguous_format),
|
||||
router_weights.clone(memory_format=torch.contiguous_format),
|
||||
),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
elif provider == "flashinfer":
|
||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
||||
lambda: dsv3_router_gemm_flashinfer(
|
||||
hidden_states.clone(memory_format=torch.contiguous_format),
|
||||
router_weights.clone(memory_format=torch.contiguous_format),
|
||||
),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
|
||||
# Calculate TFLOPS
|
||||
flops = 2 * m * n * k # multiply-adds
|
||||
tflops = flops / (ms * 1e-3) / 1e12
|
||||
|
||||
# Print shape-specific results with TFLOPS
|
||||
print(f"Time: {ms*1000:.2f} us, TFLOPS: {tflops:.2f}")
|
||||
return ms, max_ms, min_ms
|
||||
|
||||
|
||||
def get_benchmark_plot_friendly(tp_sizes):
|
||||
all_configs = create_benchmark_configs(tp_sizes)
|
||||
x_vals = list(range(len(all_configs)))
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["cfg_id"],
|
||||
x_vals=x_vals,
|
||||
line_arg="provider",
|
||||
line_vals=["sglang", "flashinfer"],
|
||||
line_names=["SGLang", "Flashinfer"],
|
||||
styles=[("blue", "-"), ("red", "-")],
|
||||
ylabel="us",
|
||||
plot_name=f"fp8-gemm-performance-comparison-tp-{"-".join(str(tp) for tp in tp_sizes)}",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(cfg_id, provider):
|
||||
m, n, k, tp_size = all_configs[cfg_id]
|
||||
ms, min_ms, max_ms = _benchmark(m, n, k, tp_size, provider)
|
||||
return ms * 1000, max_ms * 1000, min_ms * 1000 # convert to ms
|
||||
|
||||
return benchmark
|
||||
|
||||
|
||||
def get_benchmark(tp_sizes):
|
||||
all_configs = create_benchmark_configs(tp_sizes)
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=[
|
||||
"m",
|
||||
"n",
|
||||
"k",
|
||||
"tp_size",
|
||||
],
|
||||
x_vals=[list(config) for config in all_configs],
|
||||
line_arg="provider",
|
||||
line_vals=["sglang", "flashinfer"],
|
||||
line_names=["SGLang", "Flashinfer"],
|
||||
styles=[("blue", "-"), ("red", "-")],
|
||||
ylabel="us",
|
||||
plot_name=f"fp8-gemm-performance-comparison-tp-{"-".join(str(tp) for tp in tp_sizes)}",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(m, n, k, tp_size, provider):
|
||||
ms, min_ms, max_ms = _benchmark(m, n, k, tp_size, provider)
|
||||
return ms * 1000, max_ms * 1000, min_ms * 1000 # convert to ms
|
||||
|
||||
return benchmark
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] != 10:
|
||||
print("Skipping benchmark because the device is not supported")
|
||||
exit(0)
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--save-path",
|
||||
type=str,
|
||||
default="./configs/benchmark_ops/dsv3_router_gemm/",
|
||||
help="Path to save dsv3 router gemm benchmark results",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--run-correctness",
|
||||
action="store_true",
|
||||
default=True,
|
||||
help="Whether to run correctness test",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tp-sizes",
|
||||
type=int,
|
||||
nargs="+",
|
||||
default=[1],
|
||||
help="List of tensor parallelism sizes to benchmark",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--plot-friendly",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Plot x axis as the config index instead of the m",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-pdl",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Use PDL if true.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
# Set random seed for reproducibility
|
||||
torch.manual_seed(0)
|
||||
torch.cuda.manual_seed(0)
|
||||
|
||||
if args.use_pdl:
|
||||
os.environ["TRTLLM_ENABLE_PDL"] = "1"
|
||||
|
||||
# Run correctness tests on a few examples
|
||||
if args.run_correctness:
|
||||
print("Running correctness tests...")
|
||||
for m, n, k, _ in create_benchmark_configs(args.tp_sizes):
|
||||
calculate_diff(m, n, k)
|
||||
|
||||
# Get the benchmark function with the specified tp_size
|
||||
benchmark = (
|
||||
get_benchmark_plot_friendly(args.tp_sizes)
|
||||
if args.plot_friendly
|
||||
else get_benchmark(args.tp_sizes)
|
||||
)
|
||||
|
||||
print(f"Running performance benchmark for TP sizes = {args.tp_sizes}...")
|
||||
benchmark.run(print_data=True, save_path=args.save_path)
|
||||
402
third_party/sglang/benchmark/kernels/deepseek/benchmark_deepgemm_fp8_gemm.py
vendored
Normal file
402
third_party/sglang/benchmark/kernels/deepseek/benchmark_deepgemm_fp8_gemm.py
vendored
Normal file
@@ -0,0 +1,402 @@
|
||||
from typing import Tuple
|
||||
|
||||
import deep_gemm
|
||||
import tilelang
|
||||
import tilelang.language as T
|
||||
import torch
|
||||
import triton
|
||||
from deep_gemm import ceil_div
|
||||
from deep_gemm.utils.layout import get_mn_major_tma_aligned_tensor
|
||||
from vllm.model_executor.layers.quantization.utils.fp8_utils import (
|
||||
w8a8_block_fp8_matmul as vllm_w8a8_block_fp8_matmul,
|
||||
)
|
||||
|
||||
from sglang.benchmark.bench_utils import run_bench
|
||||
from sglang.srt.layers.quantization.fp8_kernel import (
|
||||
w8a8_block_fp8_matmul_deepgemm as w8a8_block_fp8_matmul,
|
||||
)
|
||||
|
||||
|
||||
# Adapted from https://github.com/tile-ai/tilelang/blob/a8cfdce92795cb861c9033573534653ee040b5ed/examples/deepseek_deepgemm/example_deepgemm_fp8_2xAcc.py#L1
|
||||
def tl_gemm(
|
||||
M,
|
||||
N,
|
||||
K,
|
||||
in_dtype,
|
||||
out_dtype,
|
||||
accum_dtype,
|
||||
):
|
||||
assert in_dtype in [
|
||||
"e4m3_float8",
|
||||
], "Currently only e4m3_float8 is supported"
|
||||
assert out_dtype in [
|
||||
"bfloat16",
|
||||
"float16",
|
||||
], "Currently only bfloat16 and float16 are supported"
|
||||
|
||||
TILE_SIZE = (128, 128, 128)
|
||||
block_M = TILE_SIZE[0]
|
||||
block_N = TILE_SIZE[1]
|
||||
block_K = TILE_SIZE[2]
|
||||
|
||||
A_shape = (M, K)
|
||||
Scales_A_shape = (M, T.ceildiv(K, block_K))
|
||||
B_shape = (N, K)
|
||||
Scales_B_shape = (T.ceildiv(N, block_N), T.ceildiv(K, block_K))
|
||||
A_shared_shape = (block_M, block_K)
|
||||
B_shared_shape = (block_N, block_K)
|
||||
C_shared_shape = (block_M, block_N)
|
||||
|
||||
@T.prim_func
|
||||
def main(
|
||||
A: T.Buffer(A_shape, in_dtype),
|
||||
scales_a: T.Buffer(Scales_A_shape, "float32"),
|
||||
B: T.Buffer(B_shape, in_dtype),
|
||||
scales_b: T.Buffer(Scales_B_shape, "float32"),
|
||||
C: T.Buffer((M, N), out_dtype),
|
||||
):
|
||||
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (
|
||||
bx,
|
||||
by,
|
||||
):
|
||||
|
||||
A_shared = T.alloc_shared(A_shared_shape, in_dtype)
|
||||
B_shared = T.alloc_shared(B_shared_shape, in_dtype)
|
||||
C_shared = T.alloc_shared(C_shared_shape, out_dtype)
|
||||
Scale_C_shared = T.alloc_shared((block_M), "float32")
|
||||
C_local = T.alloc_fragment(C_shared_shape, accum_dtype)
|
||||
C_local_accum = T.alloc_fragment(C_shared_shape, accum_dtype)
|
||||
|
||||
# Improve L2 Cache
|
||||
T.use_swizzle(panel_size=10)
|
||||
|
||||
T.clear(C_local)
|
||||
T.clear(C_local_accum)
|
||||
K_iters = T.ceildiv(K, block_K)
|
||||
for k in T.Pipelined(K_iters, num_stages=4):
|
||||
# Load A into shared memory
|
||||
T.copy(A[by * block_M, k * block_K], A_shared)
|
||||
# Load B into shared memory
|
||||
T.copy(B[bx * block_N, k * block_K], B_shared)
|
||||
# Load scale into shared memory
|
||||
Scale_B = scales_b[bx, k]
|
||||
for i in T.Parallel(block_M):
|
||||
Scale_C_shared[i] = scales_a[by * block_M + i, k] * Scale_B
|
||||
|
||||
T.gemm(A_shared, B_shared, C_local, transpose_B=True)
|
||||
# Promote to enable 2xAcc
|
||||
for i, j in T.Parallel(block_M, block_N):
|
||||
C_local_accum[i, j] += C_local[i, j] * Scale_C_shared[i]
|
||||
T.clear(C_local)
|
||||
# TMA store
|
||||
T.copy(C_local_accum, C_shared)
|
||||
T.copy(C_shared, C[by * block_M, bx * block_N])
|
||||
|
||||
return main
|
||||
|
||||
|
||||
def per_token_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
assert x.dim() == 2 and x.size(1) % 128 == 0
|
||||
m, n = x.shape
|
||||
x_view = x.view(m, -1, 128)
|
||||
x_amax = x_view.abs().float().amax(dim=2).view(m, -1).clamp(1e-4)
|
||||
return (x_view * (448.0 / x_amax.unsqueeze(2))).to(torch.float8_e4m3fn).view(
|
||||
m, n
|
||||
), (x_amax / 448.0).view(m, -1)
|
||||
|
||||
|
||||
def per_block_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
assert x.dim() == 2
|
||||
m, n = x.shape
|
||||
x_padded = torch.zeros(
|
||||
(ceil_div(m, 128) * 128, ceil_div(n, 128) * 128), dtype=x.dtype, device=x.device
|
||||
)
|
||||
x_padded[:m, :n] = x
|
||||
x_view = x_padded.view(-1, 128, x_padded.size(1) // 128, 128)
|
||||
x_amax = x_view.abs().float().amax(dim=(1, 3), keepdim=True).clamp(1e-4)
|
||||
x_scaled = (x_view * (448.0 / x_amax)).to(torch.float8_e4m3fn)
|
||||
return x_scaled.view_as(x_padded)[:m, :n].contiguous(), (x_amax / 448.0).view(
|
||||
x_view.size(0), x_view.size(2)
|
||||
)
|
||||
|
||||
|
||||
def fp8_gemm_deepgemm(
|
||||
x_fp8: torch.Tensor,
|
||||
x_scale: torch.Tensor,
|
||||
y_fp8: torch.Tensor,
|
||||
y_scale: torch.Tensor,
|
||||
m: int,
|
||||
n: int,
|
||||
k: int,
|
||||
):
|
||||
"""DeepGEMM implementation of FP8 GEMM"""
|
||||
out = torch.empty((m, n), device="cuda", dtype=torch.bfloat16)
|
||||
|
||||
# Run DeepGEMM kernel
|
||||
deep_gemm.fp8_gemm_nt((x_fp8, x_scale), (y_fp8, y_scale), out)
|
||||
return out
|
||||
|
||||
|
||||
def fp8_gemm_sglang(
|
||||
x_fp8: torch.Tensor,
|
||||
x_scale: torch.Tensor,
|
||||
y_fp8: torch.Tensor,
|
||||
y_scale: torch.Tensor,
|
||||
m: int,
|
||||
n: int,
|
||||
k: int,
|
||||
):
|
||||
"""SGLang implementation of FP8 GEMM"""
|
||||
block_size = [128, 128] # Matches the block size in per_block_cast_to_fp8
|
||||
|
||||
# Run SGLang kernel
|
||||
out = w8a8_block_fp8_matmul(
|
||||
x_fp8, y_fp8, x_scale, y_scale, block_size, torch.bfloat16
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def fp8_gemm_vllm(
|
||||
x_fp8: torch.Tensor,
|
||||
x_scale: torch.Tensor,
|
||||
y_fp8: torch.Tensor,
|
||||
y_scale: torch.Tensor,
|
||||
m: int,
|
||||
n: int,
|
||||
k: int,
|
||||
):
|
||||
"""vLLM implementation of FP8 GEMM"""
|
||||
block_size = [128, 128] # Matches the block size in per_block_cast_to_fp8
|
||||
|
||||
# Run vLLM kernel
|
||||
out = vllm_w8a8_block_fp8_matmul(
|
||||
x_fp8, y_fp8, x_scale, y_scale, block_size, torch.bfloat16
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def calculate_diff(m: int, n: int, k: int):
|
||||
x = torch.randn((m, k), device="cuda", dtype=torch.bfloat16)
|
||||
y = torch.randn((n, k), device="cuda", dtype=torch.bfloat16)
|
||||
|
||||
x_fp8, x_scale = per_token_cast_to_fp8(x.clone())
|
||||
y_fp8, y_scale = per_block_cast_to_fp8(y.clone())
|
||||
x_scale_col_major = get_mn_major_tma_aligned_tensor(x_scale.clone())
|
||||
|
||||
out_deepgemm = fp8_gemm_deepgemm(
|
||||
x_fp8.clone(),
|
||||
x_scale_col_major.clone(),
|
||||
y_fp8.clone(),
|
||||
y_scale.clone(),
|
||||
m,
|
||||
n,
|
||||
k,
|
||||
)
|
||||
out_sglang = fp8_gemm_sglang(
|
||||
x_fp8.clone(), x_scale.clone(), y_fp8.clone(), y_scale.clone(), m, n, k
|
||||
)
|
||||
|
||||
tilelang_func = tl_gemm(m, n, k, "e4m3_float8", "bfloat16", "float32")
|
||||
tilelang_kernel = tilelang.compile(tilelang_func, out_idx=[-1])
|
||||
out_tilelang = tilelang_kernel(
|
||||
x_fp8.clone(), x_scale.clone(), y_fp8.clone(), y_scale.clone()
|
||||
)
|
||||
|
||||
diff_sglang_deepgemm = torch.abs(out_deepgemm - out_sglang).mean().item()
|
||||
diff_tilelang_deepgemm = torch.abs(out_deepgemm - out_tilelang).mean().item()
|
||||
diff_tilelang_sglang = torch.abs(out_tilelang - out_sglang).mean().item()
|
||||
|
||||
print(f"Shape m={m}, n={n}, k={k}:")
|
||||
print(f"DeepGEMM output: {out_deepgemm[0, 0:5]}")
|
||||
print(f"SGLang output: {out_sglang[0, 0:5]}")
|
||||
print(f"TileLang output: {out_tilelang[0, 0:5]}")
|
||||
print(f"Mean absolute difference (SGLang-DeepGEMM): {diff_sglang_deepgemm}")
|
||||
print(f"Mean absolute difference (TileLang-DeepGEMM): {diff_tilelang_deepgemm}")
|
||||
print(f"Mean absolute difference (TileLang-SGLang): {diff_tilelang_sglang}")
|
||||
|
||||
sglang_deepgemm_match = torch.allclose(
|
||||
out_deepgemm, out_sglang, atol=1e-2, rtol=1e-2
|
||||
)
|
||||
tilelang_deepgemm_match = torch.allclose(
|
||||
out_deepgemm, out_tilelang, atol=1e-2, rtol=1e-2
|
||||
)
|
||||
tilelang_sglang_match = torch.allclose(
|
||||
out_tilelang, out_sglang, atol=1e-2, rtol=1e-2
|
||||
)
|
||||
|
||||
if sglang_deepgemm_match and tilelang_deepgemm_match and tilelang_sglang_match:
|
||||
print("✅ All implementations match\n")
|
||||
else:
|
||||
print("❌ Some implementations differ:")
|
||||
print(f" - SGLang vs DeepGEMM: {'✅' if sglang_deepgemm_match else '❌'}")
|
||||
print(f" - TileLang vs DeepGEMM: {'✅' if tilelang_deepgemm_match else '❌'}")
|
||||
print(f" - TileLang vs SGLang: {'✅' if tilelang_sglang_match else '❌'}\n")
|
||||
|
||||
|
||||
def get_weight_shapes(tp_size):
|
||||
# cannot TP
|
||||
total = [
|
||||
(512 + 64, 7168),
|
||||
((128 + 64) * 128, 7168),
|
||||
(128 * (128 + 128), 512),
|
||||
(7168, 16384),
|
||||
(7168, 18432),
|
||||
]
|
||||
# N can TP
|
||||
n_tp = [
|
||||
(18432 * 2, 7168),
|
||||
((128 + 64) * 128, 7168),
|
||||
(128 * (128 + 128), 512),
|
||||
(24576, 1536),
|
||||
(4096, 7168),
|
||||
]
|
||||
# K can TP
|
||||
k_tp = [(7168, 18432), (7168, 16384), (7168, 2048)]
|
||||
|
||||
weight_shapes = []
|
||||
for t in total:
|
||||
weight_shapes.append(t)
|
||||
for n_t in n_tp:
|
||||
new_t = (n_t[0] // tp_size, n_t[1])
|
||||
weight_shapes.append(new_t)
|
||||
for k_t in k_tp:
|
||||
new_t = (k_t[0], k_t[1] // tp_size)
|
||||
weight_shapes.append(new_t)
|
||||
|
||||
return weight_shapes
|
||||
|
||||
|
||||
def create_benchmark_configs(tp_size):
|
||||
configs = []
|
||||
weight_shapes = get_weight_shapes(tp_size)
|
||||
batch_sizes = [8, 16, 32, 64, 128, 256, 1024, 2048, 4096]
|
||||
|
||||
for n, k in weight_shapes:
|
||||
for m in batch_sizes:
|
||||
configs.append((m, n, k, tp_size))
|
||||
|
||||
return configs
|
||||
|
||||
|
||||
def get_benchmark(tp_size):
|
||||
all_configs = create_benchmark_configs(tp_size)
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["m", "n", "k", "tp_size"],
|
||||
x_vals=[list(config) for config in all_configs],
|
||||
line_arg="provider",
|
||||
line_vals=["deepgemm", "sglang", "tilelang"],
|
||||
line_names=["DeepGEMM", "SGLang", "TileLang"],
|
||||
styles=[("blue", "-"), ("red", "-"), ("green", "-")],
|
||||
ylabel="ms",
|
||||
plot_name=f"fp8-gemm-performance-comparison-tp{tp_size}",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(m, n, k, tp_size, provider):
|
||||
print(f"Shape (m={m}, n={n}, k={k}, tp={tp_size}), Provider: {provider}")
|
||||
x = torch.randn((m, k), device="cuda", dtype=torch.bfloat16)
|
||||
y = torch.randn((n, k), device="cuda", dtype=torch.bfloat16)
|
||||
|
||||
# Preprocess data before benchmarking
|
||||
x_fp8, x_scale = per_token_cast_to_fp8(x)
|
||||
y_fp8, y_scale = per_block_cast_to_fp8(y)
|
||||
x_scale_col_major = get_mn_major_tma_aligned_tensor(x_scale.clone())
|
||||
|
||||
quantiles = (0.5, 0.2, 0.8)
|
||||
|
||||
if provider == "deepgemm":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: fp8_gemm_deepgemm(
|
||||
x_fp8.clone(),
|
||||
x_scale_col_major.clone(),
|
||||
y_fp8.clone(),
|
||||
y_scale.clone(),
|
||||
m,
|
||||
n,
|
||||
k,
|
||||
),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
elif provider == "sglang":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: fp8_gemm_sglang(
|
||||
x_fp8.clone(),
|
||||
x_scale.clone(),
|
||||
y_fp8.clone(),
|
||||
y_scale.clone(),
|
||||
m,
|
||||
n,
|
||||
k,
|
||||
),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
else: # tilelang
|
||||
tilelang_func = tl_gemm(m, n, k, "e4m3_float8", "bfloat16", "float32")
|
||||
tilelang_kernel = tilelang.compile(tilelang_func, out_idx=[-1])
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: tilelang_kernel(
|
||||
x_fp8.clone(),
|
||||
x_scale.clone(),
|
||||
y_fp8.clone(),
|
||||
y_scale.clone(),
|
||||
),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
|
||||
# Calculate TFLOPS
|
||||
flops = 2 * m * n * k # multiply-adds
|
||||
tflops = flops / (ms * 1e-3) / 1e12
|
||||
|
||||
# Print shape-specific results with TFLOPS
|
||||
print(f"Time: {ms*1000:.2f} ms, TFLOPS: {tflops:.2f}")
|
||||
return ms * 1000, max_ms * 1000, min_ms * 1000 # convert to ms
|
||||
|
||||
return benchmark
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--save_path",
|
||||
type=str,
|
||||
default="./configs/benchmark_ops/fp8_gemm/",
|
||||
help="Path to save fp8 gemm benchmark results",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--run_correctness",
|
||||
action="store_true",
|
||||
default=True,
|
||||
help="Whether to run correctness test",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tp_size",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Tensor parallelism size to benchmark (default: 1)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
# Set random seed for reproducibility
|
||||
torch.manual_seed(0)
|
||||
torch.cuda.manual_seed(0)
|
||||
|
||||
# Enable TF32, adapted from https://github.com/deepseek-ai/DeepGEMM/blob/main/tests/test_core.py#L148
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
torch.backends.cudnn.allow_tf32 = True
|
||||
|
||||
# Run correctness tests on a few examples
|
||||
if args.run_correctness:
|
||||
print("Running correctness tests...")
|
||||
calculate_diff(64, 512, 7168) # Small test
|
||||
calculate_diff(64, 7168, 16384) # Medium test
|
||||
calculate_diff(64, 18432, 7168) # Large test
|
||||
|
||||
# Get the benchmark function with the specified tp_size
|
||||
benchmark = get_benchmark(args.tp_size)
|
||||
|
||||
print(f"Running performance benchmark for TP size = {args.tp_size}...")
|
||||
benchmark.run(print_data=True, save_path=args.save_path)
|
||||
330
third_party/sglang/benchmark/kernels/deepseek/benchmark_deepgemm_fp8_gemm_blackwell.py
vendored
Normal file
330
third_party/sglang/benchmark/kernels/deepseek/benchmark_deepgemm_fp8_gemm_blackwell.py
vendored
Normal file
@@ -0,0 +1,330 @@
|
||||
import argparse
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import triton
|
||||
from deep_gemm import ceil_div
|
||||
from flashinfer.gemm import gemm_fp8_nt_groupwise
|
||||
|
||||
from sglang.benchmark.bench_utils import run_bench
|
||||
from sglang.srt.layers.quantization.fp8_kernel import (
|
||||
sglang_per_token_group_quant_fp8,
|
||||
w8a8_block_fp8_matmul_deepgemm,
|
||||
)
|
||||
from sglang.srt.layers.quantization.fp8_utils import requant_weight_ue8m0
|
||||
|
||||
BLOCK_SIZE = 128
|
||||
|
||||
|
||||
def per_block_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
assert x.dim() == 2
|
||||
assert BLOCK_SIZE == 128
|
||||
m, n = x.shape
|
||||
x_padded = torch.zeros(
|
||||
(ceil_div(m, 128) * 128, ceil_div(n, 128) * 128), dtype=x.dtype, device=x.device
|
||||
)
|
||||
x_padded[:m, :n] = x
|
||||
x_view = x_padded.view(-1, 128, x_padded.size(1) // 128, 128)
|
||||
x_amax = x_view.abs().float().amax(dim=(1, 3), keepdim=True).clamp(1e-4)
|
||||
x_scaled = (x_view * (448.0 / x_amax)).to(torch.float8_e4m3fn)
|
||||
return x_scaled.view_as(x_padded)[:m, :n].contiguous(), (x_amax / 448.0).view(
|
||||
x_view.size(0), x_view.size(2)
|
||||
)
|
||||
|
||||
|
||||
def get_weight_shapes(tp_size):
|
||||
# cannot TP
|
||||
total = [
|
||||
(512 + 64, 7168),
|
||||
((128 + 64) * 128, 7168),
|
||||
(128 * (128 + 128), 512),
|
||||
(7168, 16384),
|
||||
(7168, 18432),
|
||||
]
|
||||
# N can TP
|
||||
n_tp = [
|
||||
(18432 * 2, 7168),
|
||||
((128 + 64) * 128, 7168),
|
||||
(128 * (128 + 128), 512),
|
||||
(24576, 1536),
|
||||
(4096, 7168),
|
||||
]
|
||||
# K can TP
|
||||
k_tp = [(7168, 18432), (7168, 16384), (7168, 2048)]
|
||||
|
||||
weight_shapes = []
|
||||
for t in total:
|
||||
weight_shapes.append(t)
|
||||
for n_t in n_tp:
|
||||
new_t = (n_t[0] // tp_size, n_t[1])
|
||||
weight_shapes.append(new_t)
|
||||
for k_t in k_tp:
|
||||
new_t = (k_t[0], k_t[1] // tp_size)
|
||||
weight_shapes.append(new_t)
|
||||
|
||||
return weight_shapes
|
||||
|
||||
|
||||
def create_benchmark_configs(tp_size):
|
||||
configs = []
|
||||
weight_shapes = get_weight_shapes(tp_size)
|
||||
batch_sizes = [8, 16, 32, 64, 128, 256, 1024, 2048, 4096]
|
||||
|
||||
for n, k in weight_shapes:
|
||||
for m in batch_sizes:
|
||||
configs.append((m, n, k, tp_size))
|
||||
|
||||
return configs
|
||||
|
||||
|
||||
def fp8_gemm_flashinfer(
|
||||
x_fp8: torch.Tensor,
|
||||
x_scale: torch.Tensor,
|
||||
y_fp8: torch.Tensor,
|
||||
y_scale: torch.Tensor,
|
||||
):
|
||||
"""Flashinfer implementation of FP8 GEMM"""
|
||||
output = gemm_fp8_nt_groupwise(
|
||||
x_fp8,
|
||||
y_fp8,
|
||||
x_scale,
|
||||
y_scale,
|
||||
out_dtype=torch.bfloat16,
|
||||
backend="trtllm",
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
def fp8_gemm_deepgemm_blackwell(
|
||||
x_fp8: torch.Tensor,
|
||||
x_scale: torch.Tensor,
|
||||
y_fp8: torch.Tensor,
|
||||
y_scale: torch.Tensor,
|
||||
):
|
||||
"""DeepGEMM implementation of FP8 GEMM"""
|
||||
block_size = [BLOCK_SIZE, BLOCK_SIZE]
|
||||
output = w8a8_block_fp8_matmul_deepgemm(
|
||||
x_fp8, y_fp8, x_scale, y_scale, block_size, output_dtype=torch.bfloat16
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
def check_accuracy(a, b, atol, rtol, percent):
|
||||
"""Unified accuracy checking function with detailed error reporting."""
|
||||
if not torch.isfinite(a).all():
|
||||
print("Non-finite values in reference output")
|
||||
return False
|
||||
if not torch.isfinite(b).all():
|
||||
print("Non-finite values in actual output")
|
||||
return False
|
||||
assert a.shape == b.shape, f"Shape mismatch: {a.shape} vs {b.shape}"
|
||||
|
||||
close = torch.isclose(a, b, atol=atol, rtol=rtol)
|
||||
match_ratio = close.float().mean()
|
||||
if match_ratio >= percent:
|
||||
return True
|
||||
|
||||
mismatch_percent = 1.0 - match_ratio.item()
|
||||
if mismatch_percent > 1 - percent:
|
||||
print(
|
||||
f"Mismatch percentage is {mismatch_percent:.4f} for rtol {rtol} "
|
||||
f"(threshold: {1 - percent:.4f})"
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def calculate_diff(m: int, n: int, k: int):
|
||||
x = torch.randn((m, k), device="cuda", dtype=torch.bfloat16)
|
||||
y = torch.randn((n, k), device="cuda", dtype=torch.bfloat16)
|
||||
|
||||
y_fp8, y_scale = per_block_cast_to_fp8(y)
|
||||
x_fp8, x_scale = sglang_per_token_group_quant_fp8(
|
||||
x, BLOCK_SIZE, column_major_scales=True
|
||||
)
|
||||
out_flashinfer = fp8_gemm_flashinfer(
|
||||
x_fp8,
|
||||
x_scale,
|
||||
y_fp8,
|
||||
y_scale,
|
||||
)
|
||||
|
||||
dg_x_fp8, dg_x_scale = sglang_per_token_group_quant_fp8(
|
||||
x,
|
||||
BLOCK_SIZE,
|
||||
column_major_scales=True,
|
||||
scale_tma_aligned=True,
|
||||
scale_ue8m0=True,
|
||||
)
|
||||
# We can directly quantize y here, but to mimic the behavior of the actual
|
||||
# implementations, we requant it here.
|
||||
dg_y_fp8, dg_y_scale = requant_weight_ue8m0(
|
||||
y_fp8, y_scale, [BLOCK_SIZE, BLOCK_SIZE]
|
||||
)
|
||||
out_deepgemm = fp8_gemm_deepgemm_blackwell(
|
||||
dg_x_fp8, dg_x_scale, dg_y_fp8, dg_y_scale
|
||||
)
|
||||
|
||||
print(f"Shape m={m}, n={n}, k={k}:")
|
||||
print(f"Flashinfer output: {out_flashinfer[0, 0:5]}")
|
||||
print(f"DeepGEMM output: {out_deepgemm[0, 0:5]}")
|
||||
|
||||
flashinfer_deepgemm_match = check_accuracy(
|
||||
out_flashinfer, out_deepgemm, 0.1, 0.6, 0.95
|
||||
)
|
||||
print("Correctness check:")
|
||||
print(f" - Flashinfer vs DeepGEMM: {'✅' if flashinfer_deepgemm_match else '❌'}")
|
||||
|
||||
|
||||
def _benchmark(m, n, k, tp_size, provider):
|
||||
print(f"Shape (m={m}, n={n}, k={k}, tp={tp_size}), Provider: {provider}")
|
||||
x = torch.randn((m, k), device="cuda", dtype=torch.bfloat16)
|
||||
y = torch.randn((n, k), device="cuda", dtype=torch.bfloat16)
|
||||
|
||||
# Preprocess data before benchmarking
|
||||
y_fp8, y_scale = per_block_cast_to_fp8(y)
|
||||
x_fp8, x_scale = sglang_per_token_group_quant_fp8(
|
||||
x, BLOCK_SIZE, column_major_scales=True
|
||||
)
|
||||
dg_x_fp8, dg_x_scale = sglang_per_token_group_quant_fp8(
|
||||
x,
|
||||
BLOCK_SIZE,
|
||||
column_major_scales=True,
|
||||
scale_tma_aligned=True,
|
||||
scale_ue8m0=True,
|
||||
)
|
||||
dg_y_fp8, dg_y_scale = requant_weight_ue8m0(
|
||||
y_fp8, y_scale, [BLOCK_SIZE, BLOCK_SIZE]
|
||||
)
|
||||
|
||||
quantiles = (0.5, 0.2, 0.8)
|
||||
|
||||
if provider == "deepgemm":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: fp8_gemm_deepgemm_blackwell(
|
||||
dg_x_fp8,
|
||||
dg_x_scale,
|
||||
dg_y_fp8,
|
||||
dg_y_scale,
|
||||
),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
elif provider == "flashinfer":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: fp8_gemm_flashinfer(
|
||||
x_fp8,
|
||||
x_scale,
|
||||
y_fp8,
|
||||
y_scale,
|
||||
),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
|
||||
# Calculate TFLOPS
|
||||
flops = 2 * m * n * k # multiply-adds
|
||||
tflops = flops / (ms * 1e-3) / 1e12
|
||||
|
||||
# Print shape-specific results with TFLOPS
|
||||
print(f"Time: {ms*1000:.2f} us, TFLOPS: {tflops:.2f}")
|
||||
return ms, max_ms, min_ms
|
||||
|
||||
|
||||
def get_benchmark_plot_friendly(tp_size):
|
||||
all_configs = create_benchmark_configs(tp_size)
|
||||
x_vals = list(range(len(all_configs)))
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["cfg_id"],
|
||||
x_vals=x_vals,
|
||||
line_arg="provider",
|
||||
line_vals=["deepgemm", "flashinfer"],
|
||||
line_names=["DeepGEMM", "Flashinfer"],
|
||||
styles=[("blue", "-"), ("red", "-")],
|
||||
ylabel="us",
|
||||
plot_name=f"fp8-gemm-performance-comparison-tp{tp_size}",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(cfg_id, provider):
|
||||
m, n, k, tp_size = all_configs[cfg_id]
|
||||
ms, min_ms, max_ms = _benchmark(m, n, k, tp_size, provider)
|
||||
return ms * 1000, max_ms * 1000, min_ms * 1000 # convert to ms
|
||||
|
||||
return benchmark
|
||||
|
||||
|
||||
def get_benchmark(tp_size):
|
||||
all_configs = create_benchmark_configs(tp_size)
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["m", "n", "k", "tp_size"],
|
||||
x_vals=[list(config) for config in all_configs],
|
||||
line_arg="provider",
|
||||
line_vals=["deepgemm", "flashinfer"],
|
||||
line_names=["DeepGEMM", "Flashinfer"],
|
||||
styles=[("blue", "-"), ("red", "-")],
|
||||
ylabel="us",
|
||||
plot_name=f"fp8-gemm-performance-comparison-tp{tp_size}",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(m, n, k, tp_size, provider):
|
||||
ms, min_ms, max_ms = _benchmark(m, n, k, tp_size, provider)
|
||||
return ms * 1000, max_ms * 1000, min_ms * 1000 # convert to ms
|
||||
|
||||
return benchmark
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] != 10:
|
||||
print("Skipping benchmark because the device is not supported")
|
||||
exit(0)
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--save-path",
|
||||
type=str,
|
||||
default="./configs/benchmark_ops/fp8_gemm/",
|
||||
help="Path to save fp8 gemm benchmark results",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--run-correctness",
|
||||
action="store_true",
|
||||
default=True,
|
||||
help="Whether to run correctness test",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tp-size",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Tensor parallelism size to benchmark (default: 1)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--plot-friendly",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Plot x axis as the config index instead of the m",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
# Set random seed for reproducibility
|
||||
torch.manual_seed(0)
|
||||
torch.cuda.manual_seed(0)
|
||||
|
||||
# Run correctness tests on a few examples
|
||||
if args.run_correctness:
|
||||
print("Running correctness tests...")
|
||||
calculate_diff(64, 512, 7168) # Small test
|
||||
calculate_diff(64, 7168, 16384) # Medium test
|
||||
calculate_diff(64, 18432, 7168) # Large test
|
||||
|
||||
# Get the benchmark function with the specified tp_size
|
||||
benchmark = (
|
||||
get_benchmark_plot_friendly(args.tp_size)
|
||||
if args.plot_friendly
|
||||
else get_benchmark(args.tp_size)
|
||||
)
|
||||
|
||||
print(f"Running performance benchmark for TP size = {args.tp_size}...")
|
||||
benchmark.run(print_data=True, save_path=args.save_path)
|
||||
488
third_party/sglang/benchmark/kernels/deepseek/benchmark_deepgemm_fp8_group_gemm.py
vendored
Normal file
488
third_party/sglang/benchmark/kernels/deepseek/benchmark_deepgemm_fp8_group_gemm.py
vendored
Normal file
@@ -0,0 +1,488 @@
|
||||
from typing import Tuple
|
||||
|
||||
import deep_gemm
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
from deep_gemm import calc_diff
|
||||
from deep_gemm.utils.layout import get_mn_major_tma_aligned_tensor
|
||||
|
||||
# Import shared functionality from the regular GEMM benchmark
|
||||
from sglang.benchmark.bench_utils import run_bench
|
||||
from sglang.benchmark.kernels.deepseek.benchmark_deepgemm_fp8_gemm import (
|
||||
per_block_cast_to_fp8,
|
||||
per_token_cast_to_fp8,
|
||||
)
|
||||
|
||||
|
||||
def construct_grouped_and_flat_fp8(
|
||||
x: torch.Tensor, y: torch.Tensor, num_groups: int, is_masked: bool
|
||||
) -> Tuple[
|
||||
Tuple[torch.Tensor, torch.Tensor], # grouped x_fp8
|
||||
Tuple[torch.Tensor, torch.Tensor], # grouped y_fp8
|
||||
Tuple[torch.Tensor, torch.Tensor], # flat x_fp8
|
||||
Tuple[torch.Tensor, torch.Tensor], # flat y_fp8
|
||||
torch.Tensor, # output
|
||||
torch.Tensor, # reference output
|
||||
]:
|
||||
# Verify input shapes
|
||||
m, k = x.shape
|
||||
n, k_y = y.shape
|
||||
assert k == k_y, f"Incompatible shapes: x({m}, {k}), y({n}, {k_y})"
|
||||
assert m % num_groups == 0, f"m({m}) must be divisible by num_groups({num_groups})"
|
||||
assert m % 4 == 0, f"TMA alignment error: {m}"
|
||||
|
||||
# Reshape inputs for grouped processing
|
||||
m_per_group = m // num_groups
|
||||
x_grouped = x.view(num_groups, m_per_group, k)
|
||||
y_grouped = y.unsqueeze(0).expand(num_groups, n, k)
|
||||
|
||||
# Initialize output tensors
|
||||
out = torch.empty((num_groups, m_per_group, n), device="cuda", dtype=torch.bfloat16)
|
||||
ref_out = torch.einsum("gmk,gnk->gmn", x_grouped, y_grouped)
|
||||
|
||||
# Quantize grouped tensors
|
||||
x_fp8_grouped = (
|
||||
torch.empty_like(x_grouped, dtype=torch.float8_e4m3fn),
|
||||
torch.empty(
|
||||
(num_groups, m_per_group, k // 128), device="cuda", dtype=torch.float
|
||||
),
|
||||
)
|
||||
y_fp8_grouped = (
|
||||
torch.empty_like(y_grouped, dtype=torch.float8_e4m3fn),
|
||||
torch.empty(
|
||||
(num_groups, (n + 127) // 128, k // 128), device="cuda", dtype=torch.float
|
||||
),
|
||||
)
|
||||
for i in range(num_groups):
|
||||
x_fp8_grouped[0][i], x_fp8_grouped[1][i] = per_token_cast_to_fp8(x_grouped[i])
|
||||
y_fp8_grouped[0][i], y_fp8_grouped[1][i] = per_block_cast_to_fp8(y_grouped[i])
|
||||
|
||||
# Quantize flat tensors
|
||||
x_fp8_flat = per_token_cast_to_fp8(x)
|
||||
y_fp8_flat = per_block_cast_to_fp8(y)
|
||||
|
||||
# For non-masked input, merge the group and M dims in output
|
||||
if not is_masked:
|
||||
x_fp8_grouped = (
|
||||
x_fp8_grouped[0].view(-1, k),
|
||||
per_token_cast_to_fp8(x_grouped.view(-1, k))[1],
|
||||
)
|
||||
out, ref_out = out.view(-1, n), ref_out.view(-1, n)
|
||||
|
||||
# Transpose earlier for testing
|
||||
x_fp8_grouped = (
|
||||
x_fp8_grouped[0],
|
||||
get_mn_major_tma_aligned_tensor(x_fp8_grouped[1]),
|
||||
)
|
||||
x_fp8_flat = (x_fp8_flat[0], get_mn_major_tma_aligned_tensor(x_fp8_flat[1]))
|
||||
|
||||
return x_fp8_grouped, y_fp8_grouped, x_fp8_flat, y_fp8_flat, out, ref_out
|
||||
|
||||
|
||||
# Since we don't have a group gemm kernel in SGLang/vLLM, we implemented a
|
||||
# custom kernel based on the Triton tutorial.
|
||||
# https://triton-lang.org/main/getting-started/tutorials/03-matrix-multiplication.html
|
||||
@triton.jit
|
||||
def fp8_gemm_group_triton_kernel(
|
||||
# Pointers to matrices
|
||||
a_ptr,
|
||||
b_ptr,
|
||||
c_ptr,
|
||||
# Pointers to scaling factors
|
||||
a_scale_ptr,
|
||||
b_scale_ptr,
|
||||
# Matrix dimensions
|
||||
M,
|
||||
N,
|
||||
K,
|
||||
# The stride variables represent how much to increase the ptr by when moving by 1
|
||||
# element in a particular dimension.
|
||||
stride_am,
|
||||
stride_ak,
|
||||
stride_bk,
|
||||
stride_bn,
|
||||
stride_cm,
|
||||
stride_cn,
|
||||
# Strides for scaling factors
|
||||
stride_a_scale_m,
|
||||
stride_a_scale_k,
|
||||
stride_b_scale_n,
|
||||
stride_b_scale_k,
|
||||
# Meta-parameters
|
||||
BLOCK_SIZE_M: tl.constexpr,
|
||||
BLOCK_SIZE_N: tl.constexpr,
|
||||
BLOCK_SIZE_K: tl.constexpr,
|
||||
GROUP_SIZE_M: tl.constexpr,
|
||||
):
|
||||
"""Kernel for computing the matmul C = A x B with FP8 inputs and scaling factors.
|
||||
A has shape (M, K), B has shape (K, N) and C has shape (M, N)
|
||||
|
||||
Note: Block sizes must be multiples of 32 for optimal TMA performance.
|
||||
"""
|
||||
# Map program ids to the block of C it should compute
|
||||
pid_group = tl.program_id(axis=0) # Group ID
|
||||
pid_n = tl.program_id(axis=1) # N dimension ID
|
||||
|
||||
# Compute the M block ID within this group
|
||||
group_size_m = min(M - pid_group * GROUP_SIZE_M, GROUP_SIZE_M)
|
||||
pid_m_within_group = tl.program_id(axis=2) % group_size_m
|
||||
pid_m = pid_group * GROUP_SIZE_M + pid_m_within_group
|
||||
|
||||
# Create pointers for the first blocks of A and B
|
||||
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
|
||||
offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
|
||||
offs_k = tl.arange(0, BLOCK_SIZE_K)
|
||||
a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak)
|
||||
b_ptrs = b_ptr + (offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn)
|
||||
|
||||
# Initialize accumulator
|
||||
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
|
||||
|
||||
# Main loop
|
||||
for k_block in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
|
||||
k_offset = k_block * BLOCK_SIZE_K
|
||||
|
||||
# Load the next block of A and B, with masks
|
||||
a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k_offset, other=0.0)
|
||||
b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k_offset, other=0.0)
|
||||
|
||||
# Calculate indices for scaling factors for this K block
|
||||
a_scale_ptrs = a_scale_ptr + (
|
||||
offs_am * stride_a_scale_m + k_block * stride_a_scale_k
|
||||
)
|
||||
b_scale_ptrs = b_scale_ptr + (
|
||||
pid_n * stride_b_scale_n + k_block * stride_b_scale_k
|
||||
)
|
||||
|
||||
# Perform matrix multiplication in FP8
|
||||
res = tl.dot(a, b)
|
||||
|
||||
# Load scaling factors for the current block
|
||||
a_scale = tl.load(a_scale_ptrs)[:, None] # [BLOCK_SIZE_M, 1]
|
||||
b_scale = tl.load(b_scale_ptrs)
|
||||
|
||||
# Apply scaling factors to the accumulated result
|
||||
accumulator += res * a_scale * b_scale
|
||||
|
||||
# Advance pointers
|
||||
a_ptrs += BLOCK_SIZE_K * stride_ak
|
||||
b_ptrs += BLOCK_SIZE_K * stride_bk
|
||||
|
||||
# Convert to bfloat16 for output
|
||||
c = accumulator.to(tl.bfloat16)
|
||||
|
||||
# Write back the result
|
||||
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
|
||||
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
|
||||
c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
|
||||
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
|
||||
tl.store(c_ptrs, c, mask=c_mask)
|
||||
|
||||
|
||||
def fp8_gemm_group_triton(a_tuple, b_tuple, c, num_groups):
|
||||
"""
|
||||
Perform matrix multiplication with FP8 inputs and proper scaling.
|
||||
|
||||
Args:
|
||||
a_tuple: Tuple of (quantized_tensor, scale_factors) for input A
|
||||
b_tuple: Tuple of (quantized_tensor, scale_factors) for input B
|
||||
c: Output tensor in BF16 format
|
||||
num_groups: Number of groups for grouped GEMM
|
||||
|
||||
Returns:
|
||||
Result tensor in BF16 format
|
||||
"""
|
||||
# Unpack the tuples
|
||||
a, a_scale = a_tuple
|
||||
b, b_scale = b_tuple
|
||||
|
||||
M, K = a.shape
|
||||
_, N = b.shape
|
||||
|
||||
# Configure block sizes - must be multiples of 32 for TMA alignment
|
||||
BLOCK_SIZE_M = 128
|
||||
BLOCK_SIZE_N = 128
|
||||
BLOCK_SIZE_K = 128
|
||||
|
||||
# Calculate grid dimensions
|
||||
num_pid_m = triton.cdiv(M, BLOCK_SIZE_M)
|
||||
num_pid_n = triton.cdiv(N, BLOCK_SIZE_N)
|
||||
num_groups_grid = triton.cdiv(num_pid_m, num_groups)
|
||||
|
||||
# 3D grid launch - (group, n_blocks, m_blocks_per_group)
|
||||
grid = (num_groups_grid, num_pid_n, min(num_groups, num_pid_m))
|
||||
|
||||
fp8_gemm_group_triton_kernel[grid](
|
||||
a,
|
||||
b,
|
||||
c,
|
||||
a_scale,
|
||||
b_scale,
|
||||
M,
|
||||
N,
|
||||
K,
|
||||
a.stride(0),
|
||||
a.stride(1),
|
||||
b.stride(0),
|
||||
b.stride(1),
|
||||
c.stride(0),
|
||||
c.stride(1),
|
||||
a_scale.stride(0),
|
||||
1, # Stride in the K dimension may be 1
|
||||
b_scale.stride(0),
|
||||
1 if b_scale.dim() > 1 else 0,
|
||||
BLOCK_SIZE_M=BLOCK_SIZE_M,
|
||||
BLOCK_SIZE_N=BLOCK_SIZE_N,
|
||||
BLOCK_SIZE_K=BLOCK_SIZE_K,
|
||||
GROUP_SIZE_M=num_groups,
|
||||
)
|
||||
|
||||
return c
|
||||
|
||||
|
||||
def fp8_gemm_group_deepgemm(x_fp8_grouped, y_fp8_grouped, out, m_indices):
|
||||
deep_gemm.m_grouped_fp8_gemm_nt_contiguous(
|
||||
x_fp8_grouped,
|
||||
y_fp8_grouped,
|
||||
out,
|
||||
m_indices,
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def calculate_diff(m: int, n: int, k: int, num_groups: int):
|
||||
print(f"Shape (m={m}, n={n}, k={k}")
|
||||
x = torch.randn((m, k), device="cuda", dtype=torch.bfloat16)
|
||||
y = torch.randn((n, k), device="cuda", dtype=torch.bfloat16)
|
||||
x_fp8_grouped, y_fp8_grouped, x_fp8_flat, y_fp8_flat, out, out_torch = (
|
||||
construct_grouped_and_flat_fp8(x, y, num_groups, is_masked=False)
|
||||
)
|
||||
m_per_group = m // num_groups
|
||||
out_deepgemm = out.clone()
|
||||
m_indices = torch.arange(0, num_groups, device="cuda", dtype=torch.int)
|
||||
m_indices = (
|
||||
m_indices.unsqueeze(-1).expand(num_groups, m_per_group).contiguous().view(-1)
|
||||
)
|
||||
|
||||
fp8_gemm_group_deepgemm(
|
||||
x_fp8_grouped,
|
||||
y_fp8_grouped,
|
||||
out_deepgemm,
|
||||
m_indices,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Prepare inputs for Triton
|
||||
a, a_scale = x_fp8_flat
|
||||
b, b_scale = y_fp8_flat
|
||||
b = b.T.contiguous()
|
||||
# Ensure scales are in the right format and contiguous
|
||||
a_scale, b_scale = a_scale.contiguous(), b_scale.contiguous()
|
||||
M, _ = a.shape
|
||||
_, N = b.shape
|
||||
c = torch.empty((M, N), device=a.device, dtype=torch.bfloat16)
|
||||
out_triton = fp8_gemm_group_triton((a, a_scale), (b, b_scale), c, num_groups)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
diff_torch_deepgemm = torch.abs(out_torch - out_deepgemm).mean().item()
|
||||
diff_torch_triton = torch.abs(out_torch - out_triton).mean().item()
|
||||
diff_deepgemm_triton = torch.abs(out_deepgemm - out_triton).mean().item()
|
||||
|
||||
print(f"Shape m={m}, n={n}, k={k}:")
|
||||
print(f"Torch output: {out_torch[0, 0:5]}")
|
||||
print(f"DeepGEMM output: {out_deepgemm[0, 0:5]}")
|
||||
print(f"Triton output: {out_triton[0, 0:5]}")
|
||||
print(f"Mean absolute difference (Torch-DeepGEMM): {diff_torch_deepgemm}")
|
||||
print(f"Mean absolute difference (Torch-Triton): {diff_torch_triton}")
|
||||
print(f"Mean absolute difference (DeepGEMM-Triton): {diff_deepgemm_triton}")
|
||||
|
||||
deepgemm_torch_diff = calc_diff(out_deepgemm, out_torch)
|
||||
triton_torch_diff = calc_diff(out_triton, out_torch)
|
||||
deepgemm_triton_diff = calc_diff(out_deepgemm, out_triton)
|
||||
|
||||
DIFF_THRESHOLD = 0.001
|
||||
all_match = (
|
||||
deepgemm_torch_diff < DIFF_THRESHOLD
|
||||
and triton_torch_diff < DIFF_THRESHOLD
|
||||
and deepgemm_triton_diff < DIFF_THRESHOLD
|
||||
)
|
||||
if all_match:
|
||||
print("✅ All implementations match\n")
|
||||
else:
|
||||
print("❌ Some implementations differ:")
|
||||
print(
|
||||
f" - Torch vs DeepGEMM: {'✅' if deepgemm_torch_diff < DIFF_THRESHOLD else '❌'}"
|
||||
f" - Torch vs Triton: {'✅' if triton_torch_diff < DIFF_THRESHOLD else '❌'}"
|
||||
f" - DeepGEMM vs Triton: {'✅' if deepgemm_triton_diff < DIFF_THRESHOLD else '❌'}"
|
||||
)
|
||||
|
||||
|
||||
def get_weight_shapes(tp_size):
|
||||
# cannot TP
|
||||
total = [
|
||||
(512 + 64, 7168),
|
||||
((128 + 64) * 128, 7168),
|
||||
(128 * (128 + 128), 512),
|
||||
(7168, 16384),
|
||||
(7168, 18432),
|
||||
]
|
||||
# N can TP
|
||||
n_tp = [
|
||||
(18432 * 2, 7168),
|
||||
((128 + 64) * 128, 7168),
|
||||
(128 * (128 + 128), 512),
|
||||
(24576, 1536),
|
||||
(4096, 7168),
|
||||
]
|
||||
# K can TP
|
||||
k_tp = [(7168, 18432), (7168, 16384), (7168, 2048)]
|
||||
|
||||
weight_shapes = []
|
||||
for t in total:
|
||||
weight_shapes.append(t)
|
||||
for n_t in n_tp:
|
||||
new_t = (n_t[0] // tp_size, n_t[1])
|
||||
weight_shapes.append(new_t)
|
||||
for k_t in k_tp:
|
||||
new_t = (k_t[0], k_t[1] // tp_size)
|
||||
weight_shapes.append(new_t)
|
||||
|
||||
return weight_shapes
|
||||
|
||||
|
||||
def create_benchmark_configs(tp_size):
|
||||
configs = []
|
||||
weight_shapes = get_weight_shapes(tp_size)
|
||||
batch_sizes = [2048, 4096]
|
||||
group_sizes = [4, 8]
|
||||
for n, k in weight_shapes:
|
||||
for m in batch_sizes:
|
||||
for num_groups in group_sizes:
|
||||
configs.append((m, n, k, num_groups, tp_size))
|
||||
|
||||
return configs
|
||||
|
||||
|
||||
def get_benchmark(tp_size):
|
||||
all_configs = create_benchmark_configs(tp_size)
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["m", "n", "k", "num_groups", "tp_size"],
|
||||
x_vals=[config for config in all_configs],
|
||||
line_arg="provider",
|
||||
line_vals=["deepgemm", "triton"],
|
||||
line_names=["DeepGEMM", "Triton"],
|
||||
styles=[("blue", "-"), ("red", "-")],
|
||||
ylabel="ms",
|
||||
plot_name=f"fp8-group-gemm-performance-comparison-tp{tp_size}",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(m, n, k, num_groups, tp_size, provider):
|
||||
print(
|
||||
f"Shape (m={m}, n={n}, k={k}, tp={tp_size}, num_groups={num_groups}, Provider: {provider}"
|
||||
)
|
||||
x = torch.randn((m, k), device="cuda", dtype=torch.bfloat16)
|
||||
y = torch.randn((n, k), device="cuda", dtype=torch.bfloat16)
|
||||
x_fp8_grouped, y_fp8_grouped, x_fp8_flat, y_fp8_flat, out, out_torch = (
|
||||
construct_grouped_and_flat_fp8(x, y, num_groups, is_masked=False)
|
||||
)
|
||||
m_per_group = m // num_groups
|
||||
m_indices = torch.arange(0, num_groups, device="cuda", dtype=torch.int)
|
||||
m_indices = (
|
||||
m_indices.unsqueeze(-1)
|
||||
.expand(num_groups, m_per_group)
|
||||
.contiguous()
|
||||
.view(-1)
|
||||
)
|
||||
|
||||
quantiles = (0.5, 0.2, 0.8)
|
||||
|
||||
if provider == "deepgemm":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: fp8_gemm_group_deepgemm(
|
||||
x_fp8_grouped,
|
||||
y_fp8_grouped,
|
||||
out,
|
||||
m_indices,
|
||||
),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
elif provider == "triton":
|
||||
# Prepare inputs for Triton
|
||||
# We did it outside of the lambda function to make it fair comparison like deepgemm
|
||||
a, a_scale = x_fp8_flat
|
||||
b, b_scale = y_fp8_flat
|
||||
b = b.T.contiguous()
|
||||
# Ensure scales are in the right format and contiguous
|
||||
a_scale, b_scale = a_scale.contiguous(), b_scale.contiguous()
|
||||
M, _ = a.shape
|
||||
_, N = b.shape
|
||||
c = torch.empty((M, N), device=a.device, dtype=torch.bfloat16)
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: fp8_gemm_group_triton(
|
||||
(a, a_scale),
|
||||
(b, b_scale),
|
||||
c,
|
||||
num_groups,
|
||||
),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
|
||||
# Calculate TFLOPS
|
||||
flops = 2 * m * n * k # multiply-adds
|
||||
tflops = flops / (ms * 1e-3) / 1e12
|
||||
|
||||
print(f"Time: {ms*1000:.2f} ms, TFLOPS: {tflops:.2f}")
|
||||
return ms * 1000, max_ms * 1000, min_ms * 1000 # convert to ms
|
||||
|
||||
return benchmark
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--save_path",
|
||||
type=str,
|
||||
default="./configs/benchmark_ops/fp8_group_gemm/",
|
||||
help="Path to save deepgemm fp8 group gemm benchmark results",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--run_correctness",
|
||||
action="store_true",
|
||||
help="Whether to run correctness test",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tp_size",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Tensor parallelism size to benchmark (default: 1)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
# Set random seed for reproducibility
|
||||
torch.manual_seed(0)
|
||||
torch.cuda.manual_seed(0)
|
||||
|
||||
# Enable TF32, adapted from https://github.com/deepseek-ai/DeepGEMM/blob/main/tests/test_core.py#L148
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
torch.backends.cudnn.allow_tf32 = True
|
||||
|
||||
# Run correctness tests on a few examples
|
||||
if args.run_correctness:
|
||||
print("Running correctness tests...")
|
||||
calculate_diff(8192, 7168, 4096, 4)
|
||||
calculate_diff(8192, 2048, 7168, 4)
|
||||
calculate_diff(4096, 7168, 4096, 8)
|
||||
calculate_diff(4096, 2048, 7168, 8)
|
||||
calculate_diff(4096, 576, 7168, 8)
|
||||
|
||||
# Get the benchmark function with the specified tp_size
|
||||
benchmark = get_benchmark(args.tp_size)
|
||||
|
||||
print(f"Running performance benchmark for TP size = {args.tp_size}...")
|
||||
benchmark.run(print_data=True, save_path=args.save_path)
|
||||
198
third_party/sglang/benchmark/kernels/elementwise/benchmark_concat_mla.py
vendored
Normal file
198
third_party/sglang/benchmark/kernels/elementwise/benchmark_concat_mla.py
vendored
Normal file
@@ -0,0 +1,198 @@
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
from sgl_kernel import concat_mla_k as concat_mla_k_cuda
|
||||
|
||||
from sglang.benchmark.bench_utils import run_bench
|
||||
|
||||
DEVICE = triton.runtime.driver.active.get_active_torch_device()
|
||||
|
||||
num_local_heads = 128
|
||||
qk_nope_head_dim = 128
|
||||
qk_rope_head_dim = 64
|
||||
|
||||
|
||||
def create_data(num_tokens):
|
||||
k_nope_container = torch.randn(
|
||||
(num_tokens, num_local_heads, qk_nope_head_dim + 128),
|
||||
dtype=torch.bfloat16,
|
||||
device="cuda",
|
||||
)
|
||||
k_nope = k_nope_container[:, :, :qk_nope_head_dim]
|
||||
|
||||
k_rope_container = torch.randn(
|
||||
(num_tokens, 1, 128 + qk_rope_head_dim), dtype=torch.bfloat16, device="cuda"
|
||||
)
|
||||
k_rope = k_rope_container[:, :, -qk_rope_head_dim:]
|
||||
|
||||
k = torch.empty(
|
||||
(num_tokens, num_local_heads, qk_nope_head_dim + qk_rope_head_dim),
|
||||
dtype=torch.bfloat16,
|
||||
device="cuda",
|
||||
)
|
||||
return dict(k=k, k_nope=k_nope, k_rope=k_rope)
|
||||
|
||||
|
||||
def fn_torch(k, k_nope, k_rope):
|
||||
k[..., :qk_nope_head_dim] = k_nope
|
||||
k[..., qk_nope_head_dim:] = k_rope
|
||||
|
||||
|
||||
def fn_hack_non_strided(k, k_nope, k_rope):
|
||||
k_flatten_view = k.flatten()
|
||||
k_flatten_view[: k_nope.numel()] = k_nope.flatten()
|
||||
|
||||
k2 = k_flatten_view[k_nope.numel() :].view(k_rope.numel(), -1)
|
||||
k2 = k_rope.flatten()[:, None]
|
||||
|
||||
|
||||
@torch.compile(dynamic=True)
|
||||
def fn_torch_compiled(k, k_nope, k_rope):
|
||||
return fn_torch(k, k_nope, k_rope)
|
||||
|
||||
|
||||
def fn_cuda(k, k_nope, k_rope):
|
||||
concat_mla_k_cuda(k, k_nope, k_rope)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def fn_triton_kernel(
|
||||
k_ptr,
|
||||
k_nope_ptr,
|
||||
k_rope_ptr,
|
||||
num_tokens,
|
||||
QK_NOPE_HEAD_DIM: tl.constexpr,
|
||||
QK_ROPE_HEAD_DIM: tl.constexpr,
|
||||
NUM_LOCAL_HEADS: tl.constexpr,
|
||||
K_NOPE_STRIDE_0: tl.constexpr,
|
||||
K_NOPE_STRIDE_1: tl.constexpr,
|
||||
K_STRIDE_0: tl.constexpr,
|
||||
K_STRIDE_1: tl.constexpr,
|
||||
K_ROPE_STRIDE_0: tl.constexpr,
|
||||
BLOCK_ROWS: tl.constexpr,
|
||||
):
|
||||
pid = tl.program_id(axis=0)
|
||||
|
||||
token_id = pid * BLOCK_ROWS + tl.arange(0, BLOCK_ROWS)
|
||||
token_mask = token_id < num_tokens
|
||||
|
||||
head_id = tl.arange(0, NUM_LOCAL_HEADS)
|
||||
|
||||
# nope
|
||||
nope_sub_id = tl.arange(0, QK_NOPE_HEAD_DIM)
|
||||
offs_nope = (
|
||||
token_id[:, None, None] * K_NOPE_STRIDE_0
|
||||
+ head_id[None, :, None] * K_NOPE_STRIDE_1
|
||||
+ nope_sub_id[None, None, :]
|
||||
)
|
||||
offs_k = (
|
||||
token_id[:, None, None] * K_STRIDE_0
|
||||
+ head_id[None, :, None] * K_STRIDE_1
|
||||
+ nope_sub_id[None, None, :]
|
||||
)
|
||||
vals_nope = tl.load(k_nope_ptr + offs_nope, mask=token_mask[:, None, None])
|
||||
tl.store(k_ptr + offs_k, vals_nope, mask=token_mask[:, None, None])
|
||||
|
||||
# rope
|
||||
rope_sub_id = tl.arange(0, QK_ROPE_HEAD_DIM)
|
||||
offs_rope = token_id[:, None, None] * K_ROPE_STRIDE_0 + rope_sub_id[None, None, :]
|
||||
offs_k = (
|
||||
token_id[:, None, None] * K_STRIDE_0
|
||||
+ head_id[None, :, None] * K_STRIDE_1
|
||||
+ rope_sub_id[None, None, :]
|
||||
+ QK_NOPE_HEAD_DIM
|
||||
)
|
||||
vals_rope = tl.load(k_rope_ptr + offs_rope, mask=token_mask[:, None, None])
|
||||
tl.store(k_ptr + offs_k, vals_rope, mask=token_mask[:, None, None])
|
||||
|
||||
|
||||
def fn_triton(k, k_nope, k_rope):
|
||||
assert k.device == DEVICE and k_nope.device == DEVICE and k_rope.device == DEVICE
|
||||
num_tokens, _, _ = k.shape
|
||||
grid = lambda meta: (triton.cdiv(num_tokens, meta["BLOCK_ROWS"]),)
|
||||
fn_triton_kernel[grid](
|
||||
k,
|
||||
k_nope,
|
||||
k_rope,
|
||||
num_tokens,
|
||||
QK_NOPE_HEAD_DIM=qk_nope_head_dim,
|
||||
QK_ROPE_HEAD_DIM=qk_rope_head_dim,
|
||||
NUM_LOCAL_HEADS=num_local_heads,
|
||||
K_NOPE_STRIDE_0=k_nope.stride(0),
|
||||
K_NOPE_STRIDE_1=k_nope.stride(1),
|
||||
K_STRIDE_0=k.stride(0),
|
||||
K_STRIDE_1=k.stride(1),
|
||||
K_ROPE_STRIDE_0=k_rope.stride(0),
|
||||
BLOCK_ROWS=16,
|
||||
)
|
||||
|
||||
|
||||
def execute_and_get_output(f, data):
|
||||
data["k"].zero_()
|
||||
f(**data)
|
||||
assert data["k"].sum().item() != 0
|
||||
return data["k"].clone()
|
||||
|
||||
|
||||
torch.manual_seed(0)
|
||||
data = create_data(num_tokens=32768)
|
||||
output_ref = execute_and_get_output(fn_torch, data)
|
||||
output_exp = execute_and_get_output(fn_cuda, data)
|
||||
# print(output_ref)
|
||||
# print(output_exp)
|
||||
if not torch.all(output_ref == output_exp):
|
||||
abs_delta = torch.abs(output_ref - output_exp)
|
||||
raise AssertionError(
|
||||
f"{output_ref=} {output_exp=} "
|
||||
f"{abs_delta=} "
|
||||
f"{torch.argwhere(abs_delta != 0.0)=} "
|
||||
)
|
||||
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["num_tokens"], # Argument names to use as an x-axis for the plot.
|
||||
x_vals=[
|
||||
2048,
|
||||
4096,
|
||||
8192,
|
||||
16384,
|
||||
32768,
|
||||
], # Different possible values for `x_name`.
|
||||
x_log=False, # x axis is logarithmic.
|
||||
line_arg="provider", # Argument name whose value corresponds to a different line in the plot.
|
||||
line_vals=[
|
||||
"torch",
|
||||
"torch_compiled",
|
||||
"triton",
|
||||
"hack_non_strided",
|
||||
"cuda",
|
||||
], # Possible values for `line_arg`.
|
||||
line_names=[
|
||||
"torch",
|
||||
"torch_compiled",
|
||||
"triton",
|
||||
"hack_non_strided",
|
||||
"cuda",
|
||||
], # Label name for the lines.
|
||||
plot_name="vector-add-performance", # Name for the plot. Used also as a file name for saving the plot.
|
||||
args={}, # Values for function arguments not in `x_names` and `y_name`.
|
||||
)
|
||||
)
|
||||
def benchmark(num_tokens, provider):
|
||||
data = create_data(num_tokens=num_tokens)
|
||||
quantiles = (0.5, 0.2, 0.8)
|
||||
fn = {
|
||||
"torch": fn_torch,
|
||||
"torch_compiled": fn_torch_compiled,
|
||||
"triton": fn_triton,
|
||||
"hack_non_strided": fn_hack_non_strided,
|
||||
"cuda": fn_cuda,
|
||||
}[provider]
|
||||
ms, min_ms, max_ms = run_bench(lambda: fn(**data), quantiles=quantiles)
|
||||
return ms, min_ms, max_ms
|
||||
|
||||
|
||||
torch.cuda.cudart().cudaProfilerStart()
|
||||
benchmark.run(print_data=True, show_plots=True)
|
||||
torch.cuda.cudart().cudaProfilerStop()
|
||||
102
third_party/sglang/benchmark/kernels/flashinfer_allreduce_fusion/README.md
vendored
Normal file
102
third_party/sglang/benchmark/kernels/flashinfer_allreduce_fusion/README.md
vendored
Normal file
@@ -0,0 +1,102 @@
|
||||
# FlashInfer Fused AllReduce + RMSNorm Benchmark
|
||||
|
||||
This benchmark script is modified from the [original implementation](https://github.com/vllm-project/vllm/blob/237e1fb887c7f5a579420fa0295097f24b006594/benchmarks/kernels/benchmark_fused_collective.py) by the vLLM community. It aims to compare the performance differences between FlashInfer fused operators in SGLang (trtllm_allreduce_fusion: AllReduce + Residual Add + RMSNorm + optional quantization) and conventional implementations (standard `tensor_model_parallel_all_reduce` + separate RMSNorm/quantization). Specifically, this script tests the timing performance of two implementation paths: 1) Standard AllReduce and RMSNorm executed separately; 2) FlashInfer's fused operator combining AllReduce, Residual Add, RMSNorm, and optional quantization operations.
|
||||
|
||||
This benchmark script helps us tune the ipc workspace size of the `flashinfer_allreduce_residual_rmsnorm` operator in SGLang and prepare for applications with FP8/FP4 quantized fused operators.
|
||||
|
||||
Script path: `benchmark/kernels/flashinfer_allreduce_fusion/benchmark_fused_collective.py`
|
||||
|
||||
## Feature Overview
|
||||
|
||||
- Compare average execution time (ms) and calculate speedup ratios for the following paths:
|
||||
- standard_allreduce_rmsnorm (Standard AllReduce + RMSNorm)
|
||||
- flashinfer_fused_allreduce_rmsnorm (Fused AllReduce + RMSNorm), including oneshot and twoshot modes
|
||||
- Optionally compare FP8/FP4 quantized fused paths with standard paths
|
||||
- Use CUDA Graph capture and batch replay to reduce measurement noise
|
||||
- Automatically select the faster "standard baseline" (native/compiled version) as the denominator for speedup calculation
|
||||
- Optionally export results in Markdown format
|
||||
|
||||
## Runtime Environment and Prerequisites
|
||||
|
||||
- At least 2 GPUs, and launch multi-process distributed training using `torchrun` (NCCL backend)
|
||||
- Properly install/compile sglang along with sgl-kernel and custom operators
|
||||
|
||||
## Quick Start (Command Examples)
|
||||
|
||||
The following examples use world_size=2. You can modify `--nproc_per_node` and parameters according to your machine:
|
||||
|
||||
- Regular paths only (no quantization):
|
||||
```
|
||||
torchrun --nproc_per_node=2 \
|
||||
benchmark/kernels/flashinfer_allreduce_fusion/benchmark_fused_collective.py \
|
||||
--no-quant --hidden-dim 1024 --seq-lens 512 1024 2048 4096 --trials 100
|
||||
```
|
||||
|
||||
- FP8 quantization paths only:
|
||||
```
|
||||
torchrun --nproc_per_node=2 \
|
||||
benchmark/kernels/flashinfer_allreduce_fusion/benchmark_fused_collective.py \
|
||||
--quant-fp8 --hidden-dim 1024 --seq-lens 512 1024 2048 4096 --trials 100
|
||||
```
|
||||
|
||||
- FP4 quantization paths only:
|
||||
```
|
||||
torchrun --nproc_per_node=2 \
|
||||
benchmark/kernels/flashinfer_allreduce_fusion/benchmark_fused_collective.py \
|
||||
--quant-fp4 --hidden-dim 1024 --seq-lens 512 1024 2048 4096 --trials 100
|
||||
```
|
||||
|
||||
- Larger hidden dimensions:
|
||||
```
|
||||
torchrun --nproc_per_node=2 \
|
||||
benchmark/kernels/flashinfer_allreduce_fusion/benchmark_fused_collective.py \
|
||||
--no-quant --hidden-dim 4096 --seq-lens 512 1024 2048 4096 --trials 100
|
||||
```
|
||||
|
||||
## Parameter Description
|
||||
- `--seq-lens`: List of sequence lengths to test (default: 128 512 1024 2048)
|
||||
- `--hidden-dim`: Hidden dimension (default: 8192)
|
||||
- `--dtypes`: Data type list, `float16|bfloat16|float32` (default: bfloat16)
|
||||
- `--no-residual`: Only test "no residual" scenarios (default tests both "with/without residual")
|
||||
- Mutually exclusive quantization options:
|
||||
- `--no-quant`: No quantization testing
|
||||
- `--quant-fp8`: Only FP8 quantization testing
|
||||
- `--quant-fp4`: Only FP4 quantization testing
|
||||
- `--quant-all`: Test all (default)
|
||||
- FlashInfer related:
|
||||
- `--disable-oneshot`: Disable oneshot mode (default enables oneshot and tests twoshot simultaneously)
|
||||
- Runtime configuration:
|
||||
- `--warmup`: Warmup count before graph capture and before graph replay (default 5)
|
||||
- `--trials`: Benchmark iteration count (default 20; internally each `graph.replay()` will batch replay multiple times)
|
||||
- `--output-file`: Save results as Markdown file (only rank0 takes effect)
|
||||
|
||||
## Output Example
|
||||
|
||||
Each configuration group prints a table showing average execution time and relative speedup ratios (baseline is the faster standard implementation). For example:
|
||||
```
|
||||
================================================================================
|
||||
Results: seq_len=1024, hidden_dim=1024
|
||||
dtype=torch.bfloat16, residual=yes, quant_mode=none
|
||||
================================================================================
|
||||
Operation Time (ms) Speedup
|
||||
--------------------------------------------------------------------------------
|
||||
standard_allreduce_rmsnorm 0.024 0.98x
|
||||
standard_allreduce_rmsnorm_native_compiled 0.023 baseline
|
||||
flashinfer_fused_allreduce_rmsnorm_oneshot 0.011 2.19x
|
||||
flashinfer_fused_allreduce_rmsnorm_twoshot 0.041 0.57x
|
||||
```
|
||||
|
||||
If `--output-file` is specified, all configurations will be summarized in Markdown tables in that file.
|
||||
|
||||
## Important Notes and Recommendations
|
||||
|
||||
- Distributed: The script uses `torchrun` environment variables to initialize distributed training and binds tensors/communication groups to the current rank's corresponding device.
|
||||
- World size: Requires `WORLD_SIZE > 1` to perform communication operator benchmarks. Otherwise, the script will error and prompt.
|
||||
- FlashInfer:
|
||||
- If not installed or interfaces are missing, the script will only run standard paths and provide prompts in the logs.
|
||||
- The fused operator internally uses "oneshot"/"twoshot" two trigger methods; oneshot is enabled by default and twoshot is tested simultaneously.
|
||||
- FP8/FP4:
|
||||
- FP8 uses sglang's FP8 tools and dtype, with underlying platform selection of `e4m3`/`e4m3fnuz` etc.
|
||||
- FP4 uses sgl-kernel's `scaled_fp4_quant`, requiring corresponding platform support.
|
||||
- CUDA Graph:
|
||||
- Uses sglang's `graph_capture()` to prepare capture-ready state for communication, then uses `torch.cuda.graph` to capture kernels, reducing measurement jitter.
|
||||
1305
third_party/sglang/benchmark/kernels/flashinfer_allreduce_fusion/benchmark_fused_collective.py
vendored
Normal file
1305
third_party/sglang/benchmark/kernels/flashinfer_allreduce_fusion/benchmark_fused_collective.py
vendored
Normal file
File diff suppressed because it is too large
Load Diff
210
third_party/sglang/benchmark/kernels/fused_moe_triton/README.md
vendored
Normal file
210
third_party/sglang/benchmark/kernels/fused_moe_triton/README.md
vendored
Normal file
@@ -0,0 +1,210 @@
|
||||
## Tuning Triton MoE Kernels
|
||||
|
||||
This directory contains benchmarking tools for MoE (Mixture of Experts) kernels.
|
||||
|
||||
### Overview
|
||||
|
||||
The tuning tools support both **Tensor Parallelism (TP)** and **Expert Parallelism (EP)** modes:
|
||||
|
||||
- **TP Mode**: Traditional tensor parallelism where intermediate layers are sharded across GPUs
|
||||
- **EP Mode**: Expert parallelism where experts are distributed across GPUs. Can be combined with TP mode (e.g., `--tp-size 8 --ep-size 2`)
|
||||
- **MLLM Support**: Multi-modal Large Language Models with text encoders (e.g., Llama4, Qwen3VL)
|
||||
|
||||
### Tuning Tools
|
||||
|
||||
#### 1. `tuning_fused_moe_triton.py`
|
||||
A unified tool for tuning the `fused_moe_triton` kernel. Adapted from [vllm's benchmark_moe.py](https://github.com/vllm-project/vllm/blob/main/benchmarks/kernels/benchmark_moe.py), with support for EP mode and various model architectures.
|
||||
|
||||
#### 2. `tuning_fused_moe_triton_sep.py`
|
||||
A specialized tool for separate kernel tuning, optimizing the first and second MoE kernels independently with TMA (Tensor Memory Accelerator) support.
|
||||
|
||||
### Usage Examples
|
||||
|
||||
#### Basic TP Mode Tuning
|
||||
```bash
|
||||
# Tune Mixtral-8x7B with default TP settings
|
||||
python benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py \
|
||||
--model mistralai/Mixtral-8x7B-Instruct-v0.1 \
|
||||
--tune
|
||||
|
||||
# Tune Qwen2-57B with FP8 and TP=4
|
||||
python benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py \
|
||||
--model Qwen/Qwen2-57B-A14B-Instruct \
|
||||
--tp-size 4 \
|
||||
--dtype fp8_w8a8 \
|
||||
--tune
|
||||
|
||||
# Tune DeepSeek-V3 with FP8 and TP=8
|
||||
python benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py \
|
||||
--model deepseek-ai/DeepSeek-V3-0324 \
|
||||
--tp-size 8 \
|
||||
--dtype fp8_w8a8 \
|
||||
--tune
|
||||
```
|
||||
|
||||
#### EP Mode Tuning (Expert Parallelism)
|
||||
**Note**: EP mode can be used alone or combined with TP mode. When using both, ensure `tp_size` is divisible by `ep_size`.
|
||||
|
||||
```bash
|
||||
# Tune Mixtral-8x7B with EP=2 only
|
||||
python benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py \
|
||||
--model mistralai/Mixtral-8x7B-Instruct-v0.1 \
|
||||
--tp-size 2 \
|
||||
--ep-size 2 \
|
||||
--tune
|
||||
|
||||
# Tune Qwen2-57B with TP=8 and EP=4 (combined mode)
|
||||
python benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py \
|
||||
--model Qwen/Qwen2-57B-A14B-Instruct \
|
||||
--tp-size 8 \
|
||||
--ep-size 4 \
|
||||
--dtype fp8_w8a8 \
|
||||
--tune
|
||||
```
|
||||
|
||||
#### MLLM Model Tuning (Multi-modal)
|
||||
```bash
|
||||
python benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py \
|
||||
--model Qwen/Qwen3-VL-30B-A3B-Instruct \
|
||||
--tp-size 2 \
|
||||
--tune
|
||||
```
|
||||
|
||||
#### Separate Kernel Tuning with `tuning_fused_moe_triton_sep.py`
|
||||
|
||||
This tool requires pre-generated topk_ids files and supports both TP and EP modes:
|
||||
|
||||
Edit the code file (such as srt/models/deepseek_v2.py) in the Python site package and add the logic for saving topk_ids:
|
||||
|
||||
```python
|
||||
# import get_tensor_model_parallel_rank
|
||||
# DeepseekV2MoE::forward_normal
|
||||
if hidden_states.shape[0] >= 4096 and get_tensor_model_parallel_rank() == 0:
|
||||
topk_ids_dir = xxxx
|
||||
if not hasattr(self, "save_idx"):
|
||||
self.save_idx = 0
|
||||
if self.save_idx <= 1:
|
||||
torch.save(topk_output.topk_ids, f"{topk_ids_dir}/topk_ids_layer{self.layer_id}_idx{self.save_idx}.pt")
|
||||
self.save_idx += 1
|
||||
```
|
||||
|
||||
Launch sglang server and send request using `benchmark/kernels/fused_moe_triton/tuning_client.py`
|
||||
```bash
|
||||
python benchmark/kernels/fused_moe_triton/tuning_client.py --port 8000
|
||||
```
|
||||
|
||||
```bash
|
||||
# TP Mode: Tune separate kernels with TP=4
|
||||
python benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton_sep.py \
|
||||
--model Qwen/Qwen2-57B-A14B-Instruct \
|
||||
--tp-size 4 \
|
||||
--topk-ids-dir /path/to/topk_ids \
|
||||
--tune
|
||||
|
||||
# EP Mode: Tune separate kernels with TP=4 and EP=2
|
||||
python benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton_sep.py \
|
||||
--model mistralai/Mixtral-8x7B-Instruct-v0.1 \
|
||||
--tp-size 4 \
|
||||
--ep-size 2 \
|
||||
--topk-ids-dir /path/to/topk_ids \
|
||||
--tune
|
||||
|
||||
# MLLM: Tune DeepSeek-V3 with separate kernels, TP=8 and EP=4
|
||||
python benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton_sep.py \
|
||||
--model deepseek-ai/DeepSeek-V3-0324 \
|
||||
--tp-size 8 \
|
||||
--ep-size 4 \
|
||||
--dtype fp8_w8a8 \
|
||||
--topk-ids-dir /path/to/topk_ids \
|
||||
--tune
|
||||
|
||||
# Benchmark specific config without tuning
|
||||
python benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton_sep.py \
|
||||
--model deepseek-ai/DeepSeek-V3-0324 \
|
||||
--tp-size 4 \
|
||||
--batch-size 1024 \
|
||||
--dtype fp8_w8a8 \
|
||||
--configs 128 256 128 16 8 4 \
|
||||
--topk-ids-dir /path/to/topk_ids
|
||||
```
|
||||
|
||||
#### Advanced Options
|
||||
```bash
|
||||
# Channel-wise quantization
|
||||
python benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py \
|
||||
--model meituan/DeepSeek-R1-Channel-INT8 \
|
||||
--tp-size 16 \
|
||||
--dtype int8_w8a8 \
|
||||
--per-channel-quant \
|
||||
--tune
|
||||
|
||||
# Specific batch size tuning
|
||||
python benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py \
|
||||
--model mistralai/Mixtral-8x7B-Instruct-v0.1 \
|
||||
--batch-size 2048 \
|
||||
--tune
|
||||
```
|
||||
|
||||
### Configuration Files
|
||||
|
||||
After tuning, configuration files will be generated:
|
||||
- **Standard tuning**: `E=64,N=640,device_name=NVIDIA_GeForce_RTX_4090,dtype=fp8_w8a8.json`
|
||||
- **Separate kernel tuning**: Two files for up/down kernels with TMA optimization flags
|
||||
|
||||
Move these files to `sglang/srt/layers/moe/fused_moe_triton/configs/triton_version/` directory to use them in SGLang.
|
||||
|
||||
### Supported Models
|
||||
|
||||
- **Mixtral**: mistralai/Mixtral-8x7B-Instruct-v0.1, mixtral-8x22b
|
||||
- **Qwen**: Qwen2-57B, Qwen3-235B, Qwen3VL (MLLM)
|
||||
- **DeepSeek**: DeepSeek-V2, DeepSeek-V3, DeepSeek-R1
|
||||
- **Llama**: Llama4-Vision (MLLM)
|
||||
- **DBRX**: databricks/dbrx-instruct
|
||||
- **Jamba**: ai21labs/AI21-Jamba
|
||||
- **Grok**: xai-org/grok-1
|
||||
- **GLM**: THUDM/glm-4-9b-chat
|
||||
- **Bailing**: Custom MoE models
|
||||
|
||||
### Parameters Reference
|
||||
|
||||
- `--model`: HuggingFace model name or local path
|
||||
- `--tp-size`: Tensor parallelism size (default: 2)
|
||||
- `--ep-size`: Expert parallelism size (default: 1, can be combined with TP mode, ensure tp_size is divisible by ep_size)
|
||||
- `--dtype`: Data type (`auto`, `fp8_w8a8`, `int8_w8a16`, `int8_w8a8`)
|
||||
- `--batch-size`: Specific batch size for tuning (optional)
|
||||
- `--tune`: Enable tuning mode
|
||||
- `--per-channel-quant`: Enable per-channel quantization
|
||||
- `--disable-shared-experts-fusion`: Disable shared expert fusion for some models
|
||||
- `--topk-ids-dir`: Directory containing pre-generated topk_ids (for sep tool only)
|
||||
- `--configs`: Manual config specification [BLOCK_M, BLOCK_N, BLOCK_K, GROUP_M, warps, stages]
|
||||
|
||||
### Performance Comparison Tool
|
||||
|
||||
- `benchmark_vllm_vs_sglang_fused_moe_triton.py`: A tool for comparing the performance of fused MoE kernels between vllm and sglang implementations. Supports various model architectures and data types.
|
||||
|
||||
Example usage:
|
||||
```bash
|
||||
# Compare with default settings (Mixtral model)
|
||||
python benchmark/kernels/fused_moe_triton/benchmark_vllm_vs_sglang_fused_moe_triton.py
|
||||
|
||||
# Compare with FP8 mode for Qwen2-57B
|
||||
python benchmark/kernels/fused_moe_triton/benchmark_vllm_vs_sglang_fused_moe_triton.py \
|
||||
--model Qwen/Qwen2-57B-A14B-Instruct \
|
||||
--use-fp8-w8a8
|
||||
|
||||
# Compare with custom TP size
|
||||
python benchmark/kernels/fused_moe_triton/benchmark_vllm_vs_sglang_fused_moe_triton.py \
|
||||
--model deepseek-ai/DeepSeek-V3-0324 \
|
||||
--tp-size 8
|
||||
|
||||
# Compare with custom TP size
|
||||
python benchmark/kernels/fused_moe_triton/benchmark_vllm_vs_sglang_fused_moe_triton.py \
|
||||
--model deepseek-ai/DeepSeek-V3-0324 \
|
||||
--tp-size 8
|
||||
```
|
||||
|
||||
The benchmark results will be saved as plots and data files in the specified output directory (default: `./configs/benchmark_ops/vllm_sglang_fused_moe/`).
|
||||
|
||||
- `benchmark_torch_compile_fused_moe.py`: A tool for benchmarking the performance of the fused MoE kernel with `torch.compile` and original fused MoE kernel.
|
||||
|
||||
Usage is similar to `benchmark_vllm_vs_sglang_fused_moe_triton.py`, note that `torch.compile` does not support `fp8_w8a8` and `int8_w8a8` fused_moe_kernel. Both tools now support EP mode with `--ep-size` parameter.
|
||||
250
third_party/sglang/benchmark/kernels/fused_moe_triton/benchmark_sglang_fused_moe_triton.py
vendored
Normal file
250
third_party/sglang/benchmark/kernels/fused_moe_triton/benchmark_sglang_fused_moe_triton.py
vendored
Normal file
@@ -0,0 +1,250 @@
|
||||
# python3 benchmark/kernels/fused_moe_triton/sglang_fused_moe_triton.py --model /DeepSeek-V3/ --tp-size 8
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
import triton
|
||||
from common_utils import get_model_config
|
||||
|
||||
from sglang.benchmark.bench_utils import run_bench
|
||||
from sglang.srt.distributed.parallel_state import (
|
||||
destroy_distributed_environment,
|
||||
destroy_model_parallel,
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
from sglang.srt.layers.moe.fused_moe_triton.fused_moe import (
|
||||
fused_moe as fused_moe_sglang,
|
||||
)
|
||||
from sglang.srt.layers.moe.fused_moe_triton.triton_kernels_moe import (
|
||||
triton_kernel_moe_forward,
|
||||
)
|
||||
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
|
||||
from sglang.srt.layers.moe.topk import (
|
||||
TopK,
|
||||
TopKConfig,
|
||||
TopKOutputFormat,
|
||||
select_experts,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
|
||||
|
||||
|
||||
def fused_moe_triton_api(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
):
|
||||
topk_op = TopK(
|
||||
top_k=topk,
|
||||
renormalize=False,
|
||||
use_grouped_topk=False,
|
||||
output_format=TopKOutputFormat.TRITON_KERNEL,
|
||||
)
|
||||
triton_topk_output = topk_op.forward_cuda(
|
||||
hidden_states=x,
|
||||
router_logits=input_gating,
|
||||
)
|
||||
|
||||
moe_runner_config = MoeRunnerConfig(
|
||||
inplace=False,
|
||||
)
|
||||
return triton_kernel_moe_forward(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
triton_topk_output,
|
||||
moe_runner_config,
|
||||
)
|
||||
|
||||
|
||||
def fused_moe_sglang_api(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
use_fp8_w8a8=False,
|
||||
w1_scale=None,
|
||||
w2_scale=None,
|
||||
a1_scale=None,
|
||||
a2_scale=None,
|
||||
block_shape=None,
|
||||
):
|
||||
topk_output = select_experts(
|
||||
hidden_states=x,
|
||||
router_logits=input_gating,
|
||||
topk_config=TopKConfig(top_k=topk, renormalize=False),
|
||||
)
|
||||
return fused_moe_sglang(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
topk_output,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
w1_scale=w1_scale,
|
||||
w2_scale=w2_scale,
|
||||
a1_scale=a1_scale,
|
||||
a2_scale=a2_scale,
|
||||
block_shape=block_shape,
|
||||
)
|
||||
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["batch_size"],
|
||||
x_vals=list([128, 256, 512, 1024, 2048, 4096, 8192]),
|
||||
line_arg="provider",
|
||||
line_vals=[
|
||||
"sglang_fused_moe_triton_v340",
|
||||
"sglang_fused_moe_triton",
|
||||
],
|
||||
line_names=[
|
||||
"sglang_fused_moe_triton_v340",
|
||||
"sglang_fused_moe_triton",
|
||||
],
|
||||
styles=[
|
||||
("blue", "-"),
|
||||
("green", "-"),
|
||||
],
|
||||
ylabel="Time (ms)",
|
||||
plot_name="fused-moe-performance",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(
|
||||
batch_size,
|
||||
provider,
|
||||
model_config,
|
||||
use_fp8_w8a8=False,
|
||||
use_cuda_graph: bool = False,
|
||||
):
|
||||
print(f"benchmark {provider} with batch_size={batch_size}")
|
||||
torch.set_default_device("cuda")
|
||||
torch.cuda.manual_seed_all(0)
|
||||
|
||||
num_tokens = batch_size
|
||||
num_experts = model_config["num_experts"]
|
||||
hidden_size = model_config["hidden_size"]
|
||||
shard_intermediate_size = model_config["shard_intermediate_size"]
|
||||
topk = model_config["topk"]
|
||||
dtype = model_config["dtype"]
|
||||
block_shape = model_config["block_shape"]
|
||||
|
||||
x = torch.randn(num_tokens, hidden_size, dtype=dtype)
|
||||
|
||||
w1 = torch.randn(num_experts, shard_intermediate_size, hidden_size, dtype=dtype)
|
||||
w2 = torch.randn(
|
||||
num_experts, hidden_size, shard_intermediate_size // 2, dtype=dtype
|
||||
)
|
||||
|
||||
w1_tri = w1.clone()
|
||||
w2_tri = w2.clone()
|
||||
w1_tri = w1_tri.transpose(-2, -1).contiguous()
|
||||
w2_tri = w2_tri.transpose(-2, -1).contiguous()
|
||||
|
||||
input_gating = torch.randn(num_tokens, num_experts, dtype=torch.float32)
|
||||
|
||||
if provider == "sglang_fused_moe_triton_v340":
|
||||
api_func = fused_moe_triton_api
|
||||
api_kwargs = {
|
||||
"x": x,
|
||||
"w1": w1_tri,
|
||||
"w2": w2_tri,
|
||||
"input_gating": input_gating,
|
||||
"topk": topk,
|
||||
}
|
||||
else:
|
||||
api_func = fused_moe_sglang_api
|
||||
api_kwargs = {
|
||||
"x": x,
|
||||
"w1": w1,
|
||||
"w2": w2,
|
||||
"input_gating": input_gating,
|
||||
"topk": topk,
|
||||
"use_fp8_w8a8": use_fp8_w8a8,
|
||||
"block_shape": block_shape,
|
||||
}
|
||||
|
||||
# Warmup
|
||||
for _ in range(10):
|
||||
_ = api_func(**api_kwargs)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
if use_cuda_graph:
|
||||
stream = torch.cuda.Stream()
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph, stream=stream):
|
||||
api_func(**api_kwargs)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
bench_lambda = lambda: graph.replay()
|
||||
else:
|
||||
bench_lambda = lambda: api_func(**api_kwargs)
|
||||
|
||||
quantiles = (0.5, 0.2, 0.8)
|
||||
ms, min_ms, max_ms = run_bench(bench_lambda, quantiles=quantiles)
|
||||
return ms, min_ms, max_ms
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--model", type=str, default="mistralai/Mixtral-8x7B-Instruct-v0.1"
|
||||
)
|
||||
parser.add_argument("--tp-size", "--tp", type=int, default=2)
|
||||
parser.add_argument("--ep-size", "--ep", type=int, default=1)
|
||||
parser.add_argument("--use-fp8-w8a8", action="store_true")
|
||||
parser.add_argument(
|
||||
"--use-cuda-graph", action="store_true", help="Enable CUDA Graph capture/replay"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save-path",
|
||||
type=str,
|
||||
default="./configs/benchmark_ops/sglang_fused_moe/",
|
||||
)
|
||||
parser.add_argument("--trust-remote-code", action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
# Initialize global server args (required by SGLang MoE kernels)
|
||||
server_args = ServerArgs(model_path=args.model)
|
||||
set_global_server_args_for_scheduler(server_args)
|
||||
|
||||
try:
|
||||
if not torch.distributed.is_initialized():
|
||||
torch.distributed.init_process_group(
|
||||
backend="nccl" if torch.cuda.is_available() else "gloo",
|
||||
init_method="tcp://127.0.0.1:23456",
|
||||
world_size=1,
|
||||
rank=0,
|
||||
)
|
||||
|
||||
init_distributed_environment(
|
||||
world_size=1,
|
||||
rank=0,
|
||||
distributed_init_method="tcp://127.0.0.1:23456",
|
||||
local_rank=0,
|
||||
backend="nccl" if torch.cuda.is_available() else "gloo",
|
||||
)
|
||||
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=1,
|
||||
expert_model_parallel_size=1,
|
||||
)
|
||||
|
||||
model_config = get_model_config(args.model, args.tp_size, args.ep_size)
|
||||
benchmark.run(
|
||||
show_plots=True,
|
||||
print_data=True,
|
||||
save_path=args.save_path,
|
||||
model_config=model_config,
|
||||
use_fp8_w8a8=args.use_fp8_w8a8,
|
||||
use_cuda_graph=args.use_cuda_graph,
|
||||
)
|
||||
finally:
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
306
third_party/sglang/benchmark/kernels/fused_moe_triton/benchmark_torch_compile_fused_moe.py
vendored
Normal file
306
third_party/sglang/benchmark/kernels/fused_moe_triton/benchmark_torch_compile_fused_moe.py
vendored
Normal file
@@ -0,0 +1,306 @@
|
||||
# python3 benchmark/kernels/fused_moe_triton/benchmark_torch_compile_fused_moe.py --model /DeepSeek-V3/ --tp-size 8 --use-fp8-w8a8
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
import triton
|
||||
from torch.nn import functional as F
|
||||
from transformers import AutoConfig
|
||||
|
||||
from sglang.benchmark.bench_utils import run_bench
|
||||
from sglang.srt.layers.moe.fused_moe_triton.fused_moe import (
|
||||
fused_moe as fused_moe_triton,
|
||||
)
|
||||
from sglang.srt.model_executor.cuda_graph_runner import set_torch_compile_config
|
||||
|
||||
|
||||
def get_model_config(model_name: str, tp_size: int):
|
||||
"""Get model configuration parameters"""
|
||||
config = AutoConfig.from_pretrained(model_name, trust_remote_code=True)
|
||||
|
||||
if config.architectures[0] == "DbrxForCausalLM":
|
||||
E = config.ffn_config.moe_num_experts
|
||||
topk = config.ffn_config.moe_top_k
|
||||
intermediate_size = config.ffn_config.ffn_hidden_size
|
||||
shard_intermediate_size = 2 * intermediate_size // tp_size
|
||||
elif config.architectures[0] == "JambaForCausalLM":
|
||||
E = config.num_experts
|
||||
topk = config.num_experts_per_tok
|
||||
intermediate_size = config.intermediate_size
|
||||
shard_intermediate_size = 2 * intermediate_size // tp_size
|
||||
elif config.architectures[0] == "Qwen2MoeForCausalLM":
|
||||
E = config.num_experts
|
||||
topk = config.num_experts_per_tok
|
||||
intermediate_size = config.moe_intermediate_size
|
||||
shard_intermediate_size = 2 * intermediate_size // tp_size
|
||||
elif config.architectures[0] == "Qwen3MoeForCausalLM":
|
||||
E = config.n_routed_experts
|
||||
topk = config.num_experts_per_tok
|
||||
intermediate_size = config.moe_intermediate_size
|
||||
shard_intermediate_size = 2 * intermediate_size // tp_size
|
||||
elif config.architectures[0] in ["DeepseekV2ForCausalLM", "DeepseekV3ForCausalLM"]:
|
||||
E = config.n_routed_experts
|
||||
topk = config.num_experts_per_tok
|
||||
intermediate_size = config.moe_intermediate_size
|
||||
shard_intermediate_size = 2 * intermediate_size // tp_size
|
||||
elif config.architectures[0] == "Llama4ForConditionalGeneration":
|
||||
E = config.text_config.num_local_experts
|
||||
topk = config.text_config.num_experts_per_tok
|
||||
intermediate_size = config.text_config.intermediate_size
|
||||
shard_intermediate_size = 2 * intermediate_size // tp_size
|
||||
elif config.architectures[0] in [
|
||||
"Grok1ForCausalLM",
|
||||
"Grok1ImgGen",
|
||||
"Grok1AForCausalLM",
|
||||
]:
|
||||
E = config.num_local_experts
|
||||
topk = config.num_experts_per_tok
|
||||
intermediate_size = config.moe_intermediate_size
|
||||
shard_intermediate_size = 2 * intermediate_size // tp_size
|
||||
else:
|
||||
# Default: Mixtral
|
||||
E = config.num_local_experts
|
||||
topk = config.num_experts_per_tok
|
||||
intermediate_size = config.intermediate_size
|
||||
shard_intermediate_size = 2 * intermediate_size // tp_size
|
||||
|
||||
shape_configs = {
|
||||
"num_experts": E,
|
||||
"topk": topk,
|
||||
"hidden_size": config.hidden_size,
|
||||
"shard_intermediate_size": shard_intermediate_size,
|
||||
"dtype": config.torch_dtype,
|
||||
}
|
||||
print(f"{shape_configs=}")
|
||||
return shape_configs
|
||||
|
||||
|
||||
def fused_topk_native(
|
||||
hidden_states: torch.Tensor,
|
||||
gating_output: torch.Tensor,
|
||||
topk: int,
|
||||
renormalize: bool,
|
||||
):
|
||||
assert hidden_states.shape[0] == gating_output.shape[0], "Number of tokens mismatch"
|
||||
M, _ = hidden_states.shape
|
||||
topk_weights = torch.empty(
|
||||
M, topk, dtype=torch.float32, device=hidden_states.device
|
||||
)
|
||||
topk_ids = torch.empty(M, topk, dtype=torch.int32, device=hidden_states.device)
|
||||
topk_weights = F.softmax(gating_output.float(), dim=-1)
|
||||
topk_weights, topk_ids = torch.topk(topk_weights, topk, dim=-1)
|
||||
if renormalize:
|
||||
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||
return topk_weights, topk_ids
|
||||
|
||||
|
||||
@torch.compile(dynamic=False)
|
||||
def fused_moe_torch(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
use_fp8_w8a8=False,
|
||||
w1_scale=None,
|
||||
w2_scale=None,
|
||||
a1_scale=None,
|
||||
a2_scale=None,
|
||||
) -> torch.Tensor:
|
||||
assert not use_fp8_w8a8, "Fp8_w8a8 fused_moe is not supported for torch compile"
|
||||
|
||||
topk_weights, topk_ids = fused_topk_native(
|
||||
hidden_states=x,
|
||||
gating_output=input_gating,
|
||||
topk=topk,
|
||||
renormalize=True,
|
||||
)
|
||||
w13_weights = w1[topk_ids]
|
||||
w1_weights, w3_weights = torch.chunk(w13_weights, 2, dim=2)
|
||||
w2_weights = w2[topk_ids]
|
||||
x1 = torch.einsum("ti,taoi -> tao", x, w1_weights)
|
||||
x1 = F.silu(x1)
|
||||
x3 = torch.einsum("ti, taoi -> tao", x, w3_weights)
|
||||
expert_outs = torch.einsum("tao, taio -> tai", (x1 * x3), w2_weights)
|
||||
return torch.einsum("tai,ta -> ti", expert_outs, topk_weights.to(expert_outs.dtype))
|
||||
|
||||
|
||||
def fused_moe_torch_compile(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
use_fp8_w8a8=False,
|
||||
w1_scale=None,
|
||||
w2_scale=None,
|
||||
a1_scale=None,
|
||||
a2_scale=None,
|
||||
):
|
||||
return fused_moe_torch(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
w1_scale=w1_scale,
|
||||
w2_scale=w2_scale,
|
||||
a1_scale=a1_scale,
|
||||
a2_scale=a2_scale,
|
||||
)
|
||||
|
||||
|
||||
def fused_moe_sglang_api(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
use_fp8_w8a8=False,
|
||||
w1_scale=None,
|
||||
w2_scale=None,
|
||||
a1_scale=None,
|
||||
a2_scale=None,
|
||||
):
|
||||
return fused_moe_triton(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
renormalize=True,
|
||||
inplace=True,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
w1_scale=w1_scale,
|
||||
w2_scale=w2_scale,
|
||||
a1_scale=a1_scale,
|
||||
a2_scale=a2_scale,
|
||||
)
|
||||
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["batch_size"],
|
||||
x_vals=list(range(1, 5)),
|
||||
line_arg="provider",
|
||||
line_vals=[
|
||||
"fused_moe_triton",
|
||||
"fused_moe_torch_compile",
|
||||
],
|
||||
line_names=[
|
||||
"fused_moe_triton",
|
||||
"fused_moe_torch_compile",
|
||||
],
|
||||
styles=[
|
||||
("blue", "-"),
|
||||
("green", "-"),
|
||||
],
|
||||
ylabel="Time (ms)",
|
||||
plot_name="fused-moe-performance",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(batch_size, provider, model_config, use_fp8_w8a8=False):
|
||||
print(f"benchmark {provider} with batch_size={batch_size}")
|
||||
torch.set_default_device("cuda")
|
||||
torch.cuda.manual_seed_all(0)
|
||||
set_torch_compile_config()
|
||||
|
||||
num_tokens = batch_size
|
||||
num_experts = model_config["num_experts"]
|
||||
hidden_size = model_config["hidden_size"]
|
||||
shard_intermediate_size = model_config["shard_intermediate_size"]
|
||||
topk = model_config["topk"]
|
||||
dtype = model_config["dtype"]
|
||||
|
||||
x = torch.randn(num_tokens, hidden_size, dtype=dtype)
|
||||
|
||||
if use_fp8_w8a8:
|
||||
init_dtype = dtype
|
||||
w1 = torch.randn(
|
||||
num_experts, shard_intermediate_size, hidden_size, dtype=init_dtype
|
||||
)
|
||||
w2 = torch.randn(
|
||||
num_experts, hidden_size, shard_intermediate_size // 2, dtype=init_dtype
|
||||
)
|
||||
w1 = w1.to(torch.float8_e4m3fn)
|
||||
w2 = w2.to(torch.float8_e4m3fn)
|
||||
w1_scale = torch.randn(num_experts, dtype=torch.float32)
|
||||
w2_scale = torch.randn(num_experts, dtype=torch.float32)
|
||||
a1_scale = torch.randn(1, dtype=torch.float32)
|
||||
a2_scale = torch.randn(1, dtype=torch.float32)
|
||||
else:
|
||||
w1 = torch.randn(num_experts, shard_intermediate_size, hidden_size, dtype=dtype)
|
||||
w2 = torch.randn(
|
||||
num_experts, hidden_size, shard_intermediate_size // 2, dtype=dtype
|
||||
)
|
||||
w1_scale = w2_scale = a1_scale = a2_scale = None
|
||||
|
||||
input_gating = torch.randn(num_tokens, num_experts, dtype=torch.float32)
|
||||
|
||||
# Warmup
|
||||
api_func = (
|
||||
fused_moe_torch_compile
|
||||
if provider == "fused_moe_torch_compile"
|
||||
else fused_moe_sglang_api
|
||||
)
|
||||
for _ in range(10):
|
||||
y = api_func(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
w1_scale=w1_scale,
|
||||
w2_scale=w2_scale,
|
||||
a1_scale=a1_scale,
|
||||
a2_scale=a2_scale,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
quantiles = (0.5, 0.2, 0.8)
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: api_func(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
w1_scale=w1_scale,
|
||||
w2_scale=w2_scale,
|
||||
a1_scale=a1_scale,
|
||||
a2_scale=a2_scale,
|
||||
)[0],
|
||||
quantiles=quantiles,
|
||||
)
|
||||
return ms, min_ms, max_ms
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--model", type=str, default="mistralai/Mixtral-8x7B-Instruct-v0.1"
|
||||
)
|
||||
parser.add_argument("--tp-size", type=int, default=2)
|
||||
parser.add_argument("--use-fp8-w8a8", action="store_true")
|
||||
parser.add_argument(
|
||||
"--save-path",
|
||||
type=str,
|
||||
default="./configs/benchmark_ops/fused_moe_torch_compile/",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
model_config = get_model_config(args.model, args.tp_size)
|
||||
benchmark.run(
|
||||
show_plots=True,
|
||||
print_data=True,
|
||||
save_path=args.save_path,
|
||||
model_config=model_config,
|
||||
use_fp8_w8a8=args.use_fp8_w8a8,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
265
third_party/sglang/benchmark/kernels/fused_moe_triton/benchmark_vllm_vs_sglang_fused_moe_triton.py
vendored
Normal file
265
third_party/sglang/benchmark/kernels/fused_moe_triton/benchmark_vllm_vs_sglang_fused_moe_triton.py
vendored
Normal file
@@ -0,0 +1,265 @@
|
||||
# python3 benchmark/kernels/fused_moe_triton/benchmark_vllm_vs_sglang_fused_moe_triton.py --model /DeepSeek-V3/ --tp-size 8 --use-fp8-w8a8
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
import triton
|
||||
from vllm.model_executor.layers.fused_moe.fused_moe import fused_moe as fused_moe_vllm
|
||||
|
||||
from sglang.benchmark.bench_utils import run_bench
|
||||
from sglang.srt.distributed.parallel_state import (
|
||||
destroy_distributed_environment,
|
||||
destroy_model_parallel,
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
from sglang.srt.layers.moe.fused_moe_triton.fused_moe import (
|
||||
fused_moe as fused_moe_sglang,
|
||||
)
|
||||
|
||||
from .common_utils import get_model_config
|
||||
|
||||
|
||||
def fused_moe_vllm_api(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
use_fp8_w8a8=False,
|
||||
w1_scale=None,
|
||||
w2_scale=None,
|
||||
a1_scale=None,
|
||||
a2_scale=None,
|
||||
block_shape=None,
|
||||
):
|
||||
if block_shape is not None:
|
||||
return fused_moe_vllm(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
renormalize=True,
|
||||
inplace=True,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
w1_scale=w1_scale,
|
||||
w2_scale=w2_scale,
|
||||
a1_scale=a1_scale,
|
||||
a2_scale=a2_scale,
|
||||
block_shape=block_shape,
|
||||
)
|
||||
else:
|
||||
return fused_moe_vllm(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
renormalize=True,
|
||||
inplace=True,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
w1_scale=w1_scale,
|
||||
w2_scale=w2_scale,
|
||||
a1_scale=a1_scale,
|
||||
a2_scale=a2_scale,
|
||||
)
|
||||
|
||||
|
||||
def fused_moe_sglang_api(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
use_fp8_w8a8=False,
|
||||
w1_scale=None,
|
||||
w2_scale=None,
|
||||
a1_scale=None,
|
||||
a2_scale=None,
|
||||
block_shape=None,
|
||||
):
|
||||
return fused_moe_sglang(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
renormalize=True,
|
||||
inplace=True,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
w1_scale=w1_scale,
|
||||
w2_scale=w2_scale,
|
||||
a1_scale=a1_scale,
|
||||
a2_scale=a2_scale,
|
||||
block_shape=block_shape,
|
||||
)
|
||||
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["batch_size"],
|
||||
x_vals=list(range(1, 513)),
|
||||
line_arg="provider",
|
||||
line_vals=[
|
||||
"vllm_fused_moe_triton",
|
||||
"sglang_fused_moe_triton",
|
||||
],
|
||||
line_names=[
|
||||
"vllm_fused_moe_triton",
|
||||
"sglang_fused_moe_triton",
|
||||
],
|
||||
styles=[
|
||||
("blue", "-"),
|
||||
("green", "-"),
|
||||
],
|
||||
ylabel="Time (ms)",
|
||||
plot_name="fused-moe-performance",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(batch_size, provider, model_config, use_fp8_w8a8=False):
|
||||
print(f"benchmark {provider} with batch_size={batch_size}")
|
||||
torch.set_default_device("cuda")
|
||||
torch.cuda.manual_seed_all(0)
|
||||
|
||||
num_tokens = batch_size
|
||||
num_experts = model_config["num_experts"]
|
||||
hidden_size = model_config["hidden_size"]
|
||||
shard_intermediate_size = model_config["shard_intermediate_size"]
|
||||
topk = model_config["topk"]
|
||||
dtype = model_config["dtype"]
|
||||
block_shape = model_config["block_shape"]
|
||||
|
||||
x = torch.randn(num_tokens, hidden_size, dtype=dtype)
|
||||
w1_scale = w2_scale = a1_scale = a2_scale = None
|
||||
|
||||
if use_fp8_w8a8:
|
||||
init_dtype = dtype
|
||||
w1 = torch.randn(
|
||||
num_experts, shard_intermediate_size, hidden_size, dtype=init_dtype
|
||||
)
|
||||
w2 = torch.randn(
|
||||
num_experts, hidden_size, shard_intermediate_size // 2, dtype=init_dtype
|
||||
)
|
||||
w1 = w1.to(torch.float8_e4m3fn)
|
||||
w2 = w2.to(torch.float8_e4m3fn)
|
||||
|
||||
if block_shape is None:
|
||||
w1_scale = torch.randn(num_experts, dtype=torch.float32)
|
||||
w2_scale = torch.randn(num_experts, dtype=torch.float32)
|
||||
a1_scale = torch.randn(1, dtype=torch.float32)
|
||||
a2_scale = torch.randn(1, dtype=torch.float32)
|
||||
else:
|
||||
block_n, block_k = block_shape[0], block_shape[1]
|
||||
n_tiles_w1 = (shard_intermediate_size + block_n - 1) // block_n
|
||||
n_tiles_w2 = (hidden_size + block_n - 1) // block_n
|
||||
k_tiles_w1 = (hidden_size + block_k - 1) // block_k
|
||||
k_tiles_w2 = (shard_intermediate_size // 2 + block_k - 1) // block_k
|
||||
w1_scale = torch.rand(
|
||||
(num_experts, n_tiles_w1, k_tiles_w1), dtype=torch.float32
|
||||
)
|
||||
w2_scale = torch.rand(
|
||||
(num_experts, n_tiles_w2, k_tiles_w2), dtype=torch.float32
|
||||
)
|
||||
else:
|
||||
w1 = torch.randn(num_experts, shard_intermediate_size, hidden_size, dtype=dtype)
|
||||
w2 = torch.randn(
|
||||
num_experts, hidden_size, shard_intermediate_size // 2, dtype=dtype
|
||||
)
|
||||
|
||||
input_gating = torch.randn(num_tokens, num_experts, dtype=torch.float32)
|
||||
|
||||
# Warmup
|
||||
api_func = (
|
||||
fused_moe_vllm_api
|
||||
if provider == "vllm_fused_moe_triton"
|
||||
else fused_moe_sglang_api
|
||||
)
|
||||
for _ in range(10):
|
||||
y = api_func(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
w1_scale=w1_scale,
|
||||
w2_scale=w2_scale,
|
||||
a1_scale=a1_scale,
|
||||
a2_scale=a2_scale,
|
||||
block_shape=block_shape,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
quantiles = (0.5, 0.2, 0.8)
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: api_func(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
input_gating,
|
||||
topk,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
w1_scale=w1_scale,
|
||||
w2_scale=w2_scale,
|
||||
a1_scale=a1_scale,
|
||||
a2_scale=a2_scale,
|
||||
block_shape=block_shape,
|
||||
)[0],
|
||||
quantiles=quantiles,
|
||||
)
|
||||
return ms, min_ms, max_ms
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--model", type=str, default="mistralai/Mixtral-8x7B-Instruct-v0.1"
|
||||
)
|
||||
parser.add_argument("--tp-size", "--tp", type=int, default=2)
|
||||
parser.add_argument("--ep-size", "--ep", type=int, default=1)
|
||||
parser.add_argument("--use-fp8-w8a8", action="store_true")
|
||||
parser.add_argument(
|
||||
"--save-path",
|
||||
type=str,
|
||||
default="./configs/benchmark_ops/vllm_sglang_fused_moe/",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
try:
|
||||
if not torch.distributed.is_initialized():
|
||||
torch.distributed.init_process_group(
|
||||
backend="nccl" if torch.cuda.is_available() else "gloo",
|
||||
init_method="tcp://127.0.0.1:23456",
|
||||
world_size=1,
|
||||
rank=0,
|
||||
)
|
||||
|
||||
init_distributed_environment(
|
||||
world_size=1,
|
||||
rank=0,
|
||||
distributed_init_method="tcp://127.0.0.1:23456",
|
||||
local_rank=0,
|
||||
backend="nccl" if torch.cuda.is_available() else "gloo",
|
||||
)
|
||||
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=1,
|
||||
pipeline_model_parallel_size=1,
|
||||
)
|
||||
|
||||
shape_configs = get_model_config(args.model, args.tp_size, args.ep_size)
|
||||
benchmark.run(
|
||||
show_plots=True,
|
||||
print_data=True,
|
||||
save_path=args.save_path,
|
||||
model_config=shape_configs,
|
||||
use_fp8_w8a8=args.use_fp8_w8a8,
|
||||
)
|
||||
finally:
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
288
third_party/sglang/benchmark/kernels/fused_moe_triton/common_utils.py
vendored
Normal file
288
third_party/sglang/benchmark/kernels/fused_moe_triton/common_utils.py
vendored
Normal file
@@ -0,0 +1,288 @@
|
||||
import json
|
||||
from typing import Dict, List, TypedDict
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.moe.fused_moe_triton.fused_moe import get_config_dtype_str
|
||||
from sglang.srt.layers.moe.fused_moe_triton.fused_moe_triton_config import (
|
||||
get_config_file_name,
|
||||
)
|
||||
from sglang.srt.utils import is_hip
|
||||
from sglang.srt.utils.hf_transformers_utils import get_config
|
||||
|
||||
|
||||
class BenchmarkConfig(TypedDict):
|
||||
BLOCK_SIZE_M: int
|
||||
BLOCK_SIZE_N: int
|
||||
BLOCK_SIZE_K: int
|
||||
GROUP_SIZE_M: int
|
||||
num_warps: int
|
||||
num_stages: int
|
||||
|
||||
|
||||
def calculate_shard_intermediate_size(
|
||||
intermediate_size: int, tp_size: int, ep_size: int = 1
|
||||
) -> int:
|
||||
assert tp_size % ep_size == 0
|
||||
moe_tp_size = tp_size // ep_size
|
||||
assert intermediate_size % moe_tp_size == 0
|
||||
return 2 * intermediate_size // moe_tp_size
|
||||
|
||||
|
||||
def get_model_config(
|
||||
model_name: str,
|
||||
tp_size: int,
|
||||
ep_size: int = 1,
|
||||
disable_shared_experts_fusion: bool = False,
|
||||
topk_ids_dir: str = None,
|
||||
) -> Dict:
|
||||
config = get_config(model_name, trust_remote_code=True)
|
||||
architecture = config.architectures[0]
|
||||
block_shape = None
|
||||
if (
|
||||
hasattr(config, "quantization_config")
|
||||
and "weight_block_size" in config.quantization_config
|
||||
):
|
||||
block_shape = config.quantization_config["weight_block_size"]
|
||||
assert len(block_shape) == 2
|
||||
|
||||
if (
|
||||
hasattr(config, "quantization_config")
|
||||
and "config_groups" in config.quantization_config
|
||||
):
|
||||
config_groups = config.quantization_config["config_groups"]
|
||||
# Get group_size from the first group's weights config
|
||||
first_group = next(iter(config_groups.values()), {})
|
||||
weights_config = first_group.get("weights", {})
|
||||
group_size = weights_config.get("group_size")
|
||||
block_shape = [0, group_size]
|
||||
assert len(block_shape) == 2
|
||||
# Replace config with text_config for encoder-decoder models after getting block_shape and architecture
|
||||
if hasattr(config, "text_config"):
|
||||
config = config.get_text_config()
|
||||
|
||||
hidden_size = config.hidden_size
|
||||
if architecture == "DbrxForCausalLM":
|
||||
E = config.ffn_config.moe_num_experts // ep_size
|
||||
topk = config.ffn_config.moe_top_k
|
||||
intermediate_size = config.ffn_config.ffn_hidden_size
|
||||
elif architecture == "JambaForCausalLM":
|
||||
E = config.num_experts // ep_size
|
||||
topk = config.num_experts_per_tok
|
||||
intermediate_size = config.intermediate_size
|
||||
elif architecture in [
|
||||
"Qwen2MoeForCausalLM",
|
||||
"Qwen3MoeForCausalLM",
|
||||
"Qwen3NextForCausalLM",
|
||||
"Qwen3VLMoeForConditionalGeneration",
|
||||
"Qwen3_5MoeForConditionalGeneration",
|
||||
]:
|
||||
E = config.num_experts // ep_size
|
||||
topk = config.num_experts_per_tok
|
||||
intermediate_size = config.moe_intermediate_size
|
||||
elif architecture in [
|
||||
"DeepseekV2ForCausalLM",
|
||||
"DeepseekV3ForCausalLM",
|
||||
"DeepseekV32ForCausalLM",
|
||||
"Glm4MoeForCausalLM",
|
||||
"GlmMoeDsaForCausalLM",
|
||||
"MistralLarge3ForCausalLM",
|
||||
]:
|
||||
E = (config.n_routed_experts // ep_size) + (
|
||||
0
|
||||
if disable_shared_experts_fusion
|
||||
or architecture
|
||||
not in [
|
||||
"DeepseekV3ForCausalLM",
|
||||
"DeepseekV32ForCausalLM",
|
||||
"Glm4MoeForCausalLM",
|
||||
"GlmMoeDsaForCausalLM",
|
||||
"MistralLarge3ForCausalLM",
|
||||
]
|
||||
else 1
|
||||
)
|
||||
topk = config.num_experts_per_tok + (
|
||||
0 if disable_shared_experts_fusion or topk_ids_dir is None else 1
|
||||
)
|
||||
intermediate_size = config.moe_intermediate_size
|
||||
elif architecture == "Llama4ForConditionalGeneration":
|
||||
E = config.num_local_experts // ep_size + (
|
||||
0 if disable_shared_experts_fusion else 1
|
||||
)
|
||||
topk = config.num_experts_per_tok + (
|
||||
0 if disable_shared_experts_fusion or topk_ids_dir is None else 1
|
||||
)
|
||||
intermediate_size = config.intermediate_size
|
||||
elif architecture in [
|
||||
"Grok1ForCausalLM",
|
||||
"Grok1ImgGen",
|
||||
"Grok1AForCausalLM",
|
||||
]:
|
||||
E = config.num_local_experts // ep_size
|
||||
topk = config.num_experts_per_tok
|
||||
intermediate_size = config.moe_intermediate_size
|
||||
elif architecture in [
|
||||
"BailingMoEForCausalLM",
|
||||
"BailingMoeForCausalLM",
|
||||
"BailingMoeV2ForCausalLM",
|
||||
]:
|
||||
E = config.num_experts // ep_size
|
||||
topk = config.num_experts_per_tok
|
||||
intermediate_size = config.moe_intermediate_size
|
||||
elif architecture == "NemotronHForCausalLM":
|
||||
E = config.n_routed_experts // ep_size
|
||||
topk = config.num_experts_per_tok
|
||||
intermediate_size = config.moe_intermediate_size
|
||||
hidden_size = getattr(config, "moe_latent_size", None) or hidden_size
|
||||
else:
|
||||
# Default: Mixtral
|
||||
E = config.num_local_experts // ep_size
|
||||
topk = config.num_experts_per_tok
|
||||
intermediate_size = config.intermediate_size
|
||||
|
||||
shard_intermediate_size = calculate_shard_intermediate_size(
|
||||
intermediate_size, tp_size, ep_size
|
||||
)
|
||||
|
||||
return {
|
||||
"num_experts": E,
|
||||
"topk": topk,
|
||||
"hidden_size": hidden_size,
|
||||
"shard_intermediate_size": shard_intermediate_size,
|
||||
"dtype": config.torch_dtype,
|
||||
"block_shape": block_shape,
|
||||
"architecture": architecture,
|
||||
}
|
||||
|
||||
|
||||
def get_rocm_configs_compute_bound() -> List[Dict[str, int]]:
|
||||
configs: List[BenchmarkConfig] = []
|
||||
waves_per_eu_range = 0
|
||||
for num_stages in [2]:
|
||||
for block_m in [32, 64, 128, 256]:
|
||||
for block_k in [32, 64, 128, 256]:
|
||||
for block_n in [16, 32, 64, 128, 256]:
|
||||
for num_warps in [1, 2, 4, 8]:
|
||||
for group_size in [1, 4, 8, 16, 32]:
|
||||
configs.append(
|
||||
{
|
||||
"BLOCK_SIZE_M": block_m,
|
||||
"BLOCK_SIZE_N": block_n,
|
||||
"BLOCK_SIZE_K": block_k,
|
||||
"GROUP_SIZE_M": group_size,
|
||||
"num_warps": num_warps,
|
||||
"num_stages": num_stages,
|
||||
"waves_per_eu": waves_per_eu_range,
|
||||
}
|
||||
)
|
||||
return configs
|
||||
|
||||
|
||||
def get_configs_compute_bound() -> List[Dict[str, int]]:
|
||||
configs: List[BenchmarkConfig] = []
|
||||
if is_hip():
|
||||
configs = get_rocm_configs_compute_bound()
|
||||
else:
|
||||
for num_stages in [2, 3, 4, 5]:
|
||||
for block_m in [16, 32, 64, 128, 256]:
|
||||
for block_k in [64, 128, 256]:
|
||||
for block_n in [32, 64, 128, 256]:
|
||||
for num_warps in [4, 8]:
|
||||
for group_size in [1, 16, 32, 64]:
|
||||
configs.append(
|
||||
{
|
||||
"BLOCK_SIZE_M": block_m,
|
||||
"BLOCK_SIZE_N": block_n,
|
||||
"BLOCK_SIZE_K": block_k,
|
||||
"GROUP_SIZE_M": group_size,
|
||||
"num_warps": num_warps,
|
||||
"num_stages": num_stages,
|
||||
}
|
||||
)
|
||||
return configs
|
||||
|
||||
|
||||
def sort_config(config: BenchmarkConfig) -> BenchmarkConfig:
|
||||
return {
|
||||
"BLOCK_SIZE_M": config["BLOCK_SIZE_M"],
|
||||
"BLOCK_SIZE_N": config["BLOCK_SIZE_N"],
|
||||
"BLOCK_SIZE_K": config["BLOCK_SIZE_K"],
|
||||
"GROUP_SIZE_M": config["GROUP_SIZE_M"],
|
||||
"num_warps": config["num_warps"],
|
||||
"num_stages": config["num_stages"],
|
||||
**(
|
||||
{"waves_per_eu": config["waves_per_eu"]} if "waves_per_eu" in config else {}
|
||||
),
|
||||
**({"USE_TMA": config["USE_TMA"]} if "USE_TMA" in config else {}),
|
||||
}
|
||||
|
||||
|
||||
def save_configs(
|
||||
configs: Dict[int, BenchmarkConfig],
|
||||
filename: str,
|
||||
) -> None:
|
||||
print(f"Writing best config to {filename}...")
|
||||
with open(filename, "w") as f:
|
||||
json.dump(configs, f, indent=4)
|
||||
f.write("\n")
|
||||
|
||||
|
||||
def get_config_filename(
|
||||
num_experts: int,
|
||||
shard_intermediate_size: int,
|
||||
hidden_size: int,
|
||||
topk: int,
|
||||
dtype: torch.dtype,
|
||||
use_fp8_w8a8: bool,
|
||||
use_int8_w8a8: bool,
|
||||
use_int8_w8a16: bool,
|
||||
use_int4_w4a16: bool,
|
||||
per_channel_quant: bool,
|
||||
block_shape: List[int],
|
||||
) -> str:
|
||||
dtype_str = get_config_dtype_str(
|
||||
dtype,
|
||||
use_int8_w8a16=use_int8_w8a16,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
use_int8_w8a8=use_int8_w8a8,
|
||||
use_int4_w4a16=use_int4_w4a16,
|
||||
)
|
||||
|
||||
# NOTE(woosuk): The current naming convention uses w2.shape[2], which
|
||||
# is the intermediate size after silu_and_mul.
|
||||
N = shard_intermediate_size // 2
|
||||
if use_int4_w4a16:
|
||||
N = N // 2
|
||||
|
||||
filename = get_config_file_name(
|
||||
num_experts,
|
||||
N,
|
||||
dtype_str,
|
||||
block_shape,
|
||||
per_channel_quant,
|
||||
)
|
||||
|
||||
return filename
|
||||
|
||||
|
||||
def get_default_batch_sizes() -> List[int]:
|
||||
return [
|
||||
1,
|
||||
2,
|
||||
4,
|
||||
8,
|
||||
16,
|
||||
24,
|
||||
32,
|
||||
48,
|
||||
64,
|
||||
96,
|
||||
128,
|
||||
256,
|
||||
512,
|
||||
1024,
|
||||
1536,
|
||||
2048,
|
||||
3072,
|
||||
4096,
|
||||
]
|
||||
71
third_party/sglang/benchmark/kernels/fused_moe_triton/tuning_client.py
vendored
Normal file
71
third_party/sglang/benchmark/kernels/fused_moe_triton/tuning_client.py
vendored
Normal file
@@ -0,0 +1,71 @@
|
||||
import argparse
|
||||
import os
|
||||
import time
|
||||
|
||||
import openai
|
||||
|
||||
"""
|
||||
# Edit the code file srt/models/deepseek_v2.py in the Python site package and add the logic for saving topk_ids:
|
||||
# import get_tensor_model_parallel_rank
|
||||
# DeepseekV2MoE::forward_normal
|
||||
if hidden_states.shape[0] >= 4096 and get_tensor_model_parallel_rank() == 0:
|
||||
topk_ids_dir = xxxx
|
||||
if not hasattr(self, "save_idx"):
|
||||
self.save_idx = 0
|
||||
if self.save_idx <= 1:
|
||||
torch.save(topk_output.topk_ids, f"{topk_ids_dir}/topk_ids_layer{self.layer_id}_idx{self.save_idx}.pt")
|
||||
self.save_idx += 1
|
||||
"""
|
||||
|
||||
|
||||
def read_long_prompt():
|
||||
import json
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
with open(f"{current_dir}/tuning_text.json", "r") as fp:
|
||||
text = fp.read()
|
||||
rst = json.loads(text)
|
||||
return rst["prompt"]
|
||||
|
||||
|
||||
def openai_stream_test(model, ip, port):
|
||||
client = openai.Client(base_url=f"http://{ip}:{port}/v1", api_key="None")
|
||||
qst = read_long_prompt()
|
||||
|
||||
messages = [
|
||||
{"role": "user", "content": qst},
|
||||
]
|
||||
msg2 = dict(
|
||||
model=model,
|
||||
messages=messages,
|
||||
temperature=0.6,
|
||||
top_p=0.75,
|
||||
max_tokens=100,
|
||||
)
|
||||
response = client.chat.completions.create(**msg2, stream=True)
|
||||
time_start = time.time()
|
||||
time_cost = []
|
||||
for chunk in response:
|
||||
time_end = time.time()
|
||||
# if chunk.choices[0].delta.content:
|
||||
# print(chunk.choices[0].delta.content, end="", flush=True)
|
||||
time_cost.append(time_end - time_start)
|
||||
time_start = time.time()
|
||||
|
||||
ttft = time_cost[0] + time_cost[1]
|
||||
tpot = sum(time_cost[2:]) / len(time_cost[2:])
|
||||
print(f"\nTTFT {ttft}, TPOT {tpot}")
|
||||
return ttft, tpot
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--model", type=str, default="auto")
|
||||
parser.add_argument(
|
||||
"--ip",
|
||||
type=str,
|
||||
default="127.0.0.1",
|
||||
)
|
||||
parser.add_argument("--port", type=int, default=8188)
|
||||
args = parser.parse_args()
|
||||
openai_stream_test(args.model, args.ip, args.port)
|
||||
520
third_party/sglang/benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py
vendored
Normal file
520
third_party/sglang/benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py
vendored
Normal file
@@ -0,0 +1,520 @@
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/main/benchmarks/kernels/benchmark_moe.py
|
||||
import argparse
|
||||
import time
|
||||
from contextlib import nullcontext
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
import ray
|
||||
import torch
|
||||
import triton
|
||||
from common_utils import (
|
||||
BenchmarkConfig,
|
||||
get_config_filename,
|
||||
get_configs_compute_bound,
|
||||
get_default_batch_sizes,
|
||||
get_model_config,
|
||||
save_configs,
|
||||
sort_config,
|
||||
)
|
||||
from ray.experimental.tqdm_ray import tqdm
|
||||
|
||||
from sglang.srt.layers.moe.fused_moe_triton import override_config
|
||||
from sglang.srt.layers.moe.fused_moe_triton.fused_moe import fused_moe
|
||||
from sglang.srt.layers.moe.fused_moe_triton.fused_moe_triton_config import (
|
||||
get_config_dtype_str,
|
||||
get_default_config,
|
||||
get_moe_configs,
|
||||
)
|
||||
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
|
||||
from sglang.srt.layers.moe.topk import TopKConfig, select_experts
|
||||
from sglang.srt.server_args import (
|
||||
ServerArgs,
|
||||
set_global_server_args_for_scheduler,
|
||||
)
|
||||
from sglang.srt.utils import get_device, is_hip, is_xpu
|
||||
|
||||
_is_hip = is_hip()
|
||||
_is_xpu = is_xpu()
|
||||
|
||||
|
||||
def benchmark_config(
|
||||
config: BenchmarkConfig,
|
||||
num_tokens: int,
|
||||
num_experts: int,
|
||||
shard_intermediate_size: int,
|
||||
hidden_size: int,
|
||||
topk: int,
|
||||
dtype: torch.dtype,
|
||||
use_fp8_w8a8: bool,
|
||||
use_int8_w8a8: bool,
|
||||
use_int8_w8a16: bool,
|
||||
use_int4_w4a16: bool,
|
||||
per_channel_quant: bool,
|
||||
block_shape: List[int] = None,
|
||||
num_iters: int = 100,
|
||||
) -> float:
|
||||
init_dtype = torch.float16 if use_fp8_w8a8 else dtype
|
||||
x = torch.randn(num_tokens, hidden_size, dtype=dtype)
|
||||
if use_int8_w8a16 or use_int8_w8a8:
|
||||
w1 = torch.randint(
|
||||
-127,
|
||||
127,
|
||||
(
|
||||
num_experts,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
),
|
||||
dtype=torch.int8,
|
||||
)
|
||||
w2 = torch.randint(
|
||||
-127,
|
||||
127,
|
||||
(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
shard_intermediate_size // 2,
|
||||
),
|
||||
dtype=torch.int8,
|
||||
)
|
||||
elif use_int4_w4a16:
|
||||
w1 = torch.randint(
|
||||
0,
|
||||
255,
|
||||
(
|
||||
num_experts,
|
||||
shard_intermediate_size,
|
||||
hidden_size // 2,
|
||||
),
|
||||
dtype=torch.uint8,
|
||||
)
|
||||
w2 = torch.randint(
|
||||
0,
|
||||
255,
|
||||
(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
shard_intermediate_size // 4,
|
||||
),
|
||||
dtype=torch.uint8,
|
||||
)
|
||||
else:
|
||||
w1 = torch.randn(
|
||||
num_experts, shard_intermediate_size, hidden_size, dtype=init_dtype
|
||||
)
|
||||
w2 = torch.randn(
|
||||
num_experts, hidden_size, shard_intermediate_size // 2, dtype=init_dtype
|
||||
)
|
||||
gating_output = torch.randn(num_iters, num_tokens, num_experts, dtype=torch.float32)
|
||||
|
||||
w1_scale = None
|
||||
w2_scale = None
|
||||
a1_scale = None
|
||||
a2_scale = None
|
||||
if use_int8_w8a16:
|
||||
w1_scale = torch.randn(
|
||||
(num_experts, 2 * shard_intermediate_size), dtype=torch.float32
|
||||
)
|
||||
w2_scale = torch.randn((hidden_size, num_experts), dtype=torch.float32)
|
||||
if use_int4_w4a16:
|
||||
block_n = 1 if (block_shape[0] == 0) else block_shape[0]
|
||||
block_k = block_shape[1]
|
||||
n_tiles_w1 = (shard_intermediate_size + block_n - 1) // block_n
|
||||
n_tiles_w2 = (hidden_size + block_n - 1) // block_n
|
||||
k_tiles_w1 = (hidden_size + block_k - 1) // block_k
|
||||
k_tiles_w2 = (shard_intermediate_size // 2 + block_k - 1) // block_k
|
||||
w1_scale = torch.randn(
|
||||
(num_experts, n_tiles_w1, k_tiles_w1), dtype=torch.bfloat16
|
||||
)
|
||||
w2_scale = torch.randn(
|
||||
(num_experts, n_tiles_w2, k_tiles_w2), dtype=torch.bfloat16
|
||||
)
|
||||
if use_fp8_w8a8 or use_int8_w8a8:
|
||||
if use_int8_w8a8 and block_shape is None:
|
||||
w1_scale = torch.randn(
|
||||
num_experts, shard_intermediate_size, dtype=torch.float32
|
||||
)
|
||||
w2_scale = torch.randn(num_experts, hidden_size, dtype=torch.float32)
|
||||
elif block_shape is None:
|
||||
w1_scale = torch.randn(num_experts, dtype=torch.float32)
|
||||
w2_scale = torch.randn(num_experts, dtype=torch.float32)
|
||||
a1_scale = torch.randn(1, dtype=torch.float32)
|
||||
a2_scale = torch.randn(1, dtype=torch.float32)
|
||||
else:
|
||||
block_n, block_k = block_shape[0], block_shape[1]
|
||||
n_tiles_w1 = (shard_intermediate_size + block_n - 1) // block_n
|
||||
n_tiles_w2 = (hidden_size + block_n - 1) // block_n
|
||||
k_tiles_w1 = (hidden_size + block_k - 1) // block_k
|
||||
k_tiles_w2 = (shard_intermediate_size // 2 + block_k - 1) // block_k
|
||||
w1_scale = torch.rand(
|
||||
(num_experts, n_tiles_w1, k_tiles_w1), dtype=torch.float32
|
||||
)
|
||||
w2_scale = torch.rand(
|
||||
(num_experts, n_tiles_w2, k_tiles_w2), dtype=torch.float32
|
||||
)
|
||||
|
||||
if use_fp8_w8a8:
|
||||
w1 = w1.to(torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn)
|
||||
w2 = w2.to(torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn)
|
||||
|
||||
input_gating = torch.randn(num_tokens, num_experts, dtype=torch.float32)
|
||||
topk_config = TopKConfig(
|
||||
top_k=topk,
|
||||
renormalize=True,
|
||||
)
|
||||
topk_output = select_experts(x, input_gating, topk_config)
|
||||
|
||||
def prepare(i: int):
|
||||
input_gating = gating_output[i]
|
||||
new_topk_output = select_experts(x, input_gating, topk_config)
|
||||
topk_output.topk_weights.copy_(new_topk_output.topk_weights)
|
||||
topk_output.topk_ids.copy_(new_topk_output.topk_ids)
|
||||
topk_output.router_logits.copy_(new_topk_output.router_logits)
|
||||
|
||||
def run():
|
||||
moe_runner_config = MoeRunnerConfig(
|
||||
inplace=True,
|
||||
)
|
||||
|
||||
with override_config(config):
|
||||
fused_moe(
|
||||
x,
|
||||
w1,
|
||||
w2,
|
||||
topk_output,
|
||||
moe_runner_config=moe_runner_config,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
use_int8_w8a8=use_int8_w8a8,
|
||||
use_int8_w8a16=use_int8_w8a16,
|
||||
use_int4_w4a16=use_int4_w4a16,
|
||||
w1_scale=w1_scale,
|
||||
w2_scale=w2_scale,
|
||||
a1_scale=a1_scale,
|
||||
a2_scale=a2_scale,
|
||||
per_channel_quant=per_channel_quant,
|
||||
block_shape=block_shape,
|
||||
)
|
||||
|
||||
# JIT compilation & warmup
|
||||
run()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Capture 10 invocations with CUDA graph
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
for _ in range(10):
|
||||
run()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Warmup
|
||||
for _ in range(5):
|
||||
graph.replay()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Flush L2 cache with 256 MB data
|
||||
cache_flush = torch.empty(int(256e6 // 4), dtype=torch.int, device="cuda")
|
||||
cache_flush.zero_()
|
||||
|
||||
start_events = [torch.cuda.Event(enable_timing=True) for _ in range(num_iters)]
|
||||
end_events = [torch.cuda.Event(enable_timing=True) for _ in range(num_iters)]
|
||||
|
||||
for i in range(num_iters):
|
||||
prepare(i)
|
||||
start_events[i].record()
|
||||
graph.replay()
|
||||
end_events[i].record()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
latencies: List[float] = []
|
||||
for i in range(num_iters):
|
||||
latencies.append(start_events[i].elapsed_time(end_events[i]))
|
||||
avg = sum(latencies) / (num_iters * 10) * 1000 # us
|
||||
graph.reset()
|
||||
return avg
|
||||
|
||||
|
||||
@ray.remote(num_gpus=1)
|
||||
class BenchmarkWorker:
|
||||
|
||||
def __init__(self, seed: int, server_args: ServerArgs) -> None:
|
||||
torch.set_default_device(get_device())
|
||||
torch.get_device_module().manual_seed_all(0)
|
||||
self.seed = seed
|
||||
# Get the device ID to allocate tensors and kernels
|
||||
# on the respective GPU.
|
||||
self.device_id = int(ray.get_gpu_ids()[0])
|
||||
set_global_server_args_for_scheduler(server_args)
|
||||
|
||||
def benchmark(
|
||||
self,
|
||||
num_tokens: int,
|
||||
num_experts: int,
|
||||
shard_intermediate_size: int,
|
||||
hidden_size: int,
|
||||
topk: int,
|
||||
dtype: torch.dtype,
|
||||
use_fp8_w8a8: bool,
|
||||
use_int8_w8a8: bool,
|
||||
use_int8_w8a16: bool,
|
||||
use_int4_w4a16: bool,
|
||||
per_channel_quant: bool,
|
||||
block_shape: List[int],
|
||||
) -> Tuple[Dict[str, int], float]:
|
||||
torch.cuda.manual_seed_all(0)
|
||||
dtype_str = get_config_dtype_str(
|
||||
dtype,
|
||||
use_int8_w8a16=use_int8_w8a16,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
use_int4_w4a16=use_int4_w4a16,
|
||||
)
|
||||
# NOTE(woosuk): The current naming convention uses w2.shape[2], which
|
||||
# is the intermediate size after silu_and_mul.
|
||||
block_n = block_shape[0] if block_shape else 0
|
||||
block_k = block_shape[1] if block_shape else 0
|
||||
N = shard_intermediate_size // 2
|
||||
if use_int4_w4a16:
|
||||
N = N // 2
|
||||
op_config = get_moe_configs(
|
||||
num_experts,
|
||||
N,
|
||||
dtype_str,
|
||||
block_n,
|
||||
block_k,
|
||||
per_channel_quant,
|
||||
)
|
||||
if op_config is None:
|
||||
config = get_default_config(
|
||||
num_tokens,
|
||||
num_experts,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype_str,
|
||||
False,
|
||||
block_shape,
|
||||
)
|
||||
else:
|
||||
config = op_config[min(op_config.keys(), key=lambda x: abs(x - num_tokens))]
|
||||
with torch.cuda.device(self.device_id) if is_hip() else nullcontext():
|
||||
kernel_time = benchmark_config(
|
||||
config,
|
||||
num_tokens,
|
||||
num_experts,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
per_channel_quant,
|
||||
block_shape,
|
||||
)
|
||||
return config, kernel_time
|
||||
|
||||
def tune(
|
||||
self,
|
||||
num_tokens: int,
|
||||
num_experts: int,
|
||||
shard_intermediate_size: int,
|
||||
hidden_size: int,
|
||||
topk: int,
|
||||
dtype: torch.dtype,
|
||||
use_fp8_w8a8: bool,
|
||||
use_int8_w8a8: bool,
|
||||
use_int8_w8a16: bool,
|
||||
use_int4_w4a16: bool,
|
||||
per_channel_quant: bool,
|
||||
block_shape: List[int],
|
||||
search_space: List[Dict[str, int]],
|
||||
) -> Dict[str, int]:
|
||||
best_config = None
|
||||
best_time = float("inf")
|
||||
with (
|
||||
torch.get_device_module().device(self.device_id)
|
||||
if _is_xpu or _is_hip
|
||||
else nullcontext()
|
||||
):
|
||||
for config in tqdm(search_space):
|
||||
try:
|
||||
kernel_time = benchmark_config(
|
||||
config,
|
||||
num_tokens,
|
||||
num_experts,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
per_channel_quant,
|
||||
block_shape,
|
||||
num_iters=10,
|
||||
)
|
||||
except (triton.runtime.autotuner.OutOfResources, RuntimeError):
|
||||
# Some configurations may be invalid and fail to compile.
|
||||
continue
|
||||
|
||||
if kernel_time < best_time:
|
||||
best_time = kernel_time
|
||||
best_config = config
|
||||
now = datetime.now()
|
||||
print(f"{now.ctime()}] Completed tuning for batch_size={num_tokens}")
|
||||
assert best_config is not None
|
||||
return best_config
|
||||
|
||||
|
||||
def main(args: argparse.Namespace):
|
||||
server_args = ServerArgs(
|
||||
model_path=args.model, tp_size=args.tp_size, ep_size=args.ep_size
|
||||
)
|
||||
|
||||
model_config = get_model_config(
|
||||
args.model, args.tp_size, args.ep_size, args.disable_shared_experts_fusion
|
||||
)
|
||||
|
||||
E = model_config["num_experts"]
|
||||
topk = model_config["topk"]
|
||||
hidden_size = model_config["hidden_size"]
|
||||
shard_intermediate_size = model_config["shard_intermediate_size"]
|
||||
dtype = model_config["dtype"]
|
||||
block_shape = model_config["block_shape"]
|
||||
|
||||
use_fp8_w8a8 = args.dtype == "fp8_w8a8"
|
||||
use_int8_w8a8 = args.dtype == "int8_w8a8"
|
||||
use_int8_w8a16 = args.dtype == "int8_w8a16"
|
||||
use_int4_w4a16 = args.dtype == "int4_w4a16"
|
||||
per_channel_quant = args.per_channel_quant
|
||||
|
||||
if args.batch_size is None:
|
||||
batch_sizes = get_default_batch_sizes()
|
||||
else:
|
||||
batch_sizes = [args.batch_size]
|
||||
|
||||
ray.init()
|
||||
num_gpus = int(ray.available_resources()["GPU"])
|
||||
workers = [BenchmarkWorker.remote(args.seed, server_args) for _ in range(num_gpus)]
|
||||
|
||||
def _distribute(method: str, inputs: List[Any]) -> List[Any]:
|
||||
outputs = []
|
||||
worker_idx = 0
|
||||
for input_args in inputs:
|
||||
worker = workers[worker_idx]
|
||||
worker_method = getattr(worker, method)
|
||||
output = worker_method.remote(*input_args)
|
||||
outputs.append(output)
|
||||
worker_idx = (worker_idx + 1) % num_gpus
|
||||
return ray.get(outputs)
|
||||
|
||||
if args.tune:
|
||||
search_space = get_configs_compute_bound()
|
||||
if block_shape is not None:
|
||||
block_n, block_k = block_shape[0], block_shape[1]
|
||||
search_space = [
|
||||
config
|
||||
for config in search_space
|
||||
if block_k % config["BLOCK_SIZE_K"] == 0
|
||||
]
|
||||
|
||||
filename = get_config_filename(
|
||||
E,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
per_channel_quant,
|
||||
block_shape,
|
||||
)
|
||||
print(
|
||||
f"Start tuning over {len(search_space)} configurations to create {filename}..."
|
||||
)
|
||||
|
||||
start = time.perf_counter()
|
||||
configs = _distribute(
|
||||
"tune",
|
||||
[
|
||||
(
|
||||
batch_size,
|
||||
E,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
per_channel_quant,
|
||||
block_shape,
|
||||
search_space,
|
||||
)
|
||||
for batch_size in batch_sizes
|
||||
],
|
||||
)
|
||||
best_configs = {
|
||||
M: sort_config(config) for M, config in zip(batch_sizes, configs)
|
||||
}
|
||||
save_configs(
|
||||
best_configs,
|
||||
filename,
|
||||
)
|
||||
end = time.perf_counter()
|
||||
print(f"Tuning took {end - start:.2f} seconds")
|
||||
else:
|
||||
outputs = _distribute(
|
||||
"benchmark",
|
||||
[
|
||||
(
|
||||
batch_size,
|
||||
E,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
per_channel_quant,
|
||||
block_shape,
|
||||
)
|
||||
for batch_size in batch_sizes
|
||||
],
|
||||
)
|
||||
|
||||
for batch_size, (config, kernel_time) in zip(batch_sizes, outputs):
|
||||
print(f"Batch size: {batch_size}, config: {config}")
|
||||
print(f"Kernel time: {kernel_time:.2f} us")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--model", type=str, default="mistralai/Mixtral-8x7B-Instruct-v0.1"
|
||||
)
|
||||
parser.add_argument("--tp-size", "--tp", type=int, default=2)
|
||||
parser.add_argument("--ep-size", "--ep", type=int, default=1)
|
||||
parser.add_argument(
|
||||
"--dtype",
|
||||
type=str,
|
||||
choices=["auto", "fp8_w8a8", "int8_w8a16", "int8_w8a8", "int4_w4a16"],
|
||||
default="auto",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--per-channel-quant",
|
||||
action="store_true",
|
||||
)
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
parser.add_argument("--batch-size", type=int, required=False)
|
||||
parser.add_argument("--tune", action="store_true")
|
||||
parser.add_argument("--disable-shared-experts-fusion", action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
893
third_party/sglang/benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton_sep.py
vendored
Normal file
893
third_party/sglang/benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton_sep.py
vendored
Normal file
@@ -0,0 +1,893 @@
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/main/benchmarks/kernels/benchmark_moe.py
|
||||
import argparse
|
||||
import dataclasses
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from contextlib import nullcontext
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
import ray
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
from common_utils import (
|
||||
BenchmarkConfig,
|
||||
get_config_filename,
|
||||
get_configs_compute_bound,
|
||||
get_default_batch_sizes,
|
||||
get_model_config,
|
||||
sort_config,
|
||||
)
|
||||
from ray.experimental.tqdm_ray import tqdm
|
||||
|
||||
from sglang.srt.layers.moe.fused_moe_triton.fused_moe import (
|
||||
get_config_dtype_str,
|
||||
invoke_fused_moe_kernel,
|
||||
moe_align_block_size,
|
||||
)
|
||||
from sglang.srt.layers.moe.fused_moe_triton.fused_moe_triton_config import (
|
||||
get_config_file_name,
|
||||
)
|
||||
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
|
||||
from sglang.srt.layers.moe.topk import TopKConfig, select_experts
|
||||
from sglang.srt.server_args import (
|
||||
ServerArgs,
|
||||
set_global_server_args_for_scheduler,
|
||||
)
|
||||
from sglang.srt.utils import is_hip
|
||||
|
||||
_is_hip = is_hip()
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class MoeInputs:
|
||||
topk_ids: torch.Tensor
|
||||
sorted_token_ids: torch.Tensor
|
||||
expert_ids: torch.Tensor
|
||||
num_tokens_post_padded: torch.Tensor
|
||||
|
||||
|
||||
class KernelWrapper:
|
||||
def __init__(self, moe_inputs, use_cuda_graph=True, inner_iter=10, **kwargs):
|
||||
self.func = invoke_fused_moe_kernel
|
||||
self.use_cuda_graph = use_cuda_graph
|
||||
self.moe_inputs = moe_inputs
|
||||
self.inner_iter = inner_iter
|
||||
self.kwargs = kwargs
|
||||
if use_cuda_graph:
|
||||
self.graph = self.cuda_graph_wrapper()
|
||||
else:
|
||||
self.graph = None
|
||||
|
||||
def cuda_graph_wrapper(self):
|
||||
moe_input = self.moe_inputs[0]
|
||||
self.func(
|
||||
**self.kwargs,
|
||||
topk_ids=moe_input.topk_ids,
|
||||
sorted_token_ids=moe_input.sorted_token_ids,
|
||||
expert_ids=moe_input.expert_ids,
|
||||
num_tokens_post_padded=moe_input.num_tokens_post_padded,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Capture 10 invocations with CUDA graph
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
for k in range(self.inner_iter):
|
||||
moe_input = self.moe_inputs[k]
|
||||
self.func(
|
||||
**self.kwargs,
|
||||
topk_ids=moe_input.topk_ids,
|
||||
sorted_token_ids=moe_input.sorted_token_ids,
|
||||
expert_ids=moe_input.expert_ids,
|
||||
num_tokens_post_padded=moe_input.num_tokens_post_padded,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Warmup
|
||||
for _ in range(5):
|
||||
graph.replay()
|
||||
torch.cuda.synchronize()
|
||||
return graph
|
||||
|
||||
def forward_cost(self, try_cnt=2):
|
||||
time_cost = float("inf")
|
||||
for _ in range(try_cnt):
|
||||
start_event = torch.cuda.Event(enable_timing=True)
|
||||
end_event = torch.cuda.Event(enable_timing=True)
|
||||
start_event.record()
|
||||
if self.use_cuda_graph:
|
||||
self.graph.replay()
|
||||
else:
|
||||
for k in range(self.inner_iter):
|
||||
moe_input = self.moe_inputs[k]
|
||||
self.func(
|
||||
**self.kwargs,
|
||||
topk_ids=moe_input.topk_ids,
|
||||
sorted_token_ids=moe_input.sorted_token_ids,
|
||||
expert_ids=moe_input.expert_ids,
|
||||
num_tokens_post_padded=moe_input.num_tokens_post_padded,
|
||||
)
|
||||
end_event.record()
|
||||
torch.cuda.synchronize()
|
||||
time_cost = min(time_cost, start_event.elapsed_time(end_event))
|
||||
return time_cost
|
||||
|
||||
|
||||
def load_topk_ids(topk_ids_dir, i: int):
|
||||
num_layers = 61
|
||||
dense_layers = 3
|
||||
moe_layers = num_layers - dense_layers
|
||||
return torch.load(
|
||||
f"{topk_ids_dir}/topk_ids_layer{i % moe_layers + dense_layers}_idx{i // moe_layers}.pt"
|
||||
)
|
||||
|
||||
|
||||
def benchmark_config(
|
||||
config: BenchmarkConfig,
|
||||
num_tokens: int,
|
||||
num_experts: int,
|
||||
shard_intermediate_size: int,
|
||||
hidden_size: int,
|
||||
topk: int,
|
||||
dtype: torch.dtype,
|
||||
use_fp8_w8a8: bool,
|
||||
use_int8_w8a8: bool,
|
||||
use_int8_w8a16: bool,
|
||||
use_int4_w4a16: bool,
|
||||
topk_ids_list,
|
||||
block_shape: List[int] = None,
|
||||
ep_size: int = 1,
|
||||
num_iters: int = 100,
|
||||
) -> float:
|
||||
ncu_enable = os.getenv("NCU_ENABLE", "0") == "1"
|
||||
if ncu_enable:
|
||||
num_iters = 1
|
||||
init_dtype = torch.float16 if use_fp8_w8a8 else dtype
|
||||
hidden_states = torch.randn(num_tokens, hidden_size, dtype=dtype)
|
||||
if use_int8_w8a16 or use_int8_w8a8:
|
||||
w1 = torch.randint(
|
||||
-127,
|
||||
127,
|
||||
(
|
||||
num_experts,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
),
|
||||
dtype=torch.int8,
|
||||
)
|
||||
w2 = torch.randint(
|
||||
-127,
|
||||
127,
|
||||
(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
shard_intermediate_size // 2,
|
||||
),
|
||||
dtype=torch.int8,
|
||||
)
|
||||
elif use_int4_w4a16:
|
||||
w1 = torch.randint(
|
||||
0,
|
||||
255,
|
||||
(
|
||||
num_experts,
|
||||
shard_intermediate_size,
|
||||
hidden_size // 2,
|
||||
),
|
||||
dtype=torch.uint8,
|
||||
)
|
||||
w2 = torch.randint(
|
||||
0,
|
||||
255,
|
||||
(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
shard_intermediate_size // 4,
|
||||
),
|
||||
dtype=torch.uint8,
|
||||
)
|
||||
else:
|
||||
w1 = torch.randn(
|
||||
num_experts, shard_intermediate_size, hidden_size, dtype=init_dtype
|
||||
)
|
||||
w2 = torch.randn(
|
||||
num_experts, hidden_size, shard_intermediate_size // 2, dtype=init_dtype
|
||||
)
|
||||
|
||||
w1_scale = None
|
||||
w2_scale = None
|
||||
a1_scale = None
|
||||
a2_scale = None
|
||||
if use_int8_w8a16:
|
||||
w1_scale = torch.randn(
|
||||
(num_experts, 2 * shard_intermediate_size), dtype=torch.float32
|
||||
)
|
||||
w2_scale = torch.randn((hidden_size, num_experts), dtype=torch.float32)
|
||||
if use_int4_w4a16:
|
||||
block_n = 1 if (block_shape[0] == 0) else block_shape[0]
|
||||
block_k = block_shape[1]
|
||||
n_tiles_w1 = (shard_intermediate_size + block_n - 1) // block_n
|
||||
n_tiles_w2 = (hidden_size + block_n - 1) // block_n
|
||||
k_tiles_w1 = (hidden_size + block_k - 1) // block_k
|
||||
k_tiles_w2 = (shard_intermediate_size // 2 + block_k - 1) // block_k
|
||||
w1_scale = torch.randn(
|
||||
(num_experts, n_tiles_w1, k_tiles_w1), dtype=torch.bfloat16
|
||||
)
|
||||
w2_scale = torch.randn(
|
||||
(num_experts, n_tiles_w2, k_tiles_w2), dtype=torch.bfloat16
|
||||
)
|
||||
if use_fp8_w8a8 or use_int8_w8a8:
|
||||
if use_int8_w8a8 and block_shape is None:
|
||||
w1_scale = torch.randn(
|
||||
num_experts, shard_intermediate_size, dtype=torch.float32
|
||||
)
|
||||
w2_scale = torch.randn(num_experts, hidden_size, dtype=torch.float32)
|
||||
elif block_shape is None:
|
||||
w1_scale = torch.randn(num_experts, dtype=torch.float32)
|
||||
w2_scale = torch.randn(num_experts, dtype=torch.float32)
|
||||
a1_scale = torch.randn(1, dtype=torch.float32)
|
||||
a2_scale = torch.randn(1, dtype=torch.float32)
|
||||
else:
|
||||
block_n, block_k = block_shape[0], block_shape[1]
|
||||
n_tiles_w1 = (shard_intermediate_size + block_n - 1) // block_n
|
||||
n_tiles_w2 = (hidden_size + block_n - 1) // block_n
|
||||
k_tiles_w1 = (hidden_size + block_k - 1) // block_k
|
||||
k_tiles_w2 = (shard_intermediate_size // 2 + block_k - 1) // block_k
|
||||
w1_scale = torch.rand(
|
||||
(num_experts, n_tiles_w1, k_tiles_w1), dtype=torch.float32
|
||||
)
|
||||
w2_scale = torch.rand(
|
||||
(num_experts, n_tiles_w2, k_tiles_w2), dtype=torch.float32
|
||||
)
|
||||
|
||||
if use_fp8_w8a8:
|
||||
w1 = w1.to(torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn)
|
||||
w2 = w2.to(torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn)
|
||||
|
||||
input_gating = torch.randn(num_tokens, num_experts, dtype=torch.float32)
|
||||
topk_config = TopKConfig(
|
||||
top_k=topk,
|
||||
renormalize=True,
|
||||
)
|
||||
topk_output_ = select_experts(hidden_states, input_gating, topk_config)
|
||||
sorted_token_ids_, expert_ids_, num_tokens_post_padded_ = moe_align_block_size(
|
||||
topk_output_.topk_ids, config["BLOCK_SIZE_M"], num_experts
|
||||
)
|
||||
inner_iter = 10 if not ncu_enable else 1
|
||||
moe_inputs = [
|
||||
MoeInputs(
|
||||
topk_output_.topk_ids.clone(),
|
||||
sorted_token_ids_.clone(),
|
||||
expert_ids_.clone(),
|
||||
num_tokens_post_padded_.clone(),
|
||||
)
|
||||
for _ in range(inner_iter)
|
||||
]
|
||||
M = hidden_states.shape[0]
|
||||
E, N, _ = w1.shape
|
||||
|
||||
padded_tokens = min(M * topk, E + 1) * (
|
||||
config["BLOCK_SIZE_M"] - 1
|
||||
) # if moe_use_tma else 0
|
||||
total_tokens = M * topk + padded_tokens
|
||||
cache = torch.empty(
|
||||
total_tokens * max(N, w2.shape[1]),
|
||||
device=hidden_states.device,
|
||||
dtype=hidden_states.dtype,
|
||||
)
|
||||
intermediate_cache1 = cache[: total_tokens * N].view(
|
||||
(total_tokens, N),
|
||||
)
|
||||
intermediate_cache2 = torch.empty(
|
||||
(total_tokens, N // 2),
|
||||
device=hidden_states.device,
|
||||
dtype=hidden_states.dtype,
|
||||
)
|
||||
intermediate_cache3 = cache[: M * topk * w2.shape[1]].view(
|
||||
(M, topk, w2.shape[1]),
|
||||
)
|
||||
|
||||
def prepare(i: int, inner_iter): # update inputs according to topk_ids
|
||||
for k in range(inner_iter):
|
||||
topk_ids = topk_ids_list[i * inner_iter + k]
|
||||
# With EP, saved topk_ids are global expert indices; remap to local.
|
||||
if ep_size > 1:
|
||||
topk_ids = (topk_ids // ep_size).to(
|
||||
device=moe_inputs[k].topk_ids.device,
|
||||
dtype=moe_inputs[k].topk_ids.dtype,
|
||||
)
|
||||
tokens, _topk = moe_inputs[k].topk_ids.shape
|
||||
moe_inputs[k].topk_ids.copy_(topk_ids[:tokens, :_topk])
|
||||
sorted_token_ids_, expert_ids_, num_tokens_post_padded_ = (
|
||||
moe_align_block_size(
|
||||
moe_inputs[k].topk_ids, config["BLOCK_SIZE_M"], num_experts
|
||||
)
|
||||
)
|
||||
moe_inputs[k].sorted_token_ids.copy_(sorted_token_ids_)
|
||||
moe_inputs[k].expert_ids.copy_(expert_ids_)
|
||||
moe_inputs[k].num_tokens_post_padded.copy_(num_tokens_post_padded_)
|
||||
|
||||
def get_kernel_wrapper(moe_use_tma, inner_iter, use_cuda_graph):
|
||||
compute_type = (
|
||||
tl.bfloat16 if hidden_states.dtype == torch.bfloat16 else tl.float16
|
||||
)
|
||||
moe_runner_config = MoeRunnerConfig(
|
||||
inplace=True,
|
||||
)
|
||||
apply_router_weight_on_input = moe_runner_config.apply_router_weight_on_input
|
||||
kernel0 = KernelWrapper(
|
||||
A=hidden_states,
|
||||
B=w1,
|
||||
bias=None,
|
||||
C=intermediate_cache1,
|
||||
A_scale=a1_scale,
|
||||
B_scale=w1_scale,
|
||||
B_zp=None,
|
||||
topk_weights=topk_output_.topk_weights,
|
||||
moe_inputs=moe_inputs,
|
||||
mul_routed_weight=apply_router_weight_on_input,
|
||||
top_k=topk,
|
||||
config=config,
|
||||
compute_type=compute_type,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
use_int8_w8a8=use_int8_w8a8,
|
||||
use_int8_w8a16=use_int8_w8a16,
|
||||
use_int4_w4a16=use_int4_w4a16,
|
||||
per_channel_quant=False,
|
||||
block_shape=block_shape,
|
||||
b_use_tma=moe_use_tma,
|
||||
c_sorted=moe_use_tma,
|
||||
filter_expert=False,
|
||||
use_cuda_graph=use_cuda_graph,
|
||||
inner_iter=inner_iter,
|
||||
)
|
||||
kernel1 = KernelWrapper(
|
||||
A=intermediate_cache2,
|
||||
B=w2,
|
||||
bias=None,
|
||||
C=intermediate_cache3,
|
||||
A_scale=a2_scale,
|
||||
B_scale=w2_scale,
|
||||
B_zp=None,
|
||||
topk_weights=topk_output_.topk_weights,
|
||||
moe_inputs=moe_inputs,
|
||||
mul_routed_weight=not apply_router_weight_on_input,
|
||||
top_k=1,
|
||||
config=config,
|
||||
compute_type=compute_type,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
use_int8_w8a8=use_int8_w8a8,
|
||||
use_int8_w8a16=use_int8_w8a16,
|
||||
use_int4_w4a16=use_int4_w4a16,
|
||||
per_channel_quant=False,
|
||||
block_shape=block_shape,
|
||||
a_use_tma=moe_use_tma,
|
||||
b_use_tma=moe_use_tma,
|
||||
filter_expert=False,
|
||||
use_cuda_graph=use_cuda_graph,
|
||||
inner_iter=inner_iter,
|
||||
)
|
||||
return kernel0, kernel1
|
||||
|
||||
use_cuda_graph = True if not ncu_enable else False
|
||||
|
||||
kernel0, kernel1 = get_kernel_wrapper(False, inner_iter, use_cuda_graph)
|
||||
kernel_tma0, kernel_tma1 = get_kernel_wrapper(True, inner_iter, use_cuda_graph)
|
||||
|
||||
# JIT compilation & warmup
|
||||
if not ncu_enable:
|
||||
kernel0.forward_cost()
|
||||
kernel1.forward_cost()
|
||||
kernel_tma0.forward_cost()
|
||||
kernel_tma1.forward_cost()
|
||||
|
||||
ts0 = []
|
||||
ts1 = []
|
||||
ts_tma0 = []
|
||||
ts_tma1 = []
|
||||
|
||||
for i in range(num_iters // inner_iter):
|
||||
prepare(i, inner_iter)
|
||||
ts0.append(kernel0.forward_cost())
|
||||
ts1.append(kernel1.forward_cost())
|
||||
ts_tma0.append(kernel_tma0.forward_cost())
|
||||
ts_tma1.append(kernel_tma1.forward_cost())
|
||||
torch.cuda.synchronize()
|
||||
|
||||
avg = sum(ts0) / (num_iters) * 1000 # us
|
||||
avg1 = sum(ts1) / (num_iters) * 1000 # us
|
||||
avg_tma = sum(ts_tma0) / (num_iters) * 1000 # us
|
||||
avg1_tma = sum(ts_tma1) / (num_iters) * 1000 # us
|
||||
|
||||
return avg, avg_tma, avg1, avg1_tma
|
||||
|
||||
|
||||
class BestConfigTrace:
|
||||
def __init__(self, name, down_moe=False):
|
||||
self.name = name
|
||||
self.down_moe = down_moe
|
||||
self.best_costs_m = {} # block_m: best_cost
|
||||
|
||||
def update(self, config, time_cost_all):
|
||||
block_m = config["BLOCK_SIZE_M"]
|
||||
if not self.down_moe:
|
||||
time_cost = time_cost_all[0]
|
||||
else:
|
||||
time_cost = min(time_cost_all[2], time_cost_all[3])
|
||||
if (
|
||||
block_m not in self.best_costs_m
|
||||
or time_cost < self.best_costs_m[block_m][1]
|
||||
):
|
||||
self.best_costs_m[block_m] = config, time_cost, time_cost_all
|
||||
|
||||
def time_cost(self, block_m):
|
||||
if block_m not in self.best_costs_m:
|
||||
return float("inf")
|
||||
time_cost = self.best_costs_m[block_m][1]
|
||||
return time_cost
|
||||
|
||||
def config_dict(self, block_m):
|
||||
if block_m not in self.best_costs_m:
|
||||
return {}
|
||||
config, _, time_cost_all = self.best_costs_m[block_m]
|
||||
if not self.down_moe:
|
||||
return config
|
||||
else:
|
||||
return {
|
||||
**config,
|
||||
"USE_TMA": time_cost_all[2] > time_cost_all[3],
|
||||
}
|
||||
|
||||
|
||||
class BenchmarkWorker:
|
||||
|
||||
def __init__(self, seed: int, server_args: ServerArgs) -> None:
|
||||
torch.set_default_device("cuda")
|
||||
torch.cuda.manual_seed_all(0)
|
||||
self.seed = seed
|
||||
# Get the device ID to allocate tensors and kernels
|
||||
# on the respective GPU.
|
||||
self.device_id = 0 # int(ray.get_gpu_ids()[0])
|
||||
set_global_server_args_for_scheduler(server_args)
|
||||
|
||||
def benchmark(
|
||||
self,
|
||||
num_tokens: int,
|
||||
num_experts: int,
|
||||
shard_intermediate_size: int,
|
||||
hidden_size: int,
|
||||
topk: int,
|
||||
dtype: torch.dtype,
|
||||
use_fp8_w8a8: bool,
|
||||
use_int8_w8a8: bool,
|
||||
use_int8_w8a16: bool,
|
||||
use_int4_w4a16: bool,
|
||||
block_shape: List[int],
|
||||
cfg: Dict[str, int],
|
||||
topk_ids_dir: str,
|
||||
ep_size: int = 1,
|
||||
) -> Tuple[Dict[str, int], float]:
|
||||
torch.cuda.manual_seed_all(0)
|
||||
topk_ids_list = [load_topk_ids(topk_ids_dir, i) for i in range(100)]
|
||||
with torch.cuda.device(self.device_id) if is_hip() else nullcontext():
|
||||
kernel_time = benchmark_config(
|
||||
cfg,
|
||||
num_tokens,
|
||||
num_experts,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
topk_ids_list,
|
||||
block_shape,
|
||||
ep_size=ep_size,
|
||||
)
|
||||
return cfg, kernel_time
|
||||
|
||||
def tune(
|
||||
self,
|
||||
num_tokens: int,
|
||||
num_experts: int,
|
||||
shard_intermediate_size: int,
|
||||
hidden_size: int,
|
||||
topk: int,
|
||||
dtype: torch.dtype,
|
||||
use_fp8_w8a8: bool,
|
||||
use_int8_w8a8: bool,
|
||||
use_int8_w8a16: bool,
|
||||
use_int4_w4a16: bool,
|
||||
block_shape: List[int],
|
||||
search_space: List[Dict[str, int]],
|
||||
topk_ids_dir: str,
|
||||
ep_size: int = 1,
|
||||
) -> Dict[str, int]:
|
||||
trace0 = BestConfigTrace("kernel0", down_moe=False)
|
||||
trace1 = BestConfigTrace("kernel1", down_moe=True)
|
||||
topk_ids_list = [load_topk_ids(topk_ids_dir, i) for i in range(100)]
|
||||
|
||||
with torch.cuda.device(self.device_id) if is_hip() else nullcontext():
|
||||
for config in tqdm(search_space):
|
||||
try:
|
||||
kt0_no_tma, kt0_tma, kt1_no_tma, kt1_tma = benchmark_config(
|
||||
config,
|
||||
num_tokens,
|
||||
num_experts,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
topk_ids_list,
|
||||
block_shape,
|
||||
ep_size=ep_size,
|
||||
num_iters=100,
|
||||
)
|
||||
except triton.runtime.autotuner.OutOfResources:
|
||||
# Some configurations may be invalid and fail to compile.
|
||||
continue
|
||||
trace0.update(
|
||||
config,
|
||||
(kt0_no_tma, kt0_tma, kt1_no_tma, kt1_tma),
|
||||
)
|
||||
trace1.update(
|
||||
config,
|
||||
(kt0_no_tma, kt0_tma, kt1_no_tma, kt1_tma),
|
||||
)
|
||||
|
||||
now = datetime.now()
|
||||
print(f"{now.ctime()}] Completed tuning for batch_size={num_tokens}")
|
||||
best_block_m = 16
|
||||
for block_m in (32, 64, 128, 256):
|
||||
if trace0.time_cost(block_m) + trace1.time_cost(block_m) < trace0.time_cost(
|
||||
best_block_m
|
||||
) + trace1.time_cost(best_block_m):
|
||||
best_block_m = block_m
|
||||
|
||||
return (
|
||||
trace0.config_dict(best_block_m),
|
||||
trace1.config_dict(best_block_m),
|
||||
trace0.time_cost(best_block_m),
|
||||
trace1.time_cost(best_block_m),
|
||||
)
|
||||
|
||||
def cmp_configs(
|
||||
self,
|
||||
num_tokens: List[int],
|
||||
num_experts: int,
|
||||
shard_intermediate_size: int,
|
||||
hidden_size: int,
|
||||
topk: int,
|
||||
dtype: torch.dtype,
|
||||
use_fp8_w8a8: bool,
|
||||
use_int8_w8a8: bool,
|
||||
use_int8_w8a16: bool,
|
||||
use_int4_w4a16: bool,
|
||||
block_shape: List[int],
|
||||
cmp_config_files: List[str],
|
||||
topk_ids_dir: str,
|
||||
ep_size: int = 1,
|
||||
):
|
||||
# compare performance of different configs
|
||||
cmp_configs = []
|
||||
for file in cmp_config_files:
|
||||
with open(file) as f:
|
||||
cmp_configs.append({int(key): val for key, val in json.load(f).items()})
|
||||
for i, file in enumerate(cmp_config_files):
|
||||
print(f"config {i}: {file}")
|
||||
|
||||
topk_ids_list = [load_topk_ids(topk_ids_dir, i) for i in range(100)]
|
||||
torch.cuda.manual_seed_all(0)
|
||||
with torch.cuda.device(self.device_id) if is_hip() else nullcontext():
|
||||
for bs in num_tokens:
|
||||
kernel_times = []
|
||||
cfgs = []
|
||||
for configs in cmp_configs:
|
||||
cfg_org = configs[min(configs.keys(), key=lambda x: abs(x - bs))]
|
||||
cfgs.append(cfg_org)
|
||||
cfg = cfg_org.copy()
|
||||
cfg.pop("USE_TMA", None)
|
||||
kernel_time = benchmark_config(
|
||||
cfg,
|
||||
bs,
|
||||
num_experts,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
topk_ids_list,
|
||||
block_shape,
|
||||
ep_size=ep_size,
|
||||
)
|
||||
kernel_times.append(kernel_time)
|
||||
print(f"batch_size={bs=}:")
|
||||
for i, cfg in enumerate(cfgs):
|
||||
print(f" config {i} {cfg}: {kernel_times[i]}")
|
||||
|
||||
|
||||
def save_configs_sep(
|
||||
configs: Dict[int, BenchmarkConfig],
|
||||
num_experts: int,
|
||||
shard_intermediate_size: int,
|
||||
hidden_size: int,
|
||||
topk: int,
|
||||
dtype: torch.dtype,
|
||||
use_fp8_w8a8: bool,
|
||||
use_int8_w8a8: bool,
|
||||
use_int8_w8a16: bool,
|
||||
use_int4_w4a16: bool,
|
||||
block_shape: List[int],
|
||||
down_moe: bool = False,
|
||||
) -> None:
|
||||
dtype_str = get_config_dtype_str(
|
||||
dtype,
|
||||
use_int8_w8a16=use_int8_w8a16,
|
||||
use_fp8_w8a8=use_fp8_w8a8,
|
||||
use_int8_w8a8=use_int8_w8a8,
|
||||
use_int4_w4a16=use_int4_w4a16,
|
||||
)
|
||||
|
||||
# NOTE(woosuk): The current naming convention uses w2.shape[2], which
|
||||
# is the intermediate size after silu_and_mul.
|
||||
filename = get_config_file_name(
|
||||
num_experts,
|
||||
shard_intermediate_size // 2,
|
||||
dtype_str,
|
||||
block_shape,
|
||||
down_moe=down_moe,
|
||||
)
|
||||
|
||||
print(f"Writing best config to {filename}...")
|
||||
with open(filename, "w") as f:
|
||||
json.dump(configs, f, indent=4)
|
||||
f.write("\n")
|
||||
|
||||
|
||||
def main(args: argparse.Namespace):
|
||||
print(args)
|
||||
|
||||
server_args = ServerArgs(
|
||||
model_path=args.model, tp_size=args.tp_size, ep_size=args.ep_size
|
||||
)
|
||||
|
||||
model_config = get_model_config(
|
||||
args.model,
|
||||
args.tp_size,
|
||||
args.ep_size,
|
||||
args.disable_shared_experts_fusion,
|
||||
args.topk_ids_dir,
|
||||
)
|
||||
|
||||
E = model_config["num_experts"]
|
||||
topk = model_config["topk"]
|
||||
hidden_size = model_config["hidden_size"]
|
||||
shard_intermediate_size = model_config["shard_intermediate_size"]
|
||||
dtype = model_config["dtype"]
|
||||
block_shape = model_config["block_shape"]
|
||||
|
||||
use_fp8_w8a8 = args.dtype == "fp8_w8a8"
|
||||
use_int8_w8a8 = args.dtype == "int8_w8a8"
|
||||
use_int8_w8a16 = args.dtype == "int8_w8a16"
|
||||
use_int4_w4a16 = args.dtype == "int4_w4a16"
|
||||
|
||||
topk_ids_dir = args.topk_ids_dir
|
||||
if args.batch_size is None:
|
||||
batch_sizes = get_default_batch_sizes()
|
||||
batch_sizes.reverse()
|
||||
else:
|
||||
batch_sizes = [args.batch_size]
|
||||
|
||||
if args.cmp_configs is not None:
|
||||
worker = BenchmarkWorker(args.seed, server_args)
|
||||
worker.cmp_configs(
|
||||
batch_sizes,
|
||||
E,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
block_shape,
|
||||
args.cmp_configs,
|
||||
topk_ids_dir,
|
||||
args.ep_size,
|
||||
)
|
||||
return
|
||||
|
||||
if len(batch_sizes) == 1:
|
||||
worker = BenchmarkWorker(args.seed, server_args)
|
||||
if args.tune:
|
||||
search_space = get_configs_compute_bound()
|
||||
worker.tune(
|
||||
batch_sizes[0],
|
||||
E,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
block_shape,
|
||||
search_space,
|
||||
topk_ids_dir,
|
||||
args.ep_size,
|
||||
)
|
||||
else:
|
||||
cfg = {
|
||||
"BLOCK_SIZE_M": args.configs[0],
|
||||
"BLOCK_SIZE_N": args.configs[1],
|
||||
"BLOCK_SIZE_K": args.configs[2],
|
||||
"GROUP_SIZE_M": args.configs[3],
|
||||
"num_warps": args.configs[4],
|
||||
"num_stages": args.configs[5],
|
||||
}
|
||||
|
||||
_, (t0, t0_tma, t1, t1_tma) = worker.benchmark(
|
||||
args.batch_size,
|
||||
E,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
block_shape,
|
||||
cfg,
|
||||
topk_ids_dir,
|
||||
args.ep_size,
|
||||
)
|
||||
print(f"{t0=}, {t0_tma=}, {t1=}, {t1_tma=}")
|
||||
return
|
||||
|
||||
assert args.tune
|
||||
|
||||
ray.init()
|
||||
num_gpus = int(ray.available_resources()["GPU"])
|
||||
workers = [
|
||||
ray.remote(num_gpus=1)(BenchmarkWorker).remote(args.seed, server_args)
|
||||
for _ in range(num_gpus)
|
||||
]
|
||||
|
||||
def _distribute(method: str, inputs: List[Any]) -> List[Any]:
|
||||
outputs = []
|
||||
worker_idx = 0
|
||||
for input_args in inputs:
|
||||
worker = workers[worker_idx]
|
||||
worker_method = getattr(worker, method)
|
||||
output = worker_method.remote(*input_args)
|
||||
outputs.append(output)
|
||||
worker_idx = (worker_idx + 1) % num_gpus
|
||||
return ray.get(outputs)
|
||||
|
||||
search_space = get_configs_compute_bound()
|
||||
if block_shape is not None:
|
||||
block_n, block_k = block_shape[0], block_shape[1]
|
||||
search_space = [
|
||||
config for config in search_space if block_k % config["BLOCK_SIZE_K"] == 0
|
||||
]
|
||||
filename = get_config_filename(
|
||||
E,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
False,
|
||||
block_shape,
|
||||
)
|
||||
print(
|
||||
f"Start tuning over {len(search_space)} configurations to create {filename}..."
|
||||
)
|
||||
|
||||
start = time.perf_counter()
|
||||
configs = _distribute(
|
||||
"tune",
|
||||
[
|
||||
(
|
||||
batch_size,
|
||||
E,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
block_shape,
|
||||
search_space,
|
||||
topk_ids_dir,
|
||||
args.ep_size,
|
||||
)
|
||||
for batch_size in batch_sizes
|
||||
],
|
||||
)
|
||||
print(f"{configs=}", flush=True)
|
||||
cur_time = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime())
|
||||
with open(f"tuning_result_{cur_time}.txt", "w") as f:
|
||||
print(configs, file=f)
|
||||
batch_sizes.reverse()
|
||||
configs0 = [config[0] for config in configs]
|
||||
configs1 = [config[1] for config in configs]
|
||||
configs0.reverse()
|
||||
configs1.reverse()
|
||||
best_configs0 = {M: sort_config(config) for M, config in zip(batch_sizes, configs0)}
|
||||
save_configs_sep(
|
||||
best_configs0,
|
||||
E,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
block_shape,
|
||||
)
|
||||
|
||||
best_configs1 = {M: sort_config(config) for M, config in zip(batch_sizes, configs1)}
|
||||
save_configs_sep(
|
||||
best_configs1,
|
||||
E,
|
||||
shard_intermediate_size,
|
||||
hidden_size,
|
||||
topk,
|
||||
dtype,
|
||||
use_fp8_w8a8,
|
||||
use_int8_w8a8,
|
||||
use_int8_w8a16,
|
||||
use_int4_w4a16,
|
||||
block_shape,
|
||||
down_moe=True,
|
||||
)
|
||||
end = time.perf_counter()
|
||||
print(f"Tuning took {end - start:.2f} seconds")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--model", type=str, default="mistralai/Mixtral-8x7B-Instruct-v0.1"
|
||||
)
|
||||
parser.add_argument("--tp-size", "--tp", type=int, default=2)
|
||||
parser.add_argument("--ep-size", "--ep", type=int, default=1)
|
||||
parser.add_argument(
|
||||
"--dtype",
|
||||
type=str,
|
||||
choices=["auto", "fp8_w8a8", "int8_w8a16", "int8_w8a8", "int8_w4a16"],
|
||||
default="auto",
|
||||
)
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
parser.add_argument("--batch-size", type=int, required=False)
|
||||
parser.add_argument("--tune", action="store_true")
|
||||
parser.add_argument("--disable-shared-experts-fusion", action="store_true")
|
||||
parser.add_argument("--configs", type=int, nargs="+", required=False)
|
||||
parser.add_argument("--topk-ids-dir", type=str, required=True)
|
||||
parser.add_argument("--cmp-configs", type=str, nargs="+", required=False)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
1
third_party/sglang/benchmark/kernels/fused_moe_triton/tuning_text.json
vendored
Normal file
1
third_party/sglang/benchmark/kernels/fused_moe_triton/tuning_text.json
vendored
Normal file
File diff suppressed because one or more lines are too long
92
third_party/sglang/benchmark/kernels/quantization/README.md
vendored
Normal file
92
third_party/sglang/benchmark/kernels/quantization/README.md
vendored
Normal file
@@ -0,0 +1,92 @@
|
||||
# W8A8 Block-wise Quantization Kernel Tuning
|
||||
|
||||
Auto-tune Triton FP8/INT8 block-wise quantization kernels for optimal performance.
|
||||
|
||||
## When to Use Triton FP8 Block-wise Quantization Kernel vs DeepGEMM
|
||||
|
||||
**Use Triton FP8 Block-wise Quantization Kernel when:**
|
||||
- Output dtype is NOT `bfloat16` (e.g., `float16`, `float32`)
|
||||
- DeepGEMM is disabled (environment variable `SGLANG_ENABLE_JIT_DEEPGEMM=0`)
|
||||
- Running on GPUs with compute capability < SM90 (DeepGEMM requires SM90+)
|
||||
- You need cross-platform compatibility (Triton works on both NVIDIA and AMD GPUs)
|
||||
|
||||
**Use DeepGEMM when:**
|
||||
- Output dtype is `bfloat16` AND DeepGEMM is enabled
|
||||
- Running on NVIDIA GPUs with compute capability >= SM90 (e.g., H100, H200)
|
||||
- Need maximum performance for production workloads (DeepGEMM is highly optimized for Hopper architecture)
|
||||
|
||||
**Note:** DeepGEMM requires CUDA compute capability >= 9.0 (SM90+). It is specifically optimized for NVIDIA Hopper GPUs (H100/H200).
|
||||
|
||||
The kernel selection logic in SGLang automatically chooses DeepGEMM when conditions are met (see `w8a8_block_fp8_matmul` function in `fp8_kernel.py`), otherwise falls back to Triton implementation.
|
||||
|
||||
## Quick Start
|
||||
|
||||
**Default (DeepSeek-V3):**
|
||||
```bash
|
||||
python benchmark/kernels/quantization/tuning_block_wise_kernel.py --tp-size 8
|
||||
```
|
||||
|
||||
**Custom Model (specify N and K):**
|
||||
```bash
|
||||
python benchmark/kernels/quantization/tuning_block_wise_kernel.py --N 5120 --K 25600
|
||||
```
|
||||
|
||||
## Parameters
|
||||
|
||||
- `--N`, `--K`: Weight matrix dimensions (N=output_dim, K=input_dim). If not specified, uses `--tp-size` for DeepSeek-V3
|
||||
- `--tp-size`: Tensor parallelism size for DeepSeek-V3 (default: 8)
|
||||
- `--input-type`: `fp8` or `int8` (default: fp8)
|
||||
- `--block-n`, `--block-k`: Block quantization granularity (default: 128)
|
||||
- `--batch-size`: Test single batch size (optional)
|
||||
|
||||
## How to Calculate N and K
|
||||
|
||||
For a linear layer `y = xW^T` where `x` is (M, K) and `W` is (N, K):
|
||||
- **N**: Output features (weight matrix output dimension)
|
||||
- **K**: Input features (weight matrix input dimension)
|
||||
|
||||
**Example: Qwen3-VL-32B** (hidden_size=5120, intermediate_size=25600, num_heads=64, num_kv_heads=8, head_dim=128) and TP=1
|
||||
```bash
|
||||
# QKV projection: Q(8192) + K(1024) + V(1024) = 10240
|
||||
python benchmark/kernels/quantization/tuning_block_wise_kernel.py --N 10240 --K 5120
|
||||
|
||||
# MLP gate+up (SwiGLU): 2 * intermediate_size = 51200
|
||||
python benchmark/kernels/quantization/tuning_block_wise_kernel.py --N 51200 --K 5120
|
||||
|
||||
# MLP down projection
|
||||
python benchmark/kernels/quantization/tuning_block_wise_kernel.py --N 5120 --K 25600
|
||||
|
||||
# O projection (if separate from QKV)
|
||||
python benchmark/kernels/quantization/tuning_block_wise_kernel.py --N 5120 --K 8192
|
||||
```
|
||||
|
||||
If TP=8:
|
||||
|
||||
```bash
|
||||
# QKV projection: Q(8192) + K(1024) + V(1024) = 10240 / TP=8
|
||||
python benchmark/kernels/quantization/tuning_block_wise_kernel.py --N 1280 --K 5120
|
||||
|
||||
# MLP gate+up (SwiGLU): 2 * intermediate_size = 51200 / TP=8
|
||||
python benchmark/kernels/quantization/tuning_block_wise_kernel.py --N 6400 --K 5120
|
||||
|
||||
# MLP down projection
|
||||
python benchmark/kernels/quantization/tuning_block_wise_kernel.py --N 5120 --K 3200
|
||||
|
||||
# O projection (if separate from QKV)
|
||||
python benchmark/kernels/quantization/tuning_block_wise_kernel.py --N 5120 --K 1024
|
||||
```
|
||||
|
||||
## Output
|
||||
|
||||
Generates JSON config files saved to `python/sglang/srt/layers/quantization/configs/`:
|
||||
```
|
||||
N={N},K={K},device_name={DEVICE},dtype=fp8_w8a8,block_shape=[128,128].json
|
||||
```
|
||||
|
||||
Config maps batch size to optimal kernel parameters:
|
||||
```json
|
||||
{
|
||||
"1": {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 128, ...},
|
||||
"2048": {"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 128, ...}
|
||||
}
|
||||
```
|
||||
137
third_party/sglang/benchmark/kernels/quantization/bench_fp4_quant.py
vendored
Normal file
137
third_party/sglang/benchmark/kernels/quantization/bench_fp4_quant.py
vendored
Normal file
@@ -0,0 +1,137 @@
|
||||
import argparse
|
||||
import itertools
|
||||
|
||||
import torch
|
||||
import triton
|
||||
from flashinfer import (
|
||||
scaled_fp4_grouped_quantize,
|
||||
silu_and_mul_scaled_nvfp4_experts_quantize,
|
||||
)
|
||||
from sgl_kernel.elementwise import silu_and_mul
|
||||
|
||||
from sglang.benchmark.bench_utils import run_bench
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
from sglang.srt.layers.moe.ep_moe.kernels import silu_and_mul_masked_post_quant_fwd
|
||||
|
||||
|
||||
def _test_accuracy_once(E, M, K, input_dtype, device):
|
||||
x = torch.randn(E, M, K, device=device, dtype=input_dtype)
|
||||
glb_scales = torch.ones((E,), dtype=torch.float32, device=device)
|
||||
masks = torch.full((E,), M, dtype=torch.int32, device=device)
|
||||
out, blk_scales = silu_and_mul_scaled_nvfp4_experts_quantize(x, masks, glb_scales)
|
||||
out1, blk_scales1 = scaled_fp4_grouped_quantize(
|
||||
silu_and_mul(x),
|
||||
masks,
|
||||
glb_scales,
|
||||
)
|
||||
|
||||
torch.testing.assert_close(out, out1)
|
||||
torch.testing.assert_close(blk_scales, blk_scales1)
|
||||
print(f"E: {E}, M: {M}, K: {K}, type: {input_dtype} OK")
|
||||
|
||||
|
||||
NUM_RANKS = 48
|
||||
M_PER_RANKs = [128, 256, 512, 1024]
|
||||
Ms = [M_PER_RANK * NUM_RANKS for M_PER_RANK in M_PER_RANKs]
|
||||
Ks = [2048, 4096, 7168]
|
||||
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["M", "K"],
|
||||
x_vals=list(itertools.product(Ms, Ks)),
|
||||
x_log=False,
|
||||
line_arg="provider",
|
||||
line_vals=["triton_fp8", "cuda_unfused_fp4", "cuda_fused_fp4"],
|
||||
line_names=["triton_fp8", "cuda_unfused_fp4", "cuda_fused_fp4"],
|
||||
styles=[("blue", "-"), ("orange", "-"), ("green", "-")],
|
||||
ylabel="ms",
|
||||
plot_name="fp4 quant",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(M, K, provider):
|
||||
E = 6
|
||||
device = "cuda"
|
||||
x = torch.randn(E, M, K, device=device, dtype=torch.bfloat16)
|
||||
glb_scales = torch.ones((E,), dtype=torch.float32, device=device)
|
||||
masks = torch.randint(1, 4096, (E,), dtype=torch.int32, device=device)
|
||||
fp8_out = torch.empty(
|
||||
(
|
||||
x.shape[0],
|
||||
x.shape[1],
|
||||
x.shape[2] // 2,
|
||||
),
|
||||
device=x.device,
|
||||
dtype=torch.float8_e4m3fn,
|
||||
)
|
||||
scale_block_size = 128
|
||||
fp8_scales = torch.empty(
|
||||
(
|
||||
x.shape[0],
|
||||
x.shape[1],
|
||||
x.shape[2] // 2 // scale_block_size,
|
||||
),
|
||||
device=x.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
quantiles = (0.5, 0.2, 0.8)
|
||||
if provider == "triton_fp8":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: silu_and_mul_masked_post_quant_fwd(
|
||||
x,
|
||||
fp8_out,
|
||||
fp8_scales,
|
||||
scale_block_size,
|
||||
masks,
|
||||
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
|
||||
),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
if provider == "cuda_unfused_fp4":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: scaled_fp4_grouped_quantize(
|
||||
silu_and_mul(x),
|
||||
masks,
|
||||
glb_scales,
|
||||
),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
if provider == "cuda_fused_fp4":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: silu_and_mul_scaled_nvfp4_experts_quantize(
|
||||
x,
|
||||
masks,
|
||||
glb_scales,
|
||||
),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
|
||||
return ms, min_ms, max_ms
|
||||
|
||||
|
||||
def test_accuracy():
|
||||
E = 6
|
||||
N_RANKS = 48
|
||||
Ms = [128, 256, 512, 1024]
|
||||
Ks = [2048, 4096, 7168]
|
||||
input_dtype = torch.bfloat16
|
||||
for M in Ms:
|
||||
for K in Ks:
|
||||
_test_accuracy_once(E, N_RANKS * M, K, input_dtype, "cuda")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--save_path",
|
||||
type=str,
|
||||
default="./bench_fp4_quant_res",
|
||||
help="Path to save fp4 quant benchmark results",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
test_accuracy()
|
||||
|
||||
benchmark.run(print_data=True, show_plots=True, save_path=args.save_path)
|
||||
95
third_party/sglang/benchmark/kernels/quantization/bench_int8_quant.py
vendored
Normal file
95
third_party/sglang/benchmark/kernels/quantization/bench_int8_quant.py
vendored
Normal file
@@ -0,0 +1,95 @@
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
import triton
|
||||
from vllm._custom_ops import scaled_int8_quant as vllm_scaled_int8_quant
|
||||
|
||||
from sglang.benchmark.bench_utils import run_bench
|
||||
from sglang.srt.layers.quantization.int8_kernel import per_token_quant_int8
|
||||
|
||||
|
||||
@torch.compile(backend="inductor")
|
||||
def torch_int8_quant(x):
|
||||
int8_max = torch.iinfo(torch.int8).max
|
||||
|
||||
abs_max = x.abs().max(dim=-1, keepdim=True).values
|
||||
scales = abs_max.to(torch.float32) / float(int8_max)
|
||||
|
||||
q_x = (x / scales).round().to(torch.int8)
|
||||
|
||||
return q_x, scales
|
||||
|
||||
|
||||
def _test_accuracy_once(M, K, input_dtype, device):
|
||||
x = torch.randn(M, K, dtype=input_dtype, device=device) * 5000
|
||||
out, scales, _ = vllm_scaled_int8_quant(x, symmetric=True)
|
||||
out1, scales1 = per_token_quant_int8(x)
|
||||
out2, scales2 = torch_int8_quant(x)
|
||||
torch.testing.assert_close(out, out2, atol=1, rtol=0)
|
||||
torch.testing.assert_close(out, out1, atol=1, rtol=0)
|
||||
torch.testing.assert_close(scales, scales2)
|
||||
torch.testing.assert_close(scales1, scales2)
|
||||
print(f"M: {M}, K: {K}, type: {input_dtype} OK")
|
||||
|
||||
|
||||
def test_accuracy():
|
||||
Ms = [1, 13, 128, 1024, 2048, 4096]
|
||||
Ks = [512, 1024, 2048, 8192]
|
||||
input_dtypes = [torch.float16, torch.bfloat16]
|
||||
for M in Ms:
|
||||
for K in Ks:
|
||||
for input_dtype in input_dtypes:
|
||||
_test_accuracy_once(M, K, input_dtype, "cuda")
|
||||
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["batch_size"],
|
||||
x_vals=[1, 16, 32, 64, 128, 256, 512, 1024, 2048],
|
||||
x_log=False,
|
||||
line_arg="provider",
|
||||
line_vals=["vllm op", "triton", "torch.compile"],
|
||||
line_names=["vllm op", "triton", "torch.compile"],
|
||||
styles=[("blue", "-"), ("orange", "-"), ("red", "-")],
|
||||
ylabel="ms",
|
||||
plot_name="int8 per token quant",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(batch_size, provider):
|
||||
M, K = batch_size, 16384
|
||||
x = torch.randn(M, K, dtype=torch.float16, device="cuda") * 1000
|
||||
|
||||
quantiles = (0.5, 0.2, 0.8)
|
||||
if provider == "vllm op":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: vllm_scaled_int8_quant(x, symmetric=True),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
if provider == "triton":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: per_token_quant_int8(x),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
if provider == "torch.compile":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: torch_int8_quant(x),
|
||||
quantiles=quantiles,
|
||||
)
|
||||
|
||||
return ms, min_ms, max_ms
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--save_path",
|
||||
type=str,
|
||||
default="./bench_int8_quant_res",
|
||||
help="Path to save int8 quant benchmark results",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
test_accuracy()
|
||||
|
||||
benchmark.run(print_data=True, show_plots=True, save_path=args.save_path)
|
||||
527
third_party/sglang/benchmark/kernels/quantization/tuning_block_wise_kernel.py
vendored
Normal file
527
third_party/sglang/benchmark/kernels/quantization/tuning_block_wise_kernel.py
vendored
Normal file
@@ -0,0 +1,527 @@
|
||||
# Copyright 2025 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import multiprocessing as mp
|
||||
import os
|
||||
import time
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import torch
|
||||
import triton
|
||||
from tqdm import tqdm
|
||||
|
||||
mp.set_start_method("spawn", force=True)
|
||||
|
||||
from sglang.srt.layers.quantization.fp8_kernel import (
|
||||
_w8a8_block_fp8_matmul,
|
||||
_w8a8_block_fp8_matmul_unrolledx4,
|
||||
)
|
||||
from sglang.srt.layers.quantization.int8_kernel import _w8a8_block_int8_matmul
|
||||
from sglang.srt.utils import (
|
||||
get_device,
|
||||
get_device_core_count,
|
||||
get_device_count,
|
||||
get_device_name,
|
||||
is_hip,
|
||||
)
|
||||
|
||||
_is_hip = is_hip()
|
||||
|
||||
DTYPE_MAP = {
|
||||
"float32": torch.float32,
|
||||
"float16": torch.float16,
|
||||
"half": torch.half,
|
||||
"bfloat16": torch.bfloat16,
|
||||
}
|
||||
|
||||
|
||||
def w8a8_block_matmul(
|
||||
A: torch.Tensor,
|
||||
B: torch.Tensor,
|
||||
As: torch.Tensor,
|
||||
Bs: torch.Tensor,
|
||||
block_size: List[int],
|
||||
config: Dict[str, Any],
|
||||
output_dtype: torch.dtype = torch.float16,
|
||||
) -> torch.Tensor:
|
||||
"""This function performs matrix multiplication with block-wise quantization.
|
||||
|
||||
It takes two input tensors `A` and `B` with scales `As` and `Bs`.
|
||||
The output is returned in the specified `output_dtype`.
|
||||
|
||||
Args:
|
||||
A: The input tensor, e.g., activation.
|
||||
B: The input tensor, e.g., weight.
|
||||
As: The per-token-group quantization scale for `A`.
|
||||
Bs: The per-block quantization scale for `B`.
|
||||
block_size: The block size for per-block quantization. It should be 2-dim, e.g., [128, 128].
|
||||
output_dytpe: The dtype of the returned tensor.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: The result of matmul.
|
||||
"""
|
||||
assert len(block_size) == 2
|
||||
block_n, block_k = block_size[0], block_size[1]
|
||||
|
||||
assert A.shape[-1] == B.shape[-1]
|
||||
assert A.shape[:-1] == As.shape[:-1] and A.is_contiguous()
|
||||
assert triton.cdiv(A.shape[-1], block_k) == As.shape[-1]
|
||||
M = A.numel() // A.shape[-1]
|
||||
|
||||
assert B.ndim == 2 and B.is_contiguous() and Bs.ndim == 2
|
||||
N, K = B.shape
|
||||
assert triton.cdiv(N, block_n) == Bs.shape[0]
|
||||
assert triton.cdiv(K, block_k) == Bs.shape[1]
|
||||
|
||||
C_shape = A.shape[:-1] + (N,)
|
||||
C = A.new_empty(C_shape, dtype=output_dtype)
|
||||
|
||||
needs_masking = bool(K % config["BLOCK_SIZE_K"] != 0)
|
||||
|
||||
def grid(META):
|
||||
return (
|
||||
triton.cdiv(M, META["BLOCK_SIZE_M"]) * triton.cdiv(N, META["BLOCK_SIZE_N"]),
|
||||
)
|
||||
|
||||
# Use manually unrolledx4 kernel on AMD GPU when the grid size is small.
|
||||
# Empirical testing shows the sweet spot lies when it's less than the # of
|
||||
# compute units available on the device.
|
||||
num_workgroups = triton.cdiv(M, config["BLOCK_SIZE_M"]) * triton.cdiv(
|
||||
N, config["BLOCK_SIZE_N"]
|
||||
)
|
||||
|
||||
if A.dtype == torch.float8_e4m3fnuz or A.dtype == torch.float8_e4m3fn:
|
||||
kernel = (
|
||||
_w8a8_block_fp8_matmul_unrolledx4
|
||||
if (_is_hip == True and num_workgroups <= get_device_core_count())
|
||||
else _w8a8_block_fp8_matmul
|
||||
)
|
||||
else:
|
||||
kernel = _w8a8_block_int8_matmul
|
||||
|
||||
kernel[grid](
|
||||
A,
|
||||
B,
|
||||
C,
|
||||
As,
|
||||
Bs,
|
||||
M,
|
||||
N,
|
||||
K,
|
||||
block_n,
|
||||
block_k,
|
||||
A.stride(-2),
|
||||
A.stride(-1),
|
||||
B.stride(1),
|
||||
B.stride(0),
|
||||
C.stride(-2),
|
||||
C.stride(-1),
|
||||
As.stride(-2),
|
||||
As.stride(-1),
|
||||
Bs.stride(1),
|
||||
Bs.stride(0),
|
||||
**config,
|
||||
needs_masking=needs_masking,
|
||||
)
|
||||
|
||||
return C
|
||||
|
||||
|
||||
def get_rocm_configs_compute_bound():
|
||||
configs = []
|
||||
waves_per_eu_range = 0
|
||||
for num_stages in [2]:
|
||||
for block_m in [32, 64, 128, 256]:
|
||||
for block_k in [32, 64, 128, 256]:
|
||||
for block_n in [16, 32, 64, 128, 256]:
|
||||
for num_warps in [4, 8]:
|
||||
for group_size in [1, 4, 8, 16, 32]:
|
||||
configs.append(
|
||||
{
|
||||
"BLOCK_SIZE_M": block_m,
|
||||
"BLOCK_SIZE_N": block_n,
|
||||
"BLOCK_SIZE_K": block_k,
|
||||
"GROUP_SIZE_M": group_size,
|
||||
"num_warps": num_warps,
|
||||
"num_stages": num_stages,
|
||||
"waves_per_eu": waves_per_eu_range,
|
||||
}
|
||||
)
|
||||
return configs
|
||||
|
||||
|
||||
def get_configs_compute_bound():
|
||||
configs = []
|
||||
if _is_hip:
|
||||
configs = get_rocm_configs_compute_bound()
|
||||
else:
|
||||
for num_stages in [2, 3, 4, 5]:
|
||||
for block_m in [16, 32, 64, 128, 256]:
|
||||
for block_k in [64, 128]:
|
||||
for block_n in [32, 64, 128, 256]:
|
||||
for num_warps in [4, 8]:
|
||||
for group_size in [1, 16, 32, 64]:
|
||||
configs.append(
|
||||
{
|
||||
"BLOCK_SIZE_M": block_m,
|
||||
"BLOCK_SIZE_N": block_n,
|
||||
"BLOCK_SIZE_K": block_k,
|
||||
"GROUP_SIZE_M": group_size,
|
||||
"num_warps": num_warps,
|
||||
"num_stages": num_stages,
|
||||
}
|
||||
)
|
||||
return configs
|
||||
|
||||
|
||||
def get_weight_shapes(tp_size):
|
||||
# NOTE(HandH1998): The weight shapes only works for DeepSeek-V3. Modify them, if you tune for another different model.
|
||||
# cannot TP
|
||||
total = [
|
||||
(512 + 64, 7168),
|
||||
((128 + 64) * 128, 7168),
|
||||
(128 * (128 + 128), 512),
|
||||
(7168, 16384),
|
||||
(7168, 18432),
|
||||
]
|
||||
# N can TP
|
||||
n_tp = [
|
||||
(18432 * 2, 7168),
|
||||
((128 + 64) * 128, 7168),
|
||||
(128 * (128 + 128), 512),
|
||||
(24576, 1536),
|
||||
(4096, 7168),
|
||||
]
|
||||
# K can TP
|
||||
k_tp = [(7168, 18432), (7168, 16384), (7168, 2048)]
|
||||
|
||||
weight_shapes = []
|
||||
for t in total:
|
||||
weight_shapes.append(t)
|
||||
for n_t in n_tp:
|
||||
new_t = (n_t[0] // tp_size, n_t[1])
|
||||
weight_shapes.append(new_t)
|
||||
for k_t in k_tp:
|
||||
new_t = (k_t[0], k_t[1] // tp_size)
|
||||
weight_shapes.append(new_t)
|
||||
return weight_shapes
|
||||
|
||||
|
||||
def benchmark_config(
|
||||
A, B, As, Bs, block_size, config, out_dtype=torch.float16, num_iters=10
|
||||
):
|
||||
def run():
|
||||
w8a8_block_matmul(A, B, As, Bs, block_size, config, out_dtype)
|
||||
|
||||
torch.get_device_module().synchronize()
|
||||
# JIT complication & warmup
|
||||
for _ in range(5):
|
||||
run()
|
||||
torch.get_device_module().synchronize()
|
||||
|
||||
start_event = torch.get_device_module().Event(enable_timing=True)
|
||||
end_event = torch.get_device_module().Event(enable_timing=True)
|
||||
|
||||
latencies: List[float] = []
|
||||
for i in range(num_iters):
|
||||
torch.get_device_module().synchronize()
|
||||
start_event.record()
|
||||
run()
|
||||
end_event.record()
|
||||
end_event.synchronize()
|
||||
latencies.append(start_event.elapsed_time(end_event))
|
||||
avg = sum(latencies) / (num_iters * 10) * 1000 # us
|
||||
return avg
|
||||
|
||||
|
||||
def tune(M, N, K, block_size, out_dtype, search_space, input_type):
|
||||
factor_for_scale = 1e-2
|
||||
device = get_device()
|
||||
|
||||
if input_type == "fp8":
|
||||
fp8_info = torch.finfo(
|
||||
torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn
|
||||
)
|
||||
fp8_max, fp8_min = fp8_info.max, fp8_info.min
|
||||
|
||||
A_fp32 = (
|
||||
(torch.rand(M, K, dtype=torch.float32, device=device) - 0.5) * 2 * fp8_max
|
||||
)
|
||||
A = A_fp32.clamp(min=fp8_min, max=fp8_max).to(
|
||||
torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn
|
||||
)
|
||||
|
||||
B_fp32 = (
|
||||
(torch.rand(N, K, dtype=torch.float32, device=device) - 0.5) * 2 * fp8_max
|
||||
)
|
||||
B = B_fp32.clamp(min=fp8_min, max=fp8_max).to(
|
||||
torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn
|
||||
)
|
||||
else:
|
||||
int8_info = torch.iinfo(torch.int8)
|
||||
int8_max, int8_min = int8_info.max, int8_info.min
|
||||
|
||||
A_fp32 = (
|
||||
(torch.rand(M, K, dtype=torch.float32, device=device) - 0.5) * 2 * int8_max
|
||||
)
|
||||
A = A_fp32.clamp(min=int8_min, max=int8_max).to(torch.int8)
|
||||
|
||||
B_fp32 = (
|
||||
(torch.rand(N, K, dtype=torch.float32, device=device) - 0.5) * 2 * int8_max
|
||||
)
|
||||
B = B_fp32.clamp(min=int8_min, max=int8_max).to(torch.int8)
|
||||
|
||||
block_n, block_k = block_size[0], block_size[1]
|
||||
n_tiles = (N + block_n - 1) // block_n
|
||||
k_tiles = (K + block_k - 1) // block_k
|
||||
|
||||
As = torch.rand(M, k_tiles, dtype=torch.float32, device=device) * factor_for_scale
|
||||
Bs = (
|
||||
torch.rand(n_tiles, k_tiles, dtype=torch.float32, device=device)
|
||||
* factor_for_scale
|
||||
)
|
||||
|
||||
best_config = None
|
||||
best_time = float("inf")
|
||||
for config in tqdm(search_space):
|
||||
try:
|
||||
kernel_time = benchmark_config(
|
||||
A,
|
||||
B,
|
||||
As,
|
||||
Bs,
|
||||
block_size,
|
||||
config,
|
||||
out_dtype,
|
||||
num_iters=10,
|
||||
)
|
||||
except triton.runtime.autotuner.OutOfResources:
|
||||
# Some configurations may be invalid and fail to compile.
|
||||
continue
|
||||
|
||||
if kernel_time < best_time:
|
||||
best_time = kernel_time
|
||||
best_config = config
|
||||
now = datetime.now()
|
||||
print(f"{now.ctime()}] Completed tuning for batch_size={M}")
|
||||
assert best_config is not None
|
||||
return best_config
|
||||
|
||||
|
||||
def save_configs(
|
||||
N,
|
||||
K,
|
||||
block_n,
|
||||
block_k,
|
||||
configs,
|
||||
save_path,
|
||||
input_type="fp8",
|
||||
lock=None,
|
||||
) -> None:
|
||||
os.makedirs(save_path, exist_ok=True)
|
||||
device_name = get_device_name().replace(" ", "_")
|
||||
json_file_name = f"N={N},K={K},device_name={device_name},dtype={input_type}_w8a8,block_shape=[{block_n}, {block_k}].json"
|
||||
|
||||
config_file_path = os.path.join(save_path, json_file_name)
|
||||
print(f"Writing best config to {config_file_path}...")
|
||||
|
||||
if lock is not None:
|
||||
lock.acquire()
|
||||
try:
|
||||
existing_configs = {}
|
||||
if os.path.exists(config_file_path):
|
||||
with open(config_file_path, "r") as f:
|
||||
existing_configs = json.load(f)
|
||||
existing_configs = {int(k): v for k, v in existing_configs.items()}
|
||||
|
||||
existing_configs.update(configs)
|
||||
|
||||
with open(config_file_path, "w") as f:
|
||||
json.dump(existing_configs, f, indent=4)
|
||||
f.write("\n")
|
||||
finally:
|
||||
if lock is not None:
|
||||
lock.release()
|
||||
|
||||
|
||||
def tune_on_gpu(args_dict):
|
||||
"""Run tuning on a specific GPU."""
|
||||
gpu_id = args_dict["gpu_id"]
|
||||
batch_sizes = args_dict["batch_sizes"]
|
||||
weight_shapes = args_dict["weight_shapes"]
|
||||
args = args_dict["args"]
|
||||
lock = args_dict["lock"]
|
||||
|
||||
torch.get_device_module().set_device(gpu_id)
|
||||
print(f"Starting tuning on GPU {gpu_id} with batch sizes {batch_sizes}")
|
||||
|
||||
block_n = args.block_n
|
||||
block_k = args.block_k
|
||||
out_dtype = DTYPE_MAP[args.out_dtype]
|
||||
save_path = args.save_path
|
||||
input_type = args.input_type
|
||||
|
||||
search_space = get_configs_compute_bound()
|
||||
search_space = [
|
||||
config for config in search_space if block_k % config["BLOCK_SIZE_K"] == 0
|
||||
]
|
||||
|
||||
start = time.perf_counter()
|
||||
results = {}
|
||||
for shape in tqdm(weight_shapes, desc=f"GPU {gpu_id} - Shapes"):
|
||||
N, K = shape[0], shape[1]
|
||||
print(f"[GPU {gpu_id}] Tune for weight shape of `N: {N}, K: {K}`")
|
||||
benchmark_results = [
|
||||
tune(
|
||||
batch_size,
|
||||
N,
|
||||
K,
|
||||
[block_n, block_k],
|
||||
out_dtype,
|
||||
search_space,
|
||||
input_type,
|
||||
)
|
||||
for batch_size in tqdm(batch_sizes, desc=f"GPU {gpu_id} - Batch sizes")
|
||||
]
|
||||
best_configs = {M: config for M, config in zip(batch_sizes, benchmark_results)}
|
||||
save_configs(N, K, block_n, block_k, best_configs, save_path, input_type, lock)
|
||||
|
||||
end = time.perf_counter()
|
||||
print(f"Tuning on GPU {gpu_id} took {end - start:.2f} seconds")
|
||||
|
||||
|
||||
def distribute_batch_sizes(batch_sizes, num_gpus):
|
||||
"""Distribute batch sizes across available GPUs."""
|
||||
batches_per_gpu = []
|
||||
for i in range(num_gpus):
|
||||
start_idx = i * len(batch_sizes) // num_gpus
|
||||
end_idx = (i + 1) * len(batch_sizes) // num_gpus
|
||||
batches_per_gpu.append(batch_sizes[start_idx:end_idx])
|
||||
return batches_per_gpu
|
||||
|
||||
|
||||
def main(args):
|
||||
print(args)
|
||||
|
||||
num_gpus = get_device_count()
|
||||
if num_gpus == 0:
|
||||
raise RuntimeError("No GPU available for tuning")
|
||||
print(f"Found {num_gpus} GPUs for parallel tuning")
|
||||
|
||||
torch.get_device_module().init()
|
||||
|
||||
if args.batch_size is None:
|
||||
batch_sizes = [
|
||||
1,
|
||||
2,
|
||||
4,
|
||||
8,
|
||||
16,
|
||||
24,
|
||||
32,
|
||||
48,
|
||||
64,
|
||||
96,
|
||||
128,
|
||||
256,
|
||||
512,
|
||||
1024,
|
||||
1536,
|
||||
2048,
|
||||
3072,
|
||||
4096,
|
||||
]
|
||||
else:
|
||||
batch_sizes = [args.batch_size]
|
||||
num_gpus = 1 # If only one batch size, use only one GPU
|
||||
|
||||
# Support manual N and K specification
|
||||
if args.N is not None and args.K is not None:
|
||||
weight_shapes = [(args.N, args.K)]
|
||||
print(f"Using manually specified weight shape: N={args.N}, K={args.K}")
|
||||
else:
|
||||
weight_shapes = get_weight_shapes(args.tp_size)
|
||||
print(f"Using predefined weight shapes for TP size {args.tp_size}")
|
||||
|
||||
batches_per_gpu = distribute_batch_sizes(batch_sizes, num_gpus)
|
||||
|
||||
ctx = mp.get_context("spawn")
|
||||
manager = ctx.Manager()
|
||||
lock = manager.Lock()
|
||||
|
||||
process_args = []
|
||||
for gpu_id in range(num_gpus):
|
||||
process_args.append(
|
||||
{
|
||||
"gpu_id": gpu_id,
|
||||
"batch_sizes": batches_per_gpu[gpu_id],
|
||||
"weight_shapes": weight_shapes, # Each GPU processes all weight shapes
|
||||
"args": args,
|
||||
"lock": lock,
|
||||
}
|
||||
)
|
||||
|
||||
with ctx.Pool(num_gpus) as pool:
|
||||
pool.map(tune_on_gpu, process_args)
|
||||
|
||||
print("Multi-GPU tuning completed")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument(
|
||||
"--tp-size",
|
||||
"-tp",
|
||||
type=int,
|
||||
default=8,
|
||||
help="Tensor parallelism size (ignored if --N and --K are specified)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--N",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Output dimension of weight matrix (number of columns)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--K",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Input dimension of weight matrix (number of rows)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--input-type", type=str, choices=["fp8", "int8"], default="fp8"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--out-dtype",
|
||||
type=str,
|
||||
choices=["float32", "float16", "bfloat16", "half"],
|
||||
default="float16",
|
||||
)
|
||||
parser.add_argument("--block-n", type=int, default=128)
|
||||
parser.add_argument("--block-k", type=int, default=128)
|
||||
parser.add_argument("--batch-size", type=int, required=False)
|
||||
parser.add_argument(
|
||||
"--save-path", type=str, default="python/sglang/srt/layers/quantization/configs"
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
# Validate arguments
|
||||
if (args.N is None) != (args.K is None):
|
||||
parser.error("--N and --K must be specified together or not at all")
|
||||
|
||||
main(args)
|
||||
171
third_party/sglang/benchmark/kernels/scheduler_batch/benchmark_get_last_loc_triton.py
vendored
Normal file
171
third_party/sglang/benchmark/kernels/scheduler_batch/benchmark_get_last_loc_triton.py
vendored
Normal file
@@ -0,0 +1,171 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.benchmark.bench_utils import run_bench
|
||||
|
||||
|
||||
@torch.compile(dynamic=True)
|
||||
def get_last_loc_torch(
|
||||
req_to_token: torch.Tensor,
|
||||
req_pool_indices_tensor: torch.Tensor,
|
||||
prefix_lens_tensor: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
return torch.where(
|
||||
prefix_lens_tensor > 0,
|
||||
req_to_token[req_pool_indices_tensor, prefix_lens_tensor - 1],
|
||||
torch.full_like(prefix_lens_tensor, -1),
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def get_last_loc_kernel(
|
||||
req_to_token,
|
||||
req_pool_indices_tensor,
|
||||
prefix_lens_tensor,
|
||||
result,
|
||||
num_tokens,
|
||||
req_to_token_stride,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
pid = tl.program_id(0)
|
||||
offset = tl.arange(0, BLOCK_SIZE) + pid * BLOCK_SIZE
|
||||
mask = offset < num_tokens
|
||||
|
||||
prefix_lens = tl.load(prefix_lens_tensor + offset, mask=mask, other=0)
|
||||
req_pool_indices = tl.load(req_pool_indices_tensor + offset, mask=mask, other=0)
|
||||
|
||||
token_mask = prefix_lens > 0
|
||||
token_index = req_pool_indices * req_to_token_stride + (prefix_lens - 1)
|
||||
tokens = tl.load(req_to_token + token_index, mask=token_mask, other=-1)
|
||||
|
||||
tl.store(result + offset, tokens, mask=mask)
|
||||
|
||||
|
||||
def get_last_loc_triton(
|
||||
req_to_token: torch.Tensor,
|
||||
req_pool_indices_tensor: torch.Tensor,
|
||||
prefix_lens_tensor: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
BLOCK_SIZE = 256
|
||||
num_tokens = prefix_lens_tensor.shape[0]
|
||||
result = torch.empty_like(prefix_lens_tensor)
|
||||
grid = (triton.cdiv(num_tokens, BLOCK_SIZE),)
|
||||
|
||||
get_last_loc_kernel[grid](
|
||||
req_to_token,
|
||||
req_pool_indices_tensor,
|
||||
prefix_lens_tensor,
|
||||
result,
|
||||
num_tokens,
|
||||
req_to_token.stride(0),
|
||||
BLOCK_SIZE,
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
def test_get_last_loc():
|
||||
max_batch = 4097
|
||||
max_context_len = 6148
|
||||
batch_size = 20
|
||||
|
||||
# Initialize input tensors
|
||||
req_to_token = torch.zeros(
|
||||
(max_batch, max_context_len), dtype=torch.int32, device="cuda"
|
||||
)
|
||||
req_pool_indices = torch.arange(batch_size, dtype=torch.int64, device="cuda")
|
||||
pre_lens = torch.randint(
|
||||
-max_context_len // 2,
|
||||
max_context_len,
|
||||
(batch_size,),
|
||||
dtype=torch.int64,
|
||||
device="cuda",
|
||||
)
|
||||
|
||||
last_loc_res = get_last_loc_triton(req_to_token, req_pool_indices, pre_lens)
|
||||
last_loc_ref = get_last_loc_torch(req_to_token, req_pool_indices, pre_lens)
|
||||
|
||||
# Compare results
|
||||
torch.testing.assert_close(last_loc_res, last_loc_ref)
|
||||
|
||||
|
||||
def get_benchmark():
|
||||
batch_sizes = [1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024]
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["batch_size"],
|
||||
x_vals=batch_sizes,
|
||||
line_arg="provider",
|
||||
line_vals=["reference", "triton"],
|
||||
line_names=["PyTorch", "Triton"],
|
||||
styles=[("blue", "-"), ("green", "-")],
|
||||
ylabel="us",
|
||||
plot_name="get-last-loc-performance",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(batch_size, provider):
|
||||
max_batch = 2048
|
||||
max_context_len = 16384
|
||||
|
||||
req_to_token = torch.zeros(
|
||||
(max_batch, max_context_len), dtype=torch.int32, device="cuda"
|
||||
)
|
||||
req_pool_indices = torch.arange(batch_size, dtype=torch.int64, device="cuda")
|
||||
pre_lens = torch.randint(
|
||||
-max_context_len // 2,
|
||||
max_context_len,
|
||||
(batch_size,),
|
||||
dtype=torch.int64,
|
||||
device="cuda",
|
||||
)
|
||||
|
||||
quantiles = [0.5, 0.2, 0.8]
|
||||
|
||||
if provider == "reference":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: get_last_loc_torch(req_to_token, req_pool_indices, pre_lens),
|
||||
quantiles=tuple(quantiles),
|
||||
)
|
||||
elif provider == "triton":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: get_last_loc_triton(req_to_token, req_pool_indices, pre_lens),
|
||||
quantiles=tuple(quantiles),
|
||||
)
|
||||
|
||||
return 1000 * ms, 1000 * max_ms, 1000 * min_ms
|
||||
|
||||
return benchmark
|
||||
|
||||
|
||||
def run_benchmark(save_path: str = "./configs/benchmark_ops/get_last_loc/"):
|
||||
"""Run benchmark and save results"""
|
||||
|
||||
# Ensure save path exists
|
||||
os.makedirs(save_path, exist_ok=True)
|
||||
|
||||
# Run correctness test
|
||||
test_get_last_loc()
|
||||
print("Correctness test passed!")
|
||||
|
||||
# Run performance test
|
||||
benchmark = get_benchmark()
|
||||
benchmark.run(print_data=True, save_path=save_path)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--save_path",
|
||||
type=str,
|
||||
default="./configs/benchmark_ops/get_last_loc/",
|
||||
help="Path to save benchmark results",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
run_benchmark(args.save_path)
|
||||
342
third_party/sglang/benchmark/kernels/scheduler_batch/benchmark_write_req_to_token_pool_triton.py
vendored
Normal file
342
third_party/sglang/benchmark/kernels/scheduler_batch/benchmark_write_req_to_token_pool_triton.py
vendored
Normal file
@@ -0,0 +1,342 @@
|
||||
import itertools
|
||||
import os
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.benchmark.bench_utils import run_bench
|
||||
|
||||
|
||||
@triton.jit
|
||||
def write_req_to_token_pool_triton(
|
||||
req_to_token_ptr, # [max_batch, max_context_len]
|
||||
req_pool_indices,
|
||||
pre_lens,
|
||||
seq_lens,
|
||||
extend_lens,
|
||||
out_cache_loc,
|
||||
req_to_token_ptr_stride: tl.constexpr,
|
||||
):
|
||||
BLOCK_SIZE: tl.constexpr = 512
|
||||
pid = tl.program_id(0)
|
||||
|
||||
req_pool_index = tl.load(req_pool_indices + pid)
|
||||
pre_len = tl.load(pre_lens + pid)
|
||||
seq_len = tl.load(seq_lens + pid)
|
||||
|
||||
# TODO: optimize this?
|
||||
cumsum_start = 0
|
||||
for i in range(pid):
|
||||
cumsum_start += tl.load(extend_lens + i)
|
||||
|
||||
num_loop = tl.cdiv(seq_len - pre_len, BLOCK_SIZE)
|
||||
for i in range(num_loop):
|
||||
offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
|
||||
mask = offset < (seq_len - pre_len)
|
||||
value = tl.load(out_cache_loc + cumsum_start + offset, mask=mask)
|
||||
tl.store(
|
||||
req_to_token_ptr
|
||||
+ req_pool_index * req_to_token_ptr_stride
|
||||
+ offset
|
||||
+ pre_len,
|
||||
value,
|
||||
mask=mask,
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def write_req_to_token_pool_triton_optimize(
|
||||
req_to_token_ptr, # [max_batch, max_context_len]
|
||||
req_pool_indices,
|
||||
pre_lens,
|
||||
seq_lens,
|
||||
extend_lens,
|
||||
out_cache_loc,
|
||||
req_to_token_ptr_stride: tl.constexpr,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
pid_batch = tl.program_id(0)
|
||||
pid_token = tl.program_id(1)
|
||||
|
||||
req_pool_index = tl.load(req_pool_indices + pid_batch)
|
||||
pre_len = tl.load(pre_lens + pid_batch)
|
||||
seq_len = tl.load(seq_lens + pid_batch)
|
||||
extend_len = seq_len - pre_len
|
||||
|
||||
cumsum_start = 0
|
||||
for i in range(pid_batch):
|
||||
cumsum_start += tl.load(extend_lens + i)
|
||||
|
||||
token_start = pid_token * BLOCK_SIZE
|
||||
|
||||
offset = tl.arange(0, BLOCK_SIZE)
|
||||
actual_offset = token_start + offset
|
||||
mask = actual_offset < extend_len
|
||||
|
||||
src_ptr = out_cache_loc + cumsum_start + actual_offset
|
||||
src_ptr = tl.max_contiguous(tl.multiple_of(src_ptr, BLOCK_SIZE), BLOCK_SIZE)
|
||||
value = tl.load(src_ptr, mask=mask)
|
||||
dst_ptr = (
|
||||
req_to_token_ptr
|
||||
+ req_pool_index * req_to_token_ptr_stride
|
||||
+ actual_offset
|
||||
+ pre_len
|
||||
)
|
||||
dst_ptr = tl.max_contiguous(tl.multiple_of(dst_ptr, BLOCK_SIZE), BLOCK_SIZE)
|
||||
|
||||
tl.store(dst_ptr, value, mask=mask)
|
||||
|
||||
|
||||
def write_req_to_token_pool_reference(
|
||||
req_to_token: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
pre_lens: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
extend_lens: torch.Tensor,
|
||||
out_cache_loc: torch.Tensor,
|
||||
) -> None:
|
||||
"""Reference implementation using PyTorch"""
|
||||
for i in range(len(req_pool_indices)):
|
||||
req_pool_idx = req_pool_indices[i].item()
|
||||
pre_len = pre_lens[i].item()
|
||||
seq_len = seq_lens[i].item()
|
||||
extend_len = extend_lens[i].item()
|
||||
|
||||
cumsum_start = sum(extend_lens[:i].tolist())
|
||||
|
||||
# Copy values from out_cache_loc to req_to_token
|
||||
req_to_token[req_pool_idx, pre_len:seq_len] = out_cache_loc[
|
||||
cumsum_start : cumsum_start + extend_len
|
||||
]
|
||||
|
||||
|
||||
def test_write_req_to_token_pool():
|
||||
max_batch = 4097
|
||||
max_context_len = 6148
|
||||
batch_size = 1
|
||||
extend_len = 14
|
||||
|
||||
# Initialize input tensors
|
||||
req_to_token = torch.zeros(
|
||||
(max_batch, max_context_len), dtype=torch.int32, device="cuda"
|
||||
)
|
||||
req_pool_indices = torch.tensor([42], dtype=torch.int32, device="cuda")
|
||||
pre_lens = torch.tensor([8], dtype=torch.int32, device="cuda")
|
||||
seq_lens = torch.tensor([22], dtype=torch.int32, device="cuda")
|
||||
extend_lens = torch.tensor([extend_len], dtype=torch.int32, device="cuda")
|
||||
out_cache_loc = torch.arange(extend_len, dtype=torch.int32, device="cuda")
|
||||
|
||||
# Create copies for reference implementation
|
||||
req_to_token_ref = req_to_token.clone()
|
||||
req_to_token_opt = req_to_token.clone()
|
||||
|
||||
# Run original triton kernel
|
||||
write_req_to_token_pool_triton[(batch_size,)](
|
||||
req_to_token,
|
||||
req_pool_indices,
|
||||
pre_lens,
|
||||
seq_lens,
|
||||
extend_lens,
|
||||
out_cache_loc,
|
||||
max_context_len,
|
||||
)
|
||||
|
||||
# Run optimized triton kernel
|
||||
def grid(batch_size, extend_len):
|
||||
num_token_blocks = triton.cdiv(extend_len, 512)
|
||||
return (batch_size, num_token_blocks)
|
||||
|
||||
write_req_to_token_pool_triton_optimize[grid(batch_size, extend_len)](
|
||||
req_to_token_opt,
|
||||
req_pool_indices,
|
||||
pre_lens,
|
||||
seq_lens,
|
||||
extend_lens,
|
||||
out_cache_loc,
|
||||
max_context_len,
|
||||
BLOCK_SIZE=512,
|
||||
)
|
||||
|
||||
# Run reference implementation
|
||||
write_req_to_token_pool_reference(
|
||||
req_to_token_ref,
|
||||
req_pool_indices,
|
||||
pre_lens,
|
||||
seq_lens,
|
||||
extend_lens,
|
||||
out_cache_loc,
|
||||
)
|
||||
|
||||
# Compare results
|
||||
torch.testing.assert_close(req_to_token, req_to_token_ref)
|
||||
torch.testing.assert_close(req_to_token_opt, req_to_token_ref)
|
||||
|
||||
# Test case 2: batch size > 1
|
||||
batch_size = 3
|
||||
extend_lens_list = [14, 20, 30]
|
||||
total_extend_len = sum(extend_lens_list)
|
||||
|
||||
req_to_token = torch.zeros(
|
||||
(max_batch, max_context_len), dtype=torch.int32, device="cuda"
|
||||
)
|
||||
req_pool_indices = torch.tensor([42, 100, 200], dtype=torch.int32, device="cuda")
|
||||
pre_lens = torch.tensor([8, 10, 15], dtype=torch.int32, device="cuda")
|
||||
seq_lens = torch.tensor([22, 30, 45], dtype=torch.int32, device="cuda")
|
||||
extend_lens = torch.tensor(extend_lens_list, dtype=torch.int32, device="cuda")
|
||||
out_cache_loc = torch.arange(total_extend_len, dtype=torch.int32, device="cuda")
|
||||
|
||||
req_to_token_ref = req_to_token.clone()
|
||||
req_to_token_opt = req_to_token.clone()
|
||||
|
||||
# Run original triton kernel
|
||||
write_req_to_token_pool_triton[(batch_size,)](
|
||||
req_to_token,
|
||||
req_pool_indices,
|
||||
pre_lens,
|
||||
seq_lens,
|
||||
extend_lens,
|
||||
out_cache_loc,
|
||||
max_context_len,
|
||||
)
|
||||
|
||||
# Run optimized triton kernel
|
||||
max_extend_len = max(extend_lens_list)
|
||||
write_req_to_token_pool_triton_optimize[grid(batch_size, max_extend_len)](
|
||||
req_to_token_opt,
|
||||
req_pool_indices,
|
||||
pre_lens,
|
||||
seq_lens,
|
||||
extend_lens,
|
||||
out_cache_loc,
|
||||
max_context_len,
|
||||
BLOCK_SIZE=512,
|
||||
)
|
||||
|
||||
# Run reference implementation
|
||||
write_req_to_token_pool_reference(
|
||||
req_to_token_ref,
|
||||
req_pool_indices,
|
||||
pre_lens,
|
||||
seq_lens,
|
||||
extend_lens,
|
||||
out_cache_loc,
|
||||
)
|
||||
|
||||
# Compare results
|
||||
torch.testing.assert_close(req_to_token, req_to_token_ref)
|
||||
torch.testing.assert_close(req_to_token_opt, req_to_token_ref)
|
||||
|
||||
|
||||
def get_benchmark():
|
||||
batch_sizes = [1, 2, 4, 8, 16, 32, 64, 128]
|
||||
extend_lens = [32, 64, 128, 256, 512, 1024, 2048, 4096, 8192]
|
||||
configs = list(itertools.product(batch_sizes, extend_lens))
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["batch_size", "extend_len"],
|
||||
x_vals=configs,
|
||||
line_arg="provider",
|
||||
line_vals=["reference", "triton", "triton_optimize"],
|
||||
line_names=["PyTorch", "Triton", "Triton Optimized"],
|
||||
styles=[("blue", "-"), ("green", "-"), ("red", "-")],
|
||||
ylabel="us",
|
||||
plot_name="write-req-to-token-pool-performance",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(batch_size, extend_len, provider):
|
||||
max_batch = 256
|
||||
max_context_len = 16384
|
||||
|
||||
extend_lens_list = [extend_len] * batch_size
|
||||
total_extend_len = sum(extend_lens_list)
|
||||
|
||||
req_to_token = torch.zeros(
|
||||
(max_batch, max_context_len), dtype=torch.int32, device="cuda"
|
||||
)
|
||||
req_pool_indices = torch.arange(batch_size, dtype=torch.int32, device="cuda")
|
||||
pre_lens = torch.ones(batch_size, dtype=torch.int32, device="cuda") * 8
|
||||
seq_lens = pre_lens + extend_len
|
||||
extend_lens = torch.tensor(extend_lens_list, dtype=torch.int32, device="cuda")
|
||||
out_cache_loc = torch.arange(total_extend_len, dtype=torch.int32, device="cuda")
|
||||
|
||||
quantiles = [0.5, 0.2, 0.8]
|
||||
|
||||
if provider == "reference":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: write_req_to_token_pool_reference(
|
||||
req_to_token.clone(),
|
||||
req_pool_indices,
|
||||
pre_lens,
|
||||
seq_lens,
|
||||
extend_lens,
|
||||
out_cache_loc,
|
||||
),
|
||||
quantiles=tuple(quantiles),
|
||||
)
|
||||
elif provider == "triton":
|
||||
ms, min_ms, max_ms = run_bench(
|
||||
lambda: write_req_to_token_pool_triton[(batch_size,)](
|
||||
req_to_token.clone(),
|
||||
req_pool_indices,
|
||||
pre_lens,
|
||||
seq_lens,
|
||||
extend_lens,
|
||||
out_cache_loc,
|
||||
max_context_len,
|
||||
),
|
||||
quantiles=tuple(quantiles),
|
||||
)
|
||||
else:
|
||||
|
||||
def run_optimized():
|
||||
block_size = 128 if extend_len <= 1024 else 512
|
||||
grid_config = (batch_size, triton.cdiv(extend_len, block_size))
|
||||
write_req_to_token_pool_triton_optimize[grid_config](
|
||||
req_to_token.clone(),
|
||||
req_pool_indices,
|
||||
pre_lens,
|
||||
seq_lens,
|
||||
extend_lens,
|
||||
out_cache_loc,
|
||||
max_context_len,
|
||||
BLOCK_SIZE=block_size,
|
||||
)
|
||||
|
||||
ms, min_ms, max_ms = run_bench(run_optimized, quantiles=tuple(quantiles))
|
||||
|
||||
return 1000 * ms, 1000 * max_ms, 1000 * min_ms
|
||||
|
||||
return benchmark
|
||||
|
||||
|
||||
def run_benchmark(save_path: str = "./configs/benchmark_ops/write_req_to_token_pool/"):
|
||||
"""Run benchmark and save results"""
|
||||
|
||||
# Ensure save path exists
|
||||
os.makedirs(save_path, exist_ok=True)
|
||||
|
||||
# Run correctness test
|
||||
test_write_req_to_token_pool()
|
||||
print("Correctness test passed!")
|
||||
|
||||
# Run performance test
|
||||
benchmark = get_benchmark()
|
||||
benchmark.run(print_data=True, save_path=save_path)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--save_path",
|
||||
type=str,
|
||||
default="./configs/benchmark_ops/write_req_to_token_pool/",
|
||||
help="Path to save benchmark results",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
run_benchmark(args.save_path)
|
||||
294
third_party/sglang/benchmark/kernels/sliding_window_attention_triton/bench_triton_swa_kernel.py
vendored
Normal file
294
third_party/sglang/benchmark/kernels/sliding_window_attention_triton/bench_triton_swa_kernel.py
vendored
Normal file
@@ -0,0 +1,294 @@
|
||||
import itertools
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import triton.testing as tt
|
||||
|
||||
from sglang.benchmark.bench_utils import run_bench
|
||||
from sglang.srt.layers.attention.triton_ops.extend_attention import extend_attention_fwd
|
||||
|
||||
|
||||
def extend_attention_fwd_torch(
|
||||
q: torch.Tensor, # [extend_tokens, H_Q, D]
|
||||
k: torch.Tensor, # [extend_tokens, H_KV, D]
|
||||
v: torch.Tensor, # [extend_tokens, H_KV, D]
|
||||
o: torch.Tensor, # [extend_tokens, H_Q, D]
|
||||
k_cache: torch.Tensor, # [total_tokens, H_KV, D]
|
||||
v_cache: torch.Tensor, # [total_tokens, H_KV, D]
|
||||
qo_indptr: torch.Tensor, # [B+1]
|
||||
kv_indptr: torch.Tensor, # [B+1]
|
||||
kv_indices: torch.Tensor, # [prefix_tokens]
|
||||
sliding_window_size: int,
|
||||
):
|
||||
B = qo_indptr.size(0) - 1
|
||||
_, H_Q, D = q.shape
|
||||
_, H_KV, _ = k.shape
|
||||
|
||||
group_size = H_Q // H_KV
|
||||
scale = 1.0 / D**0.5
|
||||
|
||||
for i in range(B):
|
||||
q_start = int(qo_indptr[i].item())
|
||||
q_end = int(qo_indptr[i + 1].item())
|
||||
kv_start = int(kv_indptr[i].item())
|
||||
kv_end = int(kv_indptr[i + 1].item())
|
||||
|
||||
prefix_indices = kv_indices[kv_start:kv_end]
|
||||
k_prefix = k_cache[prefix_indices] # [prefix_len, H_KV, D]
|
||||
v_prefix = v_cache[prefix_indices] # [prefix_len, H_KV, D]
|
||||
|
||||
k_extend = k[q_start:q_end] # [extend_len, H_KV, D]
|
||||
v_extend = v[q_start:q_end] # [extend_len, H_KV, D]
|
||||
q_extend = q[q_start:q_end] # [extend_len, H_Q, D]
|
||||
|
||||
k_full = torch.cat([k_prefix, k_extend], dim=0) # [total_len, H_KV, D]
|
||||
v_full = torch.cat([v_prefix, v_extend], dim=0) # [total_len, H_KV, D]
|
||||
|
||||
if group_size != 1:
|
||||
k_full_hq = k_full.repeat_interleave(
|
||||
group_size, dim=1
|
||||
) # [total_len, H_Q, D]
|
||||
v_full_hq = v_full.repeat_interleave(
|
||||
group_size, dim=1
|
||||
) # [total_len, H_Q, D]
|
||||
else:
|
||||
k_full_hq = k_full
|
||||
v_full_hq = v_full
|
||||
|
||||
prefix_len = k_prefix.size(0)
|
||||
extend_len = k_extend.size(0)
|
||||
total_len = prefix_len + extend_len
|
||||
|
||||
# causal
|
||||
pos_keys = torch.arange(total_len, device=q.device)
|
||||
t = prefix_len + torch.arange(extend_len, device=q.device) # [extend_len]
|
||||
causal_mask = pos_keys.unsqueeze(0) <= t.unsqueeze(1)
|
||||
|
||||
# sliding window
|
||||
if sliding_window_size is not None and sliding_window_size > 0:
|
||||
start = (t - (sliding_window_size)).clamp_min(0) # [extend_len]
|
||||
else:
|
||||
start = torch.zeros_like(t)
|
||||
window_mask = pos_keys.unsqueeze(0) >= start.unsqueeze(1)
|
||||
|
||||
final_mask = causal_mask & window_mask
|
||||
|
||||
attn_scores = (
|
||||
torch.einsum("qhd,khd->qhk", q_extend, k_full_hq) * scale
|
||||
) # [extend_len, H_Q, total_len]
|
||||
attn_scores = attn_scores.masked_fill(~final_mask.unsqueeze(1), float("-inf"))
|
||||
|
||||
attn_weights = F.softmax(attn_scores, dim=-1)
|
||||
o[q_start:q_end] = torch.einsum("qhk,khd->qhd", attn_weights, v_full_hq)
|
||||
|
||||
|
||||
def _build_batch(
|
||||
B, N_CTX, H_Q, H_KV, D, WINDOW_SIZE, dtype=torch.bfloat16, device="cuda"
|
||||
):
|
||||
b_seq_len_prefix = torch.randint(
|
||||
1, max(2, N_CTX // 2), (B,), dtype=torch.int32, device=device
|
||||
)
|
||||
b_seq_len_extend = torch.randint(
|
||||
1, max(2, N_CTX // 2), (B,), dtype=torch.int32, device=device
|
||||
)
|
||||
b_seq_len = b_seq_len_prefix + b_seq_len_extend
|
||||
|
||||
b_start_loc = torch.zeros((B,), dtype=torch.int32, device=device)
|
||||
b_start_loc[1:] = torch.cumsum(b_seq_len[:-1], 0)
|
||||
b_start_loc_extend = torch.zeros((B,), dtype=torch.int32, device=device)
|
||||
b_start_loc_extend[1:] = torch.cumsum(b_seq_len_extend[:-1], 0)
|
||||
|
||||
kv_indptr = torch.zeros((B + 1,), dtype=torch.int32, device=device)
|
||||
kv_indptr[1 : B + 1] = torch.cumsum(b_seq_len_prefix[:B], dim=0)
|
||||
|
||||
kv_indices = torch.zeros(
|
||||
(int(b_seq_len_prefix.sum().item()),), dtype=torch.int32, device=device
|
||||
)
|
||||
for i in range(B):
|
||||
s = kv_indptr[i].item()
|
||||
e = kv_indptr[i + 1].item()
|
||||
kv_indices[s:e] = torch.arange(
|
||||
b_start_loc[i],
|
||||
b_start_loc[i] + b_seq_len_prefix[i],
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
)
|
||||
|
||||
total_token_num = int(torch.sum(b_seq_len).item())
|
||||
extend_token_num = int(torch.sum(b_seq_len_extend).item())
|
||||
|
||||
k_buffer = torch.empty(
|
||||
(total_token_num, H_KV, D), dtype=dtype, device=device
|
||||
).normal_(mean=0.1, std=0.2)
|
||||
v_buffer = torch.empty(
|
||||
(total_token_num, H_KV, D), dtype=dtype, device=device
|
||||
).normal_(mean=0.1, std=0.2)
|
||||
|
||||
k_extend = torch.empty((extend_token_num, H_KV, D), dtype=dtype, device=device)
|
||||
v_extend = torch.empty((extend_token_num, H_KV, D), dtype=dtype, device=device)
|
||||
q_extend = torch.empty((extend_token_num, H_Q, D), dtype=dtype, device=device)
|
||||
|
||||
for i in range(B):
|
||||
extend_start_in_buffer = b_start_loc[i] + b_seq_len_prefix[i]
|
||||
extend_end_in_buffer = b_start_loc[i] + b_seq_len[i]
|
||||
extend_start = b_start_loc_extend[i]
|
||||
extend_end = b_start_loc_extend[i] + b_seq_len_extend[i]
|
||||
|
||||
k_extend[extend_start:extend_end] = k_buffer[
|
||||
extend_start_in_buffer:extend_end_in_buffer
|
||||
]
|
||||
v_extend[extend_start:extend_end] = v_buffer[
|
||||
extend_start_in_buffer:extend_end_in_buffer
|
||||
]
|
||||
q_extend[extend_start:extend_end] = torch.empty(
|
||||
(int(b_seq_len_extend[i].item()), H_Q, D), dtype=dtype, device=device
|
||||
).normal_(mean=0.1, std=0.2)
|
||||
|
||||
o_extend_triton = torch.empty(
|
||||
(extend_token_num, H_Q, D), dtype=dtype, device=device
|
||||
)
|
||||
o_extend_torch = torch.empty((extend_token_num, H_Q, D), dtype=dtype, device=device)
|
||||
|
||||
b_seq_len_extend = b_seq_len - b_seq_len_prefix
|
||||
max_len_extend = int(torch.max(b_seq_len_extend, 0)[0].item())
|
||||
qo_indptr = torch.zeros((B + 1,), dtype=torch.int32, device=device)
|
||||
qo_indptr[1 : B + 1] = torch.cumsum(b_seq_len_extend[:B], dim=0)
|
||||
|
||||
inputs = dict(
|
||||
q_extend=q_extend,
|
||||
k_extend=k_extend,
|
||||
v_extend=v_extend,
|
||||
k_buffer=k_buffer,
|
||||
v_buffer=v_buffer,
|
||||
o_extend_triton=o_extend_triton,
|
||||
o_extend_torch=o_extend_torch,
|
||||
qo_indptr=qo_indptr,
|
||||
kv_indptr=kv_indptr,
|
||||
kv_indices=kv_indices,
|
||||
max_len_extend=max_len_extend,
|
||||
WINDOW_SIZE=WINDOW_SIZE,
|
||||
)
|
||||
meta = dict(
|
||||
B=B, N_CTX=N_CTX, H_Q=H_Q, H_KV=H_KV, D=D, extend_token_num=extend_token_num
|
||||
)
|
||||
return inputs, meta
|
||||
|
||||
|
||||
def _run_triton(inputs):
|
||||
extend_attention_fwd(
|
||||
inputs["q_extend"],
|
||||
inputs["k_extend"],
|
||||
inputs["v_extend"],
|
||||
inputs["o_extend_triton"],
|
||||
inputs["k_buffer"],
|
||||
inputs["v_buffer"],
|
||||
inputs["qo_indptr"],
|
||||
inputs["kv_indptr"],
|
||||
inputs["kv_indices"],
|
||||
custom_mask=None,
|
||||
is_causal=True,
|
||||
mask_indptr=None,
|
||||
max_len_extend=inputs["max_len_extend"],
|
||||
sliding_window_size=inputs["WINDOW_SIZE"],
|
||||
)
|
||||
|
||||
|
||||
def _run_torch_ref(inputs):
|
||||
extend_attention_fwd_torch(
|
||||
inputs["q_extend"],
|
||||
inputs["k_extend"],
|
||||
inputs["v_extend"],
|
||||
inputs["o_extend_torch"],
|
||||
inputs["k_buffer"],
|
||||
inputs["v_buffer"],
|
||||
inputs["qo_indptr"],
|
||||
inputs["kv_indptr"],
|
||||
inputs["kv_indices"],
|
||||
inputs["WINDOW_SIZE"],
|
||||
)
|
||||
|
||||
|
||||
N_CTXS = [1024, 2048, 4096, 8192]
|
||||
WINDOW_SIZES = [-1, 127, 256, 512]
|
||||
|
||||
CONFIGS = list(itertools.product(N_CTXS, WINDOW_SIZES))
|
||||
|
||||
PROVIDERS = ["torch", "triton"]
|
||||
|
||||
|
||||
@tt.perf_report(
|
||||
tt.Benchmark(
|
||||
x_names=["N_CTX", "WINDOW_SIZE"],
|
||||
x_vals=CONFIGS,
|
||||
line_arg="provider",
|
||||
line_vals=PROVIDERS,
|
||||
line_names=PROVIDERS,
|
||||
ylabel="Runtime (ms)",
|
||||
plot_name="extend_attention_triton_vs_torch",
|
||||
args={
|
||||
"B": 32,
|
||||
"H_Q": 64,
|
||||
"H_KV": 8,
|
||||
"D": 128,
|
||||
"dtype": "bf16",
|
||||
"device": "cuda",
|
||||
"check_correctness": False,
|
||||
"warmup": 25,
|
||||
"rep": 100,
|
||||
},
|
||||
)
|
||||
)
|
||||
def bench(
|
||||
N_CTX,
|
||||
provider,
|
||||
B,
|
||||
H_Q,
|
||||
H_KV,
|
||||
D,
|
||||
dtype,
|
||||
device,
|
||||
WINDOW_SIZE,
|
||||
check_correctness,
|
||||
warmup,
|
||||
rep,
|
||||
):
|
||||
torch.manual_seed(0)
|
||||
torch.cuda.manual_seed(0)
|
||||
dtype_map = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}
|
||||
dt = dtype_map[dtype]
|
||||
|
||||
inputs, _ = _build_batch(
|
||||
B, N_CTX, H_Q, H_KV, D, WINDOW_SIZE, dtype=dt, device=device
|
||||
)
|
||||
|
||||
if check_correctness and provider == "triton":
|
||||
_run_triton(inputs)
|
||||
_run_torch_ref(inputs)
|
||||
torch.cuda.synchronize()
|
||||
if not torch.allclose(
|
||||
inputs["o_extend_triton"], inputs["o_extend_torch"], rtol=1e-3, atol=1e-3
|
||||
):
|
||||
raise AssertionError("Mismatch between triton and torch reference.")
|
||||
|
||||
if provider == "triton":
|
||||
ms = run_bench(
|
||||
lambda: _run_triton(inputs),
|
||||
quantiles=None,
|
||||
warmup_ms=warmup,
|
||||
rep_ms=rep,
|
||||
)[0]
|
||||
elif provider == "torch":
|
||||
ms = run_bench(
|
||||
lambda: _run_torch_ref(inputs),
|
||||
quantiles=None,
|
||||
warmup_ms=warmup,
|
||||
rep_ms=rep,
|
||||
)[0]
|
||||
else:
|
||||
raise ValueError(provider)
|
||||
|
||||
return ms
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
bench.run(print_data=True, show_plots=False)
|
||||
Reference in New Issue
Block a user