From aab07de35f7b8f9b460c7dadfff8c2cc76fd7b8f Mon Sep 17 00:00:00 2001 From: papertager <2567587994@qq.com> Date: Sat, 6 Jun 2026 11:05:16 +0800 Subject: [PATCH] feat: write Fused MoE benchmark reports --- .../python/benchmark_fused_moe_i8_tn.py | 101 +++++++++++++++++- 1 file changed, 100 insertions(+), 1 deletion(-) diff --git a/基于AI Agent开发范式的国产GPU大模型推理算子库优化/baselines/fused_moe/standalone/fused_moe_i8_tn/python/benchmark_fused_moe_i8_tn.py b/基于AI Agent开发范式的国产GPU大模型推理算子库优化/baselines/fused_moe/standalone/fused_moe_i8_tn/python/benchmark_fused_moe_i8_tn.py index 587c4f2..ee2c11e 100644 --- a/基于AI Agent开发范式的国产GPU大模型推理算子库优化/baselines/fused_moe/standalone/fused_moe_i8_tn/python/benchmark_fused_moe_i8_tn.py +++ b/基于AI Agent开发范式的国产GPU大模型推理算子库优化/baselines/fused_moe/standalone/fused_moe_i8_tn/python/benchmark_fused_moe_i8_tn.py @@ -1,5 +1,10 @@ import argparse +import csv +import json +import platform import time +from datetime import datetime +from pathlib import Path from test_fused_moe_i8_tn_pybind import fill_inputs, resolve_backends @@ -8,6 +13,41 @@ K_N = 128 K_K = 128 +def get_timestamp() -> str: + return datetime.now().strftime("%Y%m%d_%H%M%S") + + +def synchronize_device(): + try: + import torch + except ImportError: + return + + if torch.cuda.is_available(): + torch.cuda.synchronize() + + +def collect_environment() -> dict[str, object]: + env = { + "python": platform.python_version(), + "platform": platform.platform(), + } + try: + import torch + + env["torch"] = torch.__version__ + env["cuda_available"] = torch.cuda.is_available() + if torch.cuda.is_available(): + device = torch.cuda.current_device() + props = torch.cuda.get_device_properties(device) + env["device_name"] = props.name + env["device_memory_gib"] = round(props.total_memory / 1024**3, 2) + except ImportError: + env["torch"] = "unavailable" + env["cuda_available"] = False + return env + + def compute_tops(rows: int, cols: int, k_dim: int, avg_ms: float) -> float: if avg_ms <= 0.0: return 0.0 @@ -21,9 +61,11 @@ def benchmark_backend_case(backend: str, backend_fn, tag: str, num_tokens: int, for _ in range(warmup): backend_fn(a, b, scale_a, scale_b, moe_weights, token_ids, expert_ids, topk) + synchronize_device() start = time.perf_counter() for _ in range(iters): backend_fn(a, b, scale_a, scale_b, moe_weights, token_ids, expert_ids, topk) + synchronize_device() elapsed_s = time.perf_counter() - start avg_ms = elapsed_s * 1000.0 / iters @@ -32,6 +74,42 @@ def benchmark_backend_case(backend: str, backend_fn, tag: str, num_tokens: int, f"{backend}:{tag} benchmark: avg_ms={avg_ms:.6f}, " f"TOPS={tops:.6f}, warmup={warmup}, iters={iters}" ) + return { + "backend": backend, + "case": tag, + "num_tokens": num_tokens, + "topk": topk, + "rows": em, + "cols": K_N, + "k_dim": K_K, + "warmup": warmup, + "iters": iters, + "avg_ms": avg_ms, + "tops": tops, + } + + +def write_reports(records: list[dict[str, object]], output_dir: str, prefix: str) -> tuple[Path, Path]: + output_path = Path(output_dir) + output_path.mkdir(parents=True, exist_ok=True) + stem = f"{prefix}_{get_timestamp()}" + csv_path = output_path / f"{stem}.csv" + json_path = output_path / f"{stem}.json" + + fieldnames = sorted({key for record in records for key in record.keys()}) + with csv_path.open("w", newline="") as f: + writer = csv.DictWriter(f, fieldnames=fieldnames) + writer.writeheader() + writer.writerows(records) + + payload = { + "environment": collect_environment(), + "records": records, + } + with json_path.open("w") as f: + json.dump(payload, f, indent=2, sort_keys=True) + + return csv_path, json_path def parse_args(): @@ -43,6 +121,16 @@ def parse_args(): ) parser.add_argument("--warmup", type=int, default=5) parser.add_argument("--iters", type=int, default=20) + parser.add_argument( + "--output-dir", + default=".", + help="Directory used for CSV and JSON benchmark reports.", + ) + parser.add_argument( + "--case-filter", + default="", + help="Run only cases whose tag contains this substring.", + ) return parser.parse_args() @@ -53,10 +141,21 @@ def main(): ("fused_moe_i8_tn_topk2", 256, 2, 512, [0, 1, 1, 0]), ("fused_moe_i8_tn_topk3", 128, 3, 384, [0, 1, 0]), ] + if args.case_filter: + cases = [case for case in cases if args.case_filter in case[0]] + if not cases: + raise ValueError(f"no benchmark cases match --case-filter={args.case_filter!r}") backends = resolve_backends(args.backend) + records = [] for backend, backend_fn in backends: for case in cases: - benchmark_backend_case(backend, backend_fn, *case, warmup=args.warmup, iters=args.iters) + records.append( + benchmark_backend_case(backend, backend_fn, *case, warmup=args.warmup, iters=args.iters) + ) + + csv_path, json_path = write_reports(records, args.output_dir, "fused_moe_i8_tn_benchmark") + print(f"CSV report saved to {csv_path}") + print(f"JSON report saved to {json_path}") if __name__ == "__main__": -- 2.34.1