fix: simplify UMMA dump script
This commit is contained in:
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user