diff --git a/scripts/dump_umma_desc.py b/scripts/dump_umma_desc.py index 6f83e78f..933bb79a 100644 --- a/scripts/dump_umma_desc.py +++ b/scripts/dump_umma_desc.py @@ -1,82 +1,16 @@ """ -Test script to dump UMMA descriptors from the CuTeDSL FMHA. -Uses the existing FMHA kernel with a patch to print descriptors. +Minimal CuTeDSL kernel that prints UMMA descriptors for FMHA Q and K. +The kernel doesn't need to do anything useful — just construct the +SMEM layouts and dump the descriptors. """ import torch import sys sys.path.insert(0, '.') from dsv4.kernels.attention.fmha import FmhaKernel -import cute -# Construct the FMHA kernel +print("=== FMHA SMEM Layout Info ===") 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}") +print(f"q_smem_size: {kernel.q_smem_size} bytes") +print(f"k_smem_size: {kernel.k_smem_size} bytes") +print(f"smem_size: {kernel.smem_size} bytes")