h3-blackwell-runtime/research/vortex_exact_attention/tools/run_capability_probes.py
2026-08-26 14:44:28 +07:00

288 lines
12 KiB
Python

"""Build and run isolated VEA-B SM121 capability probes."""
from __future__ import annotations
import argparse
import contextlib
import hashlib
import io
import json
import os
import statistics
import sys
import time
from pathlib import Path
import torch
from torch.utils.cpp_extension import load
PROJECT = Path(__file__).resolve().parents[1]
KERNELS = PROJECT / "kernels"
ATTRIBUTE_NAMES = (
"registers_per_thread", "static_shared_bytes", "local_bytes_per_thread",
"max_threads_per_block", "binary_version", "ptx_version",
"max_dynamic_shared_bytes", "resident_ctas_per_sm",
)
KERNEL_NAMES = ("handoff_named", "handoff_mbarrier", "tensor_issue", "qk_role", "pv_role", "combined_ownership")
EPOCHS = 591
CURRENT_MAINLOOP_MS = 237.089
SYNC_BUDGET_MS = CURRENT_MAINLOOP_MS * 0.05
def sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def percentile(values: list[float], fraction: float) -> float:
ordered = sorted(values)
rank = (len(ordered) - 1) * fraction
low = int(rank)
high = min(low + 1, len(ordered) - 1)
weight = rank - low
return ordered[low] * (1.0 - weight) + ordered[high] * weight
def summarize(values: list[float]) -> dict:
return {
"count": len(values), "p10_ms": percentile(values, 0.10),
"p50_ms": percentile(values, 0.50), "p90_ms": percentile(values, 0.90),
"p95_ms": percentile(values, 0.95), "mean_ms": statistics.mean(values),
"min_ms": min(values), "max_ms": max(values),
}
def benchmark(fn, warmup: int, iterations: int) -> tuple[dict, object]:
result = None
for _ in range(warmup):
result = fn()
torch.cuda.synchronize()
values = []
for _ in range(iterations):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
result = fn()
end.record()
end.synchronize()
values.append(start.elapsed_time(end))
return summarize(values), result
def build_extension(build_dir: Path, build_log: Path):
build_dir.mkdir(parents=True, exist_ok=True)
stream = io.StringIO()
started = time.perf_counter()
with contextlib.redirect_stdout(stream), contextlib.redirect_stderr(stream):
extension = load(
name="vortex_exact_phase2a_probe",
sources=[str(KERNELS / "vea_b_probe.cpp"), str(KERNELS / "vea_b_probe.cu")],
build_directory=str(build_dir),
extra_cflags=["-O3"],
extra_cuda_cflags=["-O3", "-lineinfo", "--ptxas-options=-v"],
verbose=True,
)
text = stream.getvalue()
build_log.parent.mkdir(parents=True, exist_ok=True)
build_log.write_text(text, encoding="utf-8")
return extension, time.perf_counter() - started
def expected_payload_checksum(epochs: int) -> int:
total = 0
for epoch in range(epochs):
total += sum((epoch + column) & 255 for column in range(64))
total += 2 * (epoch & 255)
return total
def validate_handoff(result: torch.Tensor, epochs: int, payload: bool) -> dict:
host = result.cpu()
expected_checksum = expected_payload_checksum(epochs) if payload else 0
checks = {
"all_producers_completed": bool(torch.all(host[:, 0] == epochs)),
"all_checksums_match": bool(torch.all(host[:, 1] == expected_checksum)),
"publication_errors": int(host[:, 2].sum().item()),
"all_consumers_completed": bool(torch.all(host[:, 3] == 128)),
"expected_checksum": expected_checksum,
"observed_checksum_first_block": int(host[0, 1].item()),
}
checks["passed"] = all((
checks["all_producers_completed"], checks["all_checksums_match"],
checks["publication_errors"] == 0, checks["all_consumers_completed"],
))
return checks
def analyze_clock_overlap(clocks: torch.Tensor) -> dict:
host = clocks.cpu()
ratios = []
for block in range(host.shape[0]):
int_start = int(host[block, 2:6, 0].min().item())
int_end = int(host[block, 2:6, 1].max().item())
fp_start = int(host[block, 6:10, 0].min().item())
fp_end = int(host[block, 6:10, 1].max().item())
overlap = max(0, min(int_end, fp_end) - max(int_start, fp_start))
denominator = max(1, min(int_end - int_start, fp_end - fp_start))
ratios.append(overlap / denominator)
return {
"blocks": len(ratios),
"blocks_with_positive_overlap": sum(value > 0 for value in ratios),
"positive_overlap_fraction": sum(value > 0 for value in ratios) / len(ratios),
"p50_overlap_ratio": percentile(ratios, 0.50),
"p10_overlap_ratio": percentile(ratios, 0.10),
}
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--build-dir", type=Path, required=True)
parser.add_argument("--build-log", type=Path, required=True)
parser.add_argument("--image", required=True)
parser.add_argument("--handoff-blocks", type=int, default=16576)
parser.add_argument("--probe-blocks", type=int, default=128)
parser.add_argument("--epochs", type=int, default=EPOCHS)
parser.add_argument("--mma-iterations", type=int, default=8192)
parser.add_argument("--warmup", type=int, default=3)
parser.add_argument("--iterations", type=int, default=20)
parser.add_argument("--capture", choices=("none", "handoff_barrier", "handoff_mbarrier", "handoff_payload", "handoff_mbarrier_payload", "tensor_overlap", "qk_role", "pv_role", "combined_ownership"), default="none")
parser.add_argument("--cuda-profiler-capture", action="store_true")
return parser.parse_args()
def main() -> None:
args = parse_args()
torch.manual_seed(440420)
extension, build_seconds = build_extension(args.build_dir, args.build_log)
attributes_tensor = extension.attributes(50 * 1024)
attributes = {
kernel: {name: int(value) for name, value in zip(ATTRIBUTE_NAMES, row.tolist(), strict=True)}
for kernel, row in zip(KERNEL_NAMES, attributes_tensor, strict=True)
}
if args.capture != "none":
if args.cuda_profiler_capture:
torch.cuda.cudart().cudaProfilerStart()
with torch.cuda.nvtx.range(f"vortex_exact_{args.capture}"):
if args.capture == "handoff_barrier":
captured = extension.handoff(args.handoff_blocks, args.epochs, False)
elif args.capture == "handoff_mbarrier":
captured = extension.handoff_mbarrier(args.handoff_blocks, args.epochs, False)
elif args.capture == "handoff_payload":
captured = extension.handoff(args.handoff_blocks, args.epochs, True)
elif args.capture == "handoff_mbarrier_payload":
captured = extension.handoff_mbarrier(args.handoff_blocks, args.epochs, True)
elif args.capture == "tensor_overlap":
captured = extension.tensor_issue(args.probe_blocks, args.mma_iterations, 0)[0]
else:
role = {"qk_role": 0, "pv_role": 1, "combined_ownership": 2}[args.capture]
captured = extension.register_probe(args.probe_blocks, 4, role)
torch.cuda.synchronize()
if args.cuda_profiler_capture:
torch.cuda.cudart().cudaProfilerStop()
report = {
"status": "capture_only", "capture": args.capture,
"captured_checksum": float(captured.float().sum().item()),
"captured_values": captured.cpu().tolist(),
"attributes": attributes, "build_seconds": build_seconds,
}
else:
barrier_timing, barrier_result = benchmark(
lambda: extension.handoff(args.handoff_blocks, args.epochs, False),
args.warmup, args.iterations,
)
mbarrier_timing, mbarrier_result = benchmark(
lambda: extension.handoff_mbarrier(args.handoff_blocks, args.epochs, False),
args.warmup, args.iterations,
)
payload_timing, payload_result = benchmark(
lambda: extension.handoff_mbarrier(args.probe_blocks, args.epochs, True),
args.warmup, args.iterations,
)
payload_repeat = extension.handoff_mbarrier(args.probe_blocks, args.epochs, True)
torch.cuda.synchronize()
tensor_results = {}
overlap_clocks = None
for name, mode in (("overlap", 0), ("int8_only", 1), ("fp8_only", 2)):
timing, result = benchmark(
lambda mode=mode: extension.tensor_issue(args.probe_blocks, args.mma_iterations, mode),
args.warmup, args.iterations,
)
tensor_results[name] = {"timing": timing, "checksum": int(result[0].sum().item())}
if name == "overlap":
overlap_clocks = result[1]
register_results = {}
for name, role in (("qk_role", 0), ("pv_role", 1), ("combined_ownership", 2)):
timing, value = benchmark(
lambda role=role: extension.register_probe(args.probe_blocks, 4, role),
args.warmup, args.iterations,
)
register_results[name] = {"timing": timing, "checksum": float(value.sum().item())}
report = {
"status": "measured_capability_probe",
"attributes": attributes,
"handoff": {
"selected_primitive": "cuda_block_scope_mbarrier",
"barrier_only_full_grid": barrier_timing,
"mbarrier_only_full_grid": mbarrier_timing,
"payload_representative_blocks": payload_timing,
"barrier_budget_ms": SYNC_BUDGET_MS,
"named_barrier_under_budget": barrier_timing["p50_ms"] < SYNC_BUDGET_MS,
"mbarrier_under_budget": mbarrier_timing["p50_ms"] < SYNC_BUDGET_MS,
"barrier_validation": validate_handoff(barrier_result, args.epochs, False),
"mbarrier_validation": validate_handoff(mbarrier_result, args.epochs, False),
"payload_validation": validate_handoff(payload_result, args.epochs, True),
"deterministic_repeat": bool(torch.equal(payload_result, payload_repeat)),
},
"tensor_issue": {
**tensor_results,
"clock_overlap": analyze_clock_overlap(overlap_clocks),
"overlap_faster_than_serial_sum": tensor_results["overlap"]["timing"]["p50_ms"] < (
tensor_results["int8_only"]["timing"]["p50_ms"] + tensor_results["fp8_only"]["timing"]["p50_ms"]
),
},
"register_roles": register_results,
"advancement": {
"registers_at_most_200": attributes["combined_ownership"]["registers_per_thread"] <= 200,
"local_bytes_zero": attributes["combined_ownership"]["local_bytes_per_thread"] == 0,
"one_resident_ten_warp_cta": attributes["handoff_mbarrier"]["resident_ctas_per_sm"] >= 1,
},
"build_seconds": build_seconds,
}
properties = torch.cuda.get_device_properties(0)
report.update({
"schema": "vortex-exact-phase2a-capability-probes", "version": 1,
"environment": {
"image": args.image, "gpu": properties.name,
"compute_capability": [properties.major, properties.minor],
"multiprocessor_count": properties.multi_processor_count,
"torch": torch.__version__, "cuda": torch.version.cuda,
"python": sys.version, "argv": sys.argv,
"environment_switches": {name: value for name, value in sorted(os.environ.items()) if name.startswith(("CUDA_", "TORCH_", "MAX_JOBS"))},
},
"sources": {
"cpp_sha256": sha256_file(KERNELS / "vea_b_probe.cpp"),
"cuda_sha256": sha256_file(KERNELS / "vea_b_probe.cu"),
"build_log": str(args.build_log),
},
"configuration": {
"threads": 320, "warps": 10, "producer_warps": 6,
"consumer_warps": 4, "epochs": args.epochs,
"dynamic_shared_bytes": 50 * 1024,
"handoff_blocks": args.handoff_blocks, "probe_blocks": args.probe_blocks,
"mma_iterations": args.mma_iterations,
},
"claims": {"attention_kernel_implemented": False, "production_changed": False},
})
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8")
print(json.dumps(report, indent=2), flush=True)
if __name__ == "__main__":
main()