-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathverify.py
More file actions
executable file
·182 lines (149 loc) · 7.18 KB
/
Copy pathverify.py
File metadata and controls
executable file
·182 lines (149 loc) · 7.18 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
#!/usr/bin/env python3
"""Pack and verify released FlashInfer contest submissions."""
from __future__ import annotations
import argparse
import importlib.util
import os
import sys
from dataclasses import dataclass
from pathlib import Path
ROOT = Path(__file__).resolve().parent
DEFAULT_DATASET = ROOT / "data" / "flashinfer-trace"
@dataclass(frozen=True)
class Submission:
name: str
path: Path
SUBMISSIONS = {
"moe-fp8": Submission("moe-fp8", ROOT / "submissions" / "moe-fp8"),
"gdn-prefill": Submission("gdn-prefill", ROOT / "submissions" / "gdn-prefill"),
"gdn-decode": Submission("gdn-decode", ROOT / "submissions" / "gdn-decode"),
"dsa-sparse-attention": Submission("dsa-sparse-attention", ROOT / "submissions" / "dsa-sparse-attention"),
"dsa-topk-indexer": Submission("dsa-topk-indexer", ROOT / "submissions" / "dsa-topk-indexer"),
}
def dataset_path(explicit: str | None) -> Path:
raw = explicit or os.environ.get("FIB_DATASET_PATH")
return Path(raw).expanduser().resolve() if raw else DEFAULT_DATASET
def import_pack_solution(submission: Submission):
script = submission.path / "scripts" / "pack_solution.py"
if not script.exists():
raise FileNotFoundError(f"Missing pack script: {script}")
module_name = f"_submission_pack_{submission.name.replace('-', '_')}"
spec = importlib.util.spec_from_file_location(module_name, script)
if spec is None or spec.loader is None:
raise RuntimeError(f"Could not import {script}")
module = importlib.util.module_from_spec(spec)
sys.modules[module_name] = module
spec.loader.exec_module(module)
return module.pack_solution
def import_flashinfer_bench():
try:
from flashinfer_bench import Benchmark, BenchmarkConfig, Solution, TraceSet
except ModuleNotFoundError as exc:
raise ModuleNotFoundError(
"flashinfer_bench is not installed. Run `uv sync`, or install the official "
"FlashInfer benchmark package following docs/reproduction.md."
) from exc
return Benchmark, BenchmarkConfig, Solution, TraceSet
def pack_submission(submission: Submission, output_dir: Path) -> Path:
output_dir.mkdir(parents=True, exist_ok=True)
output_path = output_dir / f"{submission.name}.solution.json"
pack_solution = import_pack_solution(submission)
return Path(pack_solution(output_path))
def select_workloads(workloads: list, limit: int | None) -> list:
if limit is None:
return workloads
return workloads[: max(0, limit)]
def run_one(submission: Submission, dataset: Path, args: argparse.Namespace) -> bool:
Benchmark, BenchmarkConfig, Solution, TraceSet = import_flashinfer_bench()
print(f"\n== {submission.name} ==")
solution_json = pack_submission(submission, ROOT / "outputs" / "packed")
solution = Solution.model_validate_json(solution_json.read_text())
trace_set = TraceSet.from_path(str(dataset))
if solution.definition not in trace_set.definitions:
raise ValueError(f"Definition {solution.definition!r} not found in {dataset}")
definition = trace_set.definitions[solution.definition]
workloads = list(trace_set.workloads.get(solution.definition, []))
workloads = select_workloads(workloads, args.limit)
if not workloads:
raise ValueError(f"No workloads found for {solution.definition}")
config_kwargs = dict(
warmup_runs=args.warmup_runs,
iterations=args.iterations,
num_trials=args.num_trials,
use_isolated_runner=args.isolated_runner,
)
if solution.definition.startswith("moe_fp8"):
config_kwargs.update(atol=1.0, rtol=0.3, required_matched_ratio=0.9)
config = BenchmarkConfig(**config_kwargs)
bench_trace_set = TraceSet(
root=trace_set.root,
definitions={definition.name: definition},
solutions={definition.name: [solution]},
workloads={definition.name: workloads},
traces={definition.name: []},
)
result_trace_set = Benchmark(bench_trace_set, config).run_all(dump_traces=False)
traces = list(result_trace_set.traces.get(definition.name, []))
passed = 0
total = len(workloads)
latencies = []
speedups = []
for trace in traces:
evaluation = getattr(trace, "evaluation", None)
if evaluation is None:
continue
status = getattr(evaluation.status, "value", str(evaluation.status))
if status.lower() == "passed":
passed += 1
perf = getattr(evaluation, "performance", None)
if perf is not None:
if perf.latency_ms is not None:
latencies.append(float(perf.latency_ms))
if perf.speedup_factor is not None:
speedups.append(float(perf.speedup_factor))
mean_latency = sum(latencies) / len(latencies) if latencies else None
mean_speedup = sum(speedups) / len(speedups) if speedups else None
print(f"definition: {solution.definition}")
print(f"passed: {passed}/{total}")
if mean_latency is not None:
print(f"latency: {mean_latency:.6f} ms mean")
if mean_speedup is not None:
print(f"speedup: {mean_speedup:.4f}x mean")
return passed == total
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
group = parser.add_mutually_exclusive_group(required=True)
group.add_argument("--all", action="store_true", help="Verify all released submissions.")
group.add_argument("--submission", choices=sorted(SUBMISSIONS), help="Verify one released submission.")
parser.add_argument("--dataset", default=None, help="Path to flashinfer-trace dataset.")
parser.add_argument("--limit", type=int, default=None, help="Run only the first N workloads for each definition.")
parser.add_argument("--fast", action="store_true", help="Use a short benchmark configuration.")
parser.add_argument("--isolated-runner", action="store_true", help="Use flashinfer-bench isolated runner.")
parser.add_argument("--warmup-runs", type=int, default=None)
parser.add_argument("--iterations", type=int, default=None)
parser.add_argument("--num-trials", type=int, default=None)
return parser.parse_args()
def main() -> int:
args = parse_args()
if args.fast:
args.warmup_runs = 1 if args.warmup_runs is None else args.warmup_runs
args.iterations = 5 if args.iterations is None else args.iterations
args.num_trials = 1 if args.num_trials is None else args.num_trials
args.limit = 2 if args.limit is None else args.limit
else:
args.warmup_runs = 3 if args.warmup_runs is None else args.warmup_runs
args.iterations = 50 if args.iterations is None else args.iterations
args.num_trials = 3 if args.num_trials is None else args.num_trials
dataset = dataset_path(args.dataset)
if not dataset.exists():
raise FileNotFoundError(
f"Dataset not found: {dataset}\n"
"Run ./scripts/download_data.sh or set FIB_DATASET_PATH."
)
selected = list(SUBMISSIONS.values()) if args.all else [SUBMISSIONS[args.submission]]
ok = True
for submission in selected:
ok = run_one(submission, dataset, args) and ok
return 0 if ok else 1
if __name__ == "__main__":
raise SystemExit(main())