dump SMEM layout info
This commit is contained in:
@@ -1,7 +1,5 @@
|
||||
"""
|
||||
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.
|
||||
Dump FMHA SMEM layout info and try to extract UMMA descriptor values.
|
||||
"""
|
||||
import torch
|
||||
import sys
|
||||
@@ -9,8 +7,27 @@ sys.path.insert(0, '.')
|
||||
|
||||
from dsv4.kernels.attention.fmha import FmhaKernel
|
||||
|
||||
print("=== FMHA SMEM Layout Info ===")
|
||||
kernel = FmhaKernel(head_dim=64, use_smem_p=False, normalize=True)
|
||||
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")
|
||||
|
||||
# The _s objects are (inner_swizzle, outer_layout) tuples
|
||||
q_s = kernel.q_smem_s
|
||||
k_s = kernel.k_smem_s
|
||||
|
||||
print(f"q_smem_s type: {type(q_s)}")
|
||||
print(f"q_smem_s: {q_s}")
|
||||
print(f"k_smem_s: {k_s}")
|
||||
|
||||
# Try to access inner/outer
|
||||
if hasattr(q_s, 'inner'):
|
||||
print(f"Q inner (swizzle): {q_s.inner}")
|
||||
print(f"Q outer (layout): {q_s.outer}")
|
||||
elif isinstance(q_s, tuple):
|
||||
print(f"Q is tuple of length {len(q_s)}")
|
||||
for i, x in enumerate(q_s):
|
||||
print(f" [{i}]: {type(x)} = {x}")
|
||||
|
||||
# SMEM sizes from the kernel
|
||||
print(f"\nQ SMEM size: {kernel.q_smem_size}")
|
||||
print(f"K SMEM size: {kernel.k_smem_size}")
|
||||
print(f"V SMEM size: {kernel.v_smem_size}")
|
||||
print(f"Total SMEM: {kernel.smem_size}")
|
||||
|
||||
Reference in New Issue
Block a user