fix: simplify UMMA dump script

This commit is contained in:
2026-05-28 08:57:49 +00:00
parent 948a3f8a7a
commit fe0588d906

View File

@@ -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")