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