Files
2026-07-29 13:36:49 +08:00

455 lines
13 KiB
Python

#!/usr/bin/env python3
"""Run auditable FlashKDA correctness and latency probes on one CUDA GPU.
The upstream torch reference is imported from a pinned FlashKDA checkout. The
script deliberately uses valid synthetic A_log[H] tensors: it does not resolve
the public Kimi K3 checkpoint's A_log[128] versus H=96 inconsistency.
"""
from __future__ import annotations
import argparse
import hashlib
import importlib.metadata
import json
import math
import platform
import statistics
import subprocess
import sys
import time
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Callable
import torch
import torch.nn.functional as F
import flash_kda
import flash_kda_C
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--flashkda-source", type=Path, required=True)
parser.add_argument("--wheel", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--warmup", type=int, default=20)
parser.add_argument("--iters", type=int, default=100)
parser.add_argument("--repeats", type=int, default=3)
parser.add_argument("--seed", type=int, default=20260729)
parser.add_argument("--captured-at", default=None)
return parser.parse_args()
def sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for block in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def percentile(sorted_values: list[float], fraction: float) -> float:
if not sorted_values:
return float("nan")
index = fraction * (len(sorted_values) - 1)
lower = math.floor(index)
upper = math.ceil(index)
if lower == upper:
return sorted_values[lower]
weight = index - lower
return sorted_values[lower] * (1 - weight) + sorted_values[upper] * weight
def nvidia_smi() -> dict[str, Any]:
fields = [
"name",
"driver_version",
"memory.total",
"power.limit",
"clocks.max.sm",
"clocks.max.memory",
]
output = subprocess.check_output(
[
"nvidia-smi",
f"--query-gpu={','.join(fields)}",
"--format=csv,noheader,nounits",
],
text=True,
).strip()
values = [value.strip() for value in output.split(",")]
return dict(zip(fields, values, strict=True))
def make_inputs(
sequence_lengths: list[int],
heads: int,
seed: int,
) -> dict[str, torch.Tensor | float | None]:
batch = 1
dimension = 128
total_tokens = sum(sequence_lengths)
sequences = len(sequence_lengths)
generator = torch.Generator(device="cuda").manual_seed(seed)
q = F.normalize(
torch.randn(
(batch, total_tokens, heads, dimension),
dtype=torch.float32,
device="cuda",
generator=generator,
),
p=2,
dim=-1,
).to(torch.bfloat16)
k = F.normalize(
torch.randn(
(batch, total_tokens, heads, dimension),
dtype=torch.float32,
device="cuda",
generator=generator,
),
p=2,
dim=-1,
).to(torch.bfloat16)
v = torch.randn(
(batch, total_tokens, heads, dimension),
dtype=torch.bfloat16,
device="cuda",
generator=generator,
)
g = torch.randn(
(batch, total_tokens, heads, dimension),
dtype=torch.bfloat16,
device="cuda",
generator=generator,
)
beta = torch.randn(
(batch, total_tokens, heads),
dtype=torch.bfloat16,
device="cuda",
generator=generator,
)
a_log = torch.rand(
heads, dtype=torch.float32, device="cuda", generator=generator
)
dt_bias = torch.rand(
(heads, dimension),
dtype=torch.float32,
device="cuda",
generator=generator,
)
initial_state = torch.randn(
(sequences, heads, dimension, dimension),
dtype=torch.bfloat16,
device="cuda",
generator=generator,
)
cu_seqlens = None
if len(sequence_lengths) > 1:
offsets = [0]
for length in sequence_lengths:
offsets.append(offsets[-1] + length)
cu_seqlens = torch.tensor(offsets, dtype=torch.long, device="cuda")
return {
"q": q,
"k": k,
"v": v,
"g": g,
"beta": beta,
"A_log": a_log,
"dt_bias": dt_bias,
"initial_state": initial_state,
"cu_seqlens": cu_seqlens,
"scale": 1 / math.sqrt(dimension),
}
def run_correctness_case(
torch_ref: Callable[..., None],
name: str,
sequence_lengths: list[int],
heads: int,
seed: int,
) -> dict[str, Any]:
inputs = make_inputs(sequence_lengths, heads, seed)
q = inputs["q"]
assert isinstance(q, torch.Tensor)
initial_state = inputs["initial_state"]
assert isinstance(initial_state, torch.Tensor)
out_kernel = torch.zeros_like(q)
out_reference = torch.zeros_like(q)
state_kernel = torch.zeros_like(initial_state)
state_reference = torch.zeros_like(initial_state)
common = {
"A_log": inputs["A_log"],
"dt_bias": inputs["dt_bias"],
"lower_bound": -5.0,
"initial_state": initial_state.clone(),
"cu_seqlens": inputs["cu_seqlens"],
}
start = time.perf_counter()
flash_kda.fwd(
inputs["q"],
inputs["k"],
inputs["v"],
inputs["g"],
inputs["beta"],
inputs["scale"],
out_kernel,
final_state=state_kernel,
**common,
)
torch.cuda.synchronize()
kernel_ms = (time.perf_counter() - start) * 1000
start = time.perf_counter()
torch_ref(
inputs["q"],
inputs["k"],
inputs["v"],
inputs["g"],
inputs["beta"],
inputs["scale"],
out_reference,
final_state=state_reference,
**common,
)
torch.cuda.synchronize()
reference_ms = (time.perf_counter() - start) * 1000
output_diff = (out_kernel.float() - out_reference.float()).abs()
state_diff = (state_kernel.float() - state_reference.float()).abs()
result = {
"name": name,
"sequence_lengths": sequence_lengths,
"heads": heads,
"dimension": 128,
"kernel_ms_first_measured": kernel_ms,
"reference_ms": reference_ms,
"output_exact": bool(torch.equal(out_kernel, out_reference)),
"state_exact": bool(torch.equal(state_kernel, state_reference)),
"output_max_abs_diff": output_diff.max().item(),
"output_mean_abs_diff": output_diff.mean().item(),
"state_max_abs_diff": state_diff.max().item(),
"state_mean_abs_diff": state_diff.mean().item(),
}
if not result["output_exact"] or not result["state_exact"]:
raise AssertionError(f"exact correctness failed: {result}")
return result
def bench_events(
function: Callable[[], None],
warmup: int,
iters: int,
repeats: int,
) -> list[float]:
for _ in range(max(warmup, 1)):
function()
torch.cuda.synchronize()
elapsed: list[float] = []
for _ in range(repeats):
starts = [torch.cuda.Event(enable_timing=True) for _ in range(iters)]
ends = [torch.cuda.Event(enable_timing=True) for _ in range(iters)]
torch.cuda.synchronize()
for index in range(iters):
starts[index].record()
function()
ends[index].record()
torch.cuda.synchronize()
elapsed.extend(
start.elapsed_time(end) for start, end in zip(starts, ends, strict=True)
)
return elapsed
def summarize_latency(
values: list[float],
total_tokens: int,
) -> dict[str, Any]:
ordered = sorted(float(value) for value in values)
mean = statistics.fmean(ordered)
return {
"samples": len(ordered),
"mean_ms": mean,
"min_ms": ordered[0],
"p50_ms": percentile(ordered, 0.50),
"p95_ms": percentile(ordered, 0.95),
"max_ms": ordered[-1],
"sequence_tokens_per_second": total_tokens / (mean / 1000),
}
def run_benchmark_case(
name: str,
sequence_lengths: list[int],
heads: int,
seed: int,
warmup: int,
iters: int,
repeats: int,
) -> dict[str, Any]:
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
inputs = make_inputs(sequence_lengths, heads, seed)
q = inputs["q"]
initial_state = inputs["initial_state"]
assert isinstance(q, torch.Tensor)
assert isinstance(initial_state, torch.Tensor)
out = torch.zeros_like(q)
state_bf16 = torch.zeros_like(initial_state)
initial_fp32 = initial_state.float()
state_fp32 = torch.zeros_like(initial_fp32)
common = {
"A_log": inputs["A_log"],
"dt_bias": inputs["dt_bias"],
"lower_bound": -5.0,
"cu_seqlens": inputs["cu_seqlens"],
}
def invoke(initial: torch.Tensor | None, final: torch.Tensor | None) -> None:
flash_kda.fwd(
inputs["q"],
inputs["k"],
inputs["v"],
inputs["g"],
inputs["beta"],
inputs["scale"],
out,
initial_state=initial,
final_state=final,
**common,
)
variants = {
"bf16_state": lambda: invoke(initial_state, state_bf16),
"no_state": lambda: invoke(None, None),
"fp32_state": lambda: invoke(initial_fp32, state_fp32),
}
timings = {
variant: summarize_latency(
bench_events(function, warmup, iters, repeats),
sum(sequence_lengths),
)
for variant, function in variants.items()
}
torch.cuda.synchronize()
return {
"name": name,
"sequence_lengths": sequence_lengths,
"total_tokens": sum(sequence_lengths),
"heads": heads,
"dimension": 128,
"warmup": warmup,
"iters": iters,
"repeats": repeats,
"timings": timings,
"output_abs_mean_after_last_run": out.float().abs().mean().item(),
"peak_allocated_mib": torch.cuda.max_memory_allocated() / 2**20,
"peak_reserved_mib": torch.cuda.max_memory_reserved() / 2**20,
}
def main() -> None:
args = parse_args()
if not torch.cuda.is_available():
raise RuntimeError("CUDA GPU is required")
if torch.cuda.get_device_capability(0) < (9, 0):
raise RuntimeError("FlashKDA requires SM90 or newer")
tests_dir = args.flashkda_source / "tests"
sys.path.insert(0, str(tests_dir))
from torch_ref import torch_ref
revision = subprocess.check_output(
["git", "-C", str(args.flashkda_source), "rev-parse", "HEAD"],
text=True,
).strip()
torch.manual_seed(args.seed)
torch.cuda.manual_seed_all(args.seed)
correctness_specs = [
("one_chunk", [16], 1),
("partial_tail", [17], 1),
("two_chunks", [32], 1),
("multi_chunk_multi_head", [65], 2),
("k3_head_count", [17], 96),
("varlen_partial_chunks", [17, 31], 2),
]
correctness = [
run_correctness_case(torch_ref, name, lengths, heads, args.seed + index)
for index, (name, lengths, heads) in enumerate(correctness_specs)
]
benchmark_specs = [
("teaching_scale", [512], 8),
("intermediate", [2048], 32),
("k3_fixed_shape", [8192], 96),
("k3_varlen_shape", [1300, 547, 2048, 963, 271, 3063], 96),
]
benchmarks = [
run_benchmark_case(
name,
lengths,
heads,
args.seed + 100 + index,
args.warmup,
args.iters,
args.repeats,
)
for index, (name, lengths, heads) in enumerate(benchmark_specs)
]
captured_at = args.captured_at or datetime.now(timezone.utc).isoformat()
result = {
"schema_version": 1,
"captured_at": captured_at,
"evidence_identity": "X / local execution on deterministic synthetic tensors",
"boundary": {
"k3_checkpoint_loaded": False,
"real_token_hidden_states": False,
"a_log_shape_conflict_resolved": False,
"benchmark_comparison": "local FlashKDA timings only; author H20/GB200 tables remain separate",
},
"provenance": {
"flashkda_revision": revision,
"flashkda_package": importlib.metadata.version("flash-kda"),
"wheel_filename": args.wheel.name,
"wheel_sha256": sha256(args.wheel),
"runner": str(Path(__file__).relative_to(Path.cwd())),
},
"environment": {
"python": platform.python_version(),
"platform": platform.platform(),
"libc": list(platform.libc_ver()),
"torch": torch.__version__,
"torch_cuda": torch.version.cuda,
"gpu": torch.cuda.get_device_name(0),
"capability": list(torch.cuda.get_device_capability(0)),
"nvidia_smi": nvidia_smi(),
"flash_kda_module": flash_kda.__file__,
"flash_kda_extension": flash_kda_C.__file__,
},
"correctness": {
"all_exact": all(
case["output_exact"] and case["state_exact"]
for case in correctness
),
"cases": correctness,
},
"benchmarks": benchmarks,
}
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(result, indent=2) + "\n")
print(json.dumps(result, indent=2))
if __name__ == "__main__":
main()