Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
52 changes: 16 additions & 36 deletions src/backend/cuda/runtime/cuda_device_api.cc
Original file line number Diff line number Diff line change
Expand Up @@ -602,15 +602,10 @@ TVM_FFI_STATIC_INIT_BLOCK() {
auto is_valid_swizzle =
swizzle_kind == CU_TENSOR_MAP_SWIZZLE_NONE || swizzle_kind == CU_TENSOR_MAP_SWIZZLE_32B ||
swizzle_kind == CU_TENSOR_MAP_SWIZZLE_64B || swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B;
#ifdef CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B
is_valid_swizzle = is_valid_swizzle || swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B;
#endif
#ifdef CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B_FLIP_8B
is_valid_swizzle =
is_valid_swizzle || swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B_FLIP_8B;
#endif
#ifdef CU_TENSOR_MAP_SWIZZLE_128B_ATOM_64B
is_valid_swizzle = is_valid_swizzle || swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B_ATOM_64B;
#if (CUDA_VERSION >= 12080)
is_valid_swizzle = is_valid_swizzle || swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B ||
swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B_FLIP_8B ||
swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B_ATOM_64B;
#endif
TVM_FFI_ICHECK(is_valid_swizzle)
<< "Unsupported swizzle enum value: " << static_cast<int>(swizzle_kind);
Expand All @@ -628,43 +623,28 @@ TVM_FFI_STATIC_INIT_BLOCK() {
<< "Unsupported oobFill enum value: " << static_cast<int>(oob_fill_kind);

bool is_packed_16u4_align8 = false;
#ifdef CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B
is_packed_16u4_align8 = cu_dtype == CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B;
#endif
bool is_packed_16u4_align16 = false;
#ifdef CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN16B
is_packed_16u4_align16 = cu_dtype == CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN16B;
#endif
bool is_packed_16u6_align16 = false;
#ifdef CU_TENSOR_MAP_DATA_TYPE_16U6_ALIGN16B
#if (CUDA_VERSION >= 12080)
is_packed_16u4_align8 = cu_dtype == CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B;
is_packed_16u4_align16 = cu_dtype == CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN16B;
is_packed_16u6_align16 = cu_dtype == CU_TENSOR_MAP_DATA_TYPE_16U6_ALIGN16B;
#endif
auto is_packed_align16 = is_packed_16u4_align16 || is_packed_16u6_align16;
auto is_packed_dtype = is_packed_16u4_align8 || is_packed_align16;
auto is_floating_dtype = cu_dtype == CU_TENSOR_MAP_DATA_TYPE_FLOAT16 ||
cu_dtype == CU_TENSOR_MAP_DATA_TYPE_FLOAT32 ||
cu_dtype == CU_TENSOR_MAP_DATA_TYPE_FLOAT64 ||
cu_dtype == CU_TENSOR_MAP_DATA_TYPE_BFLOAT16;
#ifdef CU_TENSOR_MAP_DATA_TYPE_FLOAT32_FTZ
is_floating_dtype = is_floating_dtype || cu_dtype == CU_TENSOR_MAP_DATA_TYPE_FLOAT32_FTZ;
#endif
#ifdef CU_TENSOR_MAP_DATA_TYPE_TFLOAT32
is_floating_dtype = is_floating_dtype || cu_dtype == CU_TENSOR_MAP_DATA_TYPE_TFLOAT32;
#endif
#ifdef CU_TENSOR_MAP_DATA_TYPE_TFLOAT32_FTZ
is_floating_dtype = is_floating_dtype || cu_dtype == CU_TENSOR_MAP_DATA_TYPE_TFLOAT32_FTZ;
#endif
cu_dtype == CU_TENSOR_MAP_DATA_TYPE_BFLOAT16 ||
cu_dtype == CU_TENSOR_MAP_DATA_TYPE_FLOAT32_FTZ ||
cu_dtype == CU_TENSOR_MAP_DATA_TYPE_TFLOAT32 ||
cu_dtype == CU_TENSOR_MAP_DATA_TYPE_TFLOAT32_FTZ;

auto is_128b_swizzle = swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B;
#ifdef CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B
is_128b_swizzle = is_128b_swizzle || swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B;
#endif
#ifdef CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B_FLIP_8B
is_128b_swizzle =
is_128b_swizzle || swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B_FLIP_8B;
#endif
#ifdef CU_TENSOR_MAP_SWIZZLE_128B_ATOM_64B
is_128b_swizzle = is_128b_swizzle || swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B_ATOM_64B;
#if (CUDA_VERSION >= 12080)
is_128b_swizzle = is_128b_swizzle || swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B ||
swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B_FLIP_8B ||
swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B_ATOM_64B;
#endif

// Host-side validation for documented cuTensorMapEncodeTiled requirements.
Expand Down Expand Up @@ -702,7 +682,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
if (is_packed_16u4_align16) {
bool supported_swizzle =
swizzle_kind == CU_TENSOR_MAP_SWIZZLE_NONE || swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B;
#ifdef CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B
#if (CUDA_VERSION >= 12080)
supported_swizzle = supported_swizzle || swizzle_kind == CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B;
#endif
TVM_FFI_ICHECK(supported_swizzle)
Expand Down
16 changes: 12 additions & 4 deletions tests/python/tirx/test_tirx_kernels_registry_correctness.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@ def _load_workloads():
_DISTRIBUTED_KERNELS = frozenset(
{"allgather_gemm", "deepgemm_fp8_fp4_mega_moe", "gemm_reduce_scatter"}
)
_MEGA_MOE_KERNELS = frozenset({"deepgemm_fp8_fp4_mega_moe", "sm100_fp8_fp4_mega_moe"})
_XDIST_CUDA_DEVICE = None


Expand Down Expand Up @@ -157,10 +158,12 @@ def _registry_gpu_lock(kernel_name, config):
)
yield
finally:
torch.cuda.empty_cache()
for lock_file in reversed(lock_files):
fcntl.flock(lock_file.fileno(), fcntl.LOCK_UN)
lock_file.close()
try:
torch.cuda.empty_cache()
finally:
for lock_file in reversed(lock_files):
fcntl.flock(lock_file.fileno(), fcntl.LOCK_UN)
lock_file.close()


@pytest.mark.parametrize(("kernel_name", "config"), _manifest_kernel_config_cases())
Expand All @@ -172,5 +175,10 @@ def test_manifest_tirx_kernel_correctness(kernel_name, config):
pytest.skip(
f"requires {required_devices} CUDA devices, but only {visible_devices} are visible"
)
if kernel_name in _MEGA_MOE_KERNELS:
pytest.skip(
"MegaMoE requires its dedicated multi-process scheduler; this suite's "
"processes own CUDA contexts that its physical-device assignment rejects"
)
with _registry_gpu_lock(kernel_name, config):
kernel_runner.run_kernel_test(kernel_name, config, registry=_KERNELS)
Loading