#!/usr/bin/env python3 """Math SDPA with the DiT's real input layout: [B, L, H, D] permuted to [B, H, L, D] (non-contiguous, exactly what diffusers' native attention backend passes) vs the same tensors made contiguous. Speed and error vs an fp32 reference, masked, at the cabin prompt's real length. usage: sdpa_layout.py [L] """ import sys import time import torch import torch.nn.functional as F from torch.nn.attention import SDPBackend, sdpa_kernel L = int(sys.argv[1]) if len(sys.argv) > 1 else 5759 H, D, dev = 30, 128, "cuda" g = torch.Generator(device=dev).manual_seed(0) blhd = [torch.randn(1, L, H, D, device=dev, dtype=torch.bfloat16, generator=g) for _ in range(3)] q, k, v = (x.permute(0, 2, 1, 3) for x in blhd) # views, as diffusers passes them qc, kc, vc = (x.contiguous() for x in (q, k, v)) mask = torch.ones(1, 1, 1, L, dtype=torch.bool, device=dev) mask[..., L - 64:] = False with sdpa_kernel(SDPBackend.MATH): ref = F.scaled_dot_product_attention(qc.float(), kc.float(), vc.float(), attn_mask=mask) def run(tag, a, b, c, bf16_reduction): torch.backends.cuda.allow_fp16_bf16_reduction_math_sdp(bf16_reduction) with sdpa_kernel(SDPBackend.MATH): out = F.scaled_dot_product_attention(a, b, c, attn_mask=mask) torch.cuda.synchronize() t0 = time.perf_counter() for _ in range(3): out = F.scaled_dot_product_attention(a, b, c, attn_mask=mask) torch.cuda.synchronize() ms = (time.perf_counter() - t0) / 3 * 1000 rel = ((out.float() - ref).norm() / ref.norm()).item() exact = torch.equal(out, base) if base is not None else None print(f" {tag:34s} {ms:8.2f} ms rel_l2 {rel:.3e} identical_to_default: {exact}") return out base = None print(f"torch {torch.__version__} | L={L} | q strides {tuple(q.stride())} contiguous={q.is_contiguous()}") base = run("permuted views (DiT today), fp32", q, k, v, False) run("contiguous, fp32 (math unchanged)", qc, kc, vc, False) run("permuted views, bf16 reduction", q, k, v, True) run("contiguous, bf16 reduction", qc, kc, vc, True) torch.backends.cuda.allow_fp16_bf16_reduction_math_sdp(False)