diff --git a/scripts/dump_umma_desc.py b/scripts/dump_umma_desc.py new file mode 100644 index 00000000..6f83e78f --- /dev/null +++ b/scripts/dump_umma_desc.py @@ -0,0 +1,82 @@ +""" +Test script to dump UMMA descriptors from the CuTeDSL FMHA. +Uses the existing FMHA kernel with a patch to print descriptors. +""" +import torch +import sys +sys.path.insert(0, '.') + +from dsv4.kernels.attention.fmha import FmhaKernel +import cute + +# Construct the FMHA kernel +kernel = FmhaKernel(head_dim=64, use_smem_p=False, normalize=True) + +# Get the SMEM layouts +q_layout = kernel.q_smem_layout_staged +k_layout = kernel.k_smem_layout_staged +v_layout = kernel.v_smem_layout_staged + +print(f"Q layout: {q_layout}") +print(f"K layout: {k_layout}") +print(f"V layout: {v_layout}") + +# Get the SMEM sizes +print(f"Q SMEM: {kernel.q_smem_size} bytes") +print(f"K SMEM: {kernel.k_smem_size} bytes") +print(f"V SMEM: {kernel.v_smem_size} bytes") +print(f"Total SMEM: {kernel.smem_size} bytes") + +# Try to construct UMMA descriptors +try: + from cute.arch.mma_sm100_umma import make_umma_desc + from cutlass.utils.blackwell_helpers import OperandMajorMode + + # Create dummy SMEM tensors + q_ptr = cute.make_smem_ptr(cutlass.BFloat16, 0) + q_tensor = cute.make_tensor(q_ptr, q_layout) + k_tensor = cute.make_tensor(cute.make_smem_ptr(cutlass.BFloat16, kernel.q_smem_size), k_layout) + + desc_q = make_umma_desc(OperandMajorMode.MN_MAJOR, q_tensor) + desc_k = make_umma_desc(OperandMajorMode.K_MAJOR, k_tensor) + + print(f"\ndesc_q = 0x{desc_q:016x}") + print(f"desc_k = 0x{desc_k:016x}") + + # Decode descriptors + sa_q = desc_q & 0x3FFF + lbo_q = (desc_q >> 16) & 0x3FFF + sbo_q = (desc_q >> 32) & 0x3FFF + ver_q = (desc_q >> 46) & 0x3 + lt_q = (desc_q >> 61) & 0x7 + print(f"Q: start_addr={sa_q} LBO={lbo_q} SBO={sbo_q} ver={ver_q} layout={lt_q}") + + sa_k = desc_k & 0x3FFF + lbo_k = (desc_k >> 16) & 0x3FFF + sbo_k = (desc_k >> 32) & 0x3FFF + ver_k = (desc_k >> 46) & 0x3 + lt_k = (desc_k >> 61) & 0x7 + print(f"K: start_addr={sa_k} LBO={lbo_k} SBO={sbo_k} ver={ver_k} layout={lt_k}") +except Exception as e: + print(f"\nCouldn't construct descriptors: {e}") + import traceback + traceback.print_exc() + + # Try to print layout strides instead + print("\nQ layout strides:") + try: + for m in [0, 1, 2, 3, 7, 8, 127]: + for n in [0, 1, 2, 3, 7, 8, 15, 31, 63]: + offset = q_layout(m, n) + print(f" ({m},{n}) -> {offset}") + except Exception as e2: + print(f" Error: {e2}") + + print("\nK layout strides:") + try: + for m in [0, 1, 2, 3, 7, 8, 127]: + for n in [0, 1, 2, 3, 7, 8, 15, 31, 63]: + offset = k_layout(m, n) + print(f" ({m},{n}) -> {offset}") + except Exception as e2: + print(f" Error: {e2}")