#!/usr/bin/env python3 """cuDNN Frontend graph boundary, adapted from NVIDIA's official quick start. Coverage is parser_only. Import, graph build, execution, and numerics require a pinned CUDA/cuDNN/PyTorch environment and have not been observed here. """ import cudnn import torch def run() -> dict[str, object]: batch, m, n, k = 16, 32, 64, 128 a = torch.randn(batch, m, k, device="cuda", dtype=torch.bfloat16) b = torch.randn(1, k, n, device="cuda", dtype=torch.bfloat16) # Host: describe a persistent operation graph and its precision contract. with cudnn.Graph( io_data_type=torch.bfloat16, compute_data_type=torch.float32, inputs=["matmul::A", "matmul::B"], outputs=["out"], ) as graph: output = graph.matmul(name="matmul", A=a, B=b) output.set_name("out").set_output(True) # Device: the built graph chooses supported cuDNN engine(s) and launches. candidate = graph(a, b, handle=cudnn.create_handle()) reference = torch.matmul(a.float(), b.float()).to(torch.bfloat16) return { "shape": list(candidate.shape), "matches": bool(torch.allclose(candidate, reference, atol=1e-2, rtol=1e-2)), "max_abs_error": float((candidate - reference).abs().max()), } if __name__ == "__main__": print(run())