liangsu9988's picture
Promote latest kernel artifacts to main
8c8128e verified
Raw
History Blame Contribute Delete
5.54 kB
# Copyright (c) 2025, Tri Dao.
from typing import Tuple
from functools import lru_cache
import torch
try:
from triton.tools.disasm import extract
except ImportError:
extract = None
import cutlass
import cutlass.cute as cute
from cutlass.cutlass_dsl import NumericMeta
from cutlass.cute.runtime import from_dlpack
StaticTypes = (cutlass.Constexpr, NumericMeta, int, bool, str, float, type(None))
load_cubin_module_data_og = cutlass.base_dsl.runtime.cuda.load_cubin_module_data
cute_compile_og = cute.compile
torch2cute_dtype_map = {
torch.float16: cutlass.Float16,
torch.bfloat16: cutlass.BFloat16,
torch.float32: cutlass.Float32,
torch.float8_e4m3fn: cutlass.Float8E4M3FN,
torch.float8_e5m2: cutlass.Float8E5M2,
}
@lru_cache
def get_max_active_clusters(cluster_size):
return cutlass.utils.HardwareInfo().get_max_active_clusters(cluster_size=cluster_size)
@lru_cache
def get_device_capacity(device: torch.device = None) -> Tuple[int, int]:
return torch.cuda.get_device_capability(device)
def assume_strides_aligned(t):
"""Assume all strides except the last are divisible by 128 bits.
Python int strides (e.g., stride=0 from GQA expand) are kept as-is
since they're static and don't need alignment assumptions.
"""
divby = 128 // t.element_type.width
strides = tuple(s if isinstance(s, int) else cute.assume(s, divby=divby) for s in t.stride[:-1])
return (*strides, t.stride[-1])
def assume_tensor_aligned(t):
"""Rebuild a tensor with 128-bit aligned stride assumptions. Passes through None."""
if t is None:
return None
return cute.make_tensor(t.iterator, cute.make_layout(t.shape, stride=assume_strides_aligned(t)))
def to_cute_tensor(t, assumed_align=16, leading_dim=-1, fully_dynamic=False, enable_tvm_ffi=True):
"""Convert torch tensor to cute tensor for TVM FFI. leading_dim=-1 defaults to t.ndim-1."""
# NOTE: torch 2.9.1 doesn't support fp8 via DLPack but 2.11.0 nightly does
# currently export raw bytes as uint8 and tell cutlass correct type
# can directly export as fp8 when torch supports it
if t.dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
tensor = from_dlpack(
t.view(torch.uint8).detach(),
assumed_align=assumed_align,
enable_tvm_ffi=enable_tvm_ffi,
)
tensor.element_type = (
cutlass.Float8E4M3FN if t.dtype == torch.float8_e4m3fn else cutlass.Float8E5M2
)
else:
tensor = from_dlpack(t.detach(), assumed_align=assumed_align, enable_tvm_ffi=enable_tvm_ffi)
if fully_dynamic:
return tensor.mark_layout_dynamic()
if leading_dim == -1:
leading_dim = t.ndim - 1
return tensor.mark_layout_dynamic(leading_dim=leading_dim)
def to_cute_aux_tensor(t, enable_tvm_ffi=True):
"""Convert torch tensor to cute tensor for TVM FFI, tailored to FlexAttention aux tensors.
This allows the user to specify alignment and leading dimension for aux tensors used in
custom score_mod callables.
"""
assumed_align: int = getattr(t, "__assumed_align__", None)
leading_dim: int = getattr(t, "__leading_dim__", None)
fully_dynamic: bool = leading_dim is None
return to_cute_tensor(
t,
assumed_align=assumed_align,
leading_dim=leading_dim,
fully_dynamic=fully_dynamic,
enable_tvm_ffi=enable_tvm_ffi,
)
def get_aux_tensor_metadata(aux_tensors):
return tuple(
(
getattr(t, "__assumed_align__", 0),
getattr(t, "__leading_dim__", -1),
hasattr(t, "__leading_dim__"),
)
for t in aux_tensors
)
def get_broadcast_dims(tensor: torch.Tensor) -> Tuple[bool, ...]:
"""Return tuple of bools indicating which dims have stride=0 (broadcast).
This is useful for compile keys since CuTe's mark_layout_dynamic() keeps
stride=0 as static, meaning kernels compiled with different broadcast
patterns are not interchangeable.
"""
return tuple(s == 0 for s in tensor.stride())
# credit: monellz (https://github.com/NVIDIA/cutlass/issues/2658#issuecomment-3630564264)
def dump_kernel_attributes(compiled_kernel):
from cuda.bindings import driver
from cutlass.utils import HardwareInfo
import torch
device_id = torch.cuda.current_device()
hardware_info = HardwareInfo(device_id=device_id)
cubin_data = compiled_kernel.artifacts.CUBIN
assert cubin_data is not None, "cubin_data is None, need '--keep-cubin' option when compiling"
cuda_library = hardware_info._checkCudaErrors(
driver.cuLibraryLoadData(cubin_data, None, None, 0, None, None, 0)
)
kernels = hardware_info._checkCudaErrors(driver.cuLibraryEnumerateKernels(1, cuda_library))
kernel = hardware_info._checkCudaErrors(driver.cuKernelGetFunction(kernels[0]))
# more metrics: https://docs.nvidia.com/cuda/cuda-driver-api/group__CUDA__EXEC.html#group__CUDA__EXEC_1g5e92a1b0d8d1b82cb00dcfb2de15961b
local_size_bytes = hardware_info._checkCudaErrors(
driver.cuFuncGetAttribute(
driver.CUfunction_attribute.CU_FUNC_ATTRIBUTE_LOCAL_SIZE_BYTES,
kernel,
)
)
num_regs = hardware_info._checkCudaErrors(
driver.cuFuncGetAttribute(
driver.CUfunction_attribute.CU_FUNC_ATTRIBUTE_NUM_REGS,
kernel,
)
)
print("--- Kernel Info ---")
print(f"local_size_bytes: {local_size_bytes}")
print(f"num_regs: {num_regs}")
print("--- End Kernel Info ---")