Inside NVIDIA’s cuDNN Graph API: Fusion, Autotuning, and Plan Reuse with cuDNN Frontend


@section("5. SDPA (Flash Attention) with causal masking")
def sdpa_demo():
   if not HAS_SDPA:
       raise RuntimeError(f"fused SDPA needs SM80+ (Ampere), this GPU is sm_{SM}")
   b, h, s, d = 4, 16, 1024, 64
   scale = 1.0 / math.sqrt(d)
   SDPA_FLOPS = 4 * b * h * s * s * d * 0.5
   q = torch.randn(b, h, s, d, device=DEV, dtype=DTYPE)
   k = torch.randn(b, h, s, d, device=DEV, dtype=DTYPE)
   v = torch.randn(b, h, s, d, device=DEV, dtype=DTYPE)
   o = torch.empty(b, h, s, d, device=DEV, dtype=DTYPE)
   g = cudnn.pygraph(
       handle=HANDLE, name="sdpa",
       io_data_type=TORCH2CUDNN[DTYPE],
       intermediate_data_type=cudnn.data_type.FLOAT,
       compute_data_type=cudnn.data_type.FLOAT,
   )
   Q, Kt, V = tensor_of(g, q, "Q"), tensor_of(g, k, "K"), tensor_of(g, v, "V")
   causal = True
   try:
       O, _stats = g.sdpa(name="sdpa", q=Q, k=Kt, v=V,
                          is_inference=True, attn_scale=scale, use_causal_mask=True)
   except TypeError:
       try:
           O, _stats = g.sdpa(name="sdpa", q=Q, k=Kt, v=V,
                              is_inference=True, attn_scale=scale,
                              diagonal_alignment=cudnn.diagonal_alignment.TOP_LEFT,
                              right_bound=0)
       except Exception:
           causal = False
           O, _stats = g.sdpa(name="sdpa", q=Q, k=Kt, v=V,
                              is_inference=True, attn_scale=scale)
   print(f"    causal masking: {causal}")
   O.set_output(True).set_data_type(TORCH2CUDNN[DTYPE])
   O.set_dim(list(o.size())).set_stride(list(o.stride()))
   build(g)
   ws = workspace_for(g)
   pack = {Q: q, Kt: k, V: v, O: o}
   g.execute(pack, ws)
   torch.cuda.synchronize()
   ref = torch.nn.functional.scaled_dot_product_attention(q, k, v, is_causal=causal, scale=scale)
   rel = ((o.float() - ref.float()).abs().max() / ref.float().abs().max()).item()
   print(f"    shape   : b{b} h{h} s{s} d{d}   workspace {ws.numel()/1024:.1f} KiB")
   print(f"    rel err : {rel:.2e}")
   ms = bench(lambda: g.execute(pack, ws))
   ms_t = bench(lambda: torch.nn.functional.scaled_dot_product_attention(
       q, k, v, is_causal=causal, scale=scale))
   print()
   report("cuDNN FE SDPA", ms, SDPA_FLOPS)
   report("torch SDPA (backend's choice)", ms_t, SDPA_FLOPS)
   print("    Note: torch may already be dispatching to cuDNN or FlashAttention,")
   print("    so parity here is the expected, healthy outcome.")
   return f"{ms:.3f} ms, {tflops(SDPA_FLOPS, ms):.1f} TFLOP/s"
sdpa_demo()
@section("6. Serialize a built graph, reload it, execute by UID")
def serialization():
   Bsz, M, Kd, Nd = 8, 256, 512, 256
   a = torch.randn(Bsz, M, Kd, device=DEV, dtype=DTYPE)
   bm = torch.randn(Bsz, Kd, Nd, device=DEV, dtype=DTYPE)
   out = torch.empty(Bsz, M, Nd, device=DEV, dtype=DTYPE)
   UID_A, UID_B, UID_C = 1, 2, 3
   g = cudnn.pygraph(
       handle=HANDLE, name="serializable_mm",
       io_data_type=TORCH2CUDNN[DTYPE],
       intermediate_data_type=cudnn.data_type.FLOAT,
       compute_data_type=cudnn.data_type.FLOAT,
   )
   A = tensor_of(g, a, "A").set_uid(UID_A)
   Bt = tensor_of(g, bm, "B").set_uid(UID_B)
   C = g.matmul(A=A, B=Bt, compute_data_type=cudnn.data_type.FLOAT)
   C.set_output(True).set_data_type(TORCH2CUDNN[DTYPE]).set_uid(UID_C)
   t0 = time.perf_counter()
   build(g)
   cold_ms = (time.perf_counter() - t0) * 1e3
   blob = g.serialize()
   print(f"    cold build      : {cold_ms:.1f} ms")
   print(f"    serialized plan : {len(blob)} bytes (cache this to disk / ship it)")
   t0 = time.perf_counter()
   g2 = cudnn.pygraph()
   try:
       g2.deserialize(HANDLE, blob)
   except TypeError:
       g2.deserialize(blob)
   warm_ms = (time.perf_counter() - t0) * 1e3
   print(f"    deserialize     : {warm_ms:.1f} ms  -> {cold_ms/max(warm_ms,1e-6):.1f}x faster startup")
   ws = torch.empty(max(g2.get_workspace_size(), 1), device=DEV, dtype=torch.uint8)
   g2.execute({UID_A: a, UID_B: bm, UID_C: out}, ws, handle=HANDLE)
   torch.cuda.synchronize()
   ref = torch.bmm(a.float(), bm.float())
   rel = ((out.float() - ref).abs().max() / ref.abs().max()).item()
   print(f"    rel err after reload: {rel:.2e}")
   return f"{len(blob)} B blob, reload {cold_ms/max(warm_ms,1e-6):.1f}x faster than rebuild"
serialization()



Source link

  • Related Posts

    Google Releases Gemini 3.8 Live and 3.8 Live Extended Thinking for Production Grade Voice Agents

    Google has introduced Gemini 3.8 Live and Gemini 3.8 Live Extended Thinking, its most advanced live dialogue models to date. Both are native speech to speech models built for real…

    Agent-net Open Sources Webagent: A Go Harness That Turns Any Website into a Guarded AI Agent

    Agent-net, the team building an agent-to-agent marketplace where AI agents discover, trust, and pay each other, has released Webagent, an open source harness for standing up public-facing business agents. So,…

    Leave a Reply

    Your email address will not be published. Required fields are marked *