#include #include #include #include #include #include #include #include #include #include namespace py = pybind11; #define LT_CHECK(call) \ do { \ const cublasStatus_t status_ = (call); \ if (status_ != CUBLAS_STATUS_SUCCESS) \ throw std::runtime_error(std::string(#call) + " failed: " + \ std::to_string(static_cast(status_))); \ } while (0) namespace { thread_local cublasLtHandle_t handle = nullptr; cublasLtHandle_t get_handle() { if (!handle) LT_CHECK(cublasLtCreate(&handle)); return handle; } void require_cuda_contiguous(const torch::Tensor& tensor, const char* name) { TORCH_CHECK(tensor.is_cuda(), name, " must be CUDA"); TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous"); } void require_same_device(const torch::Tensor& tensor, const torch::Tensor& reference, const char* name) { TORCH_CHECK(tensor.device() == reference.device(), name, " must be on the same CUDA device as activation_qdata"); } struct Problem { cublasLtMatmulDesc_t operation = nullptr; cublasLtMatrixLayout_t a = nullptr, b = nullptr, c = nullptr, d = nullptr; Problem(const torch::Tensor& activation_qdata, const torch::Tensor& activation_block_scale, const torch::Tensor& weight_qdata, const torch::Tensor& weight_block_scale) { require_cuda_contiguous(activation_qdata, "activation_qdata"); require_cuda_contiguous(activation_block_scale, "activation_block_scale"); require_cuda_contiguous(weight_qdata, "weight_qdata"); require_cuda_contiguous(weight_block_scale, "weight_block_scale"); require_same_device(activation_block_scale, activation_qdata, "activation_block_scale"); require_same_device(weight_qdata, activation_qdata, "weight_qdata"); require_same_device(weight_block_scale, activation_qdata, "weight_block_scale"); TORCH_CHECK(activation_qdata.scalar_type() == at::kByte && weight_qdata.scalar_type() == at::kByte, "packed NVFP4 operands must use uint8 storage"); TORCH_CHECK(activation_qdata.dim() == 2 && weight_qdata.dim() == 2, "packed NVFP4 operands must be rank two"); TORCH_CHECK(activation_block_scale.element_size() == 1 && weight_block_scale.element_size() == 1, "NVFP4 block scales must use one-byte E4M3 storage"); TORCH_CHECK(activation_qdata.size(1) == weight_qdata.size(1), "packed K dimensions differ"); // This is deliberately the same column-major reinterpretation used by // Comfy Kitchen 0.2.31: weight is Lt A, activation is Lt B, and D is D^T. const int64_t m = weight_qdata.size(0); // row-major N const int64_t n = activation_qdata.size(0); // row-major padded M const int64_t k = activation_qdata.size(1) * 2; TORCH_CHECK(activation_block_scale.numel() >= n * (k / 16), "activation_block_scale storage is too small"); TORCH_CHECK(weight_block_scale.numel() >= m * (k / 16), "weight_block_scale storage is too small"); LT_CHECK(cublasLtMatmulDescCreate(&operation, CUBLAS_COMPUTE_32F, CUDA_R_32F)); cublasLtMatmulMatrixScale_t scale_mode = CUBLASLT_MATMUL_MATRIX_SCALE_VEC16_UE4M3; LT_CHECK(cublasLtMatmulDescSetAttribute( operation, CUBLASLT_MATMUL_DESC_A_SCALE_MODE, &scale_mode, sizeof(scale_mode))); LT_CHECK(cublasLtMatmulDescSetAttribute( operation, CUBLASLT_MATMUL_DESC_B_SCALE_MODE, &scale_mode, sizeof(scale_mode))); const cublasOperation_t transa = CUBLAS_OP_T; const cublasOperation_t transb = CUBLAS_OP_N; LT_CHECK(cublasLtMatmulDescSetAttribute( operation, CUBLASLT_MATMUL_DESC_TRANSA, &transa, sizeof(transa))); LT_CHECK(cublasLtMatmulDescSetAttribute( operation, CUBLASLT_MATMUL_DESC_TRANSB, &transb, sizeof(transb))); const void* a_scale = weight_block_scale.data_ptr(); const void* b_scale = activation_block_scale.data_ptr(); LT_CHECK(cublasLtMatmulDescSetAttribute( operation, CUBLASLT_MATMUL_DESC_A_SCALE_POINTER, &a_scale, sizeof(a_scale))); LT_CHECK(cublasLtMatmulDescSetAttribute( operation, CUBLASLT_MATMUL_DESC_B_SCALE_POINTER, &b_scale, sizeof(b_scale))); const cublasDataType_t scale_type = CUDA_R_32F; LT_CHECK(cublasLtMatmulDescSetAttribute( operation, CUBLASLT_MATMUL_DESC_SCALE_TYPE, &scale_type, sizeof(scale_type))); const cublasLtPointerMode_t pointer_mode = CUBLASLT_POINTER_MODE_DEVICE; LT_CHECK(cublasLtMatmulDescSetAttribute( operation, CUBLASLT_MATMUL_DESC_POINTER_MODE, &pointer_mode, sizeof(pointer_mode))); const cublasLtEpilogue_t epilogue = CUBLASLT_EPILOGUE_DEFAULT; LT_CHECK(cublasLtMatmulDescSetAttribute( operation, CUBLASLT_MATMUL_DESC_EPILOGUE, &epilogue, sizeof(epilogue))); LT_CHECK(cublasLtMatrixLayoutCreate(&a, CUDA_R_4F_E2M1, k, m, k)); LT_CHECK(cublasLtMatrixLayoutCreate(&b, CUDA_R_4F_E2M1, k, n, k)); LT_CHECK(cublasLtMatrixLayoutCreate(&c, CUDA_R_16BF, m, n, m)); LT_CHECK(cublasLtMatrixLayoutCreate(&d, CUDA_R_16BF, m, n, m)); } ~Problem() { if (d) cublasLtMatrixLayoutDestroy(d); if (c) cublasLtMatrixLayoutDestroy(c); if (b) cublasLtMatrixLayoutDestroy(b); if (a) cublasLtMatrixLayoutDestroy(a); if (operation) cublasLtMatmulDescDestroy(operation); } }; template bool config_get(const cublasLtMatmulAlgo_t& algo, cublasLtMatmulAlgoConfigAttributes_t attr, T* value) { size_t written = 0; return cublasLtMatmulAlgoConfigGetAttribute(&algo, attr, value, sizeof(T), &written) == CUBLAS_STATUS_SUCCESS && written == sizeof(T); } template void put_config(py::dict& result, const char* name, const cublasLtMatmulAlgo_t& algo, cublasLtMatmulAlgoConfigAttributes_t attr) { T value{}; if (config_get(algo, attr, &value)) result[name] = value; } template void put_cap_scalar(py::dict& caps, const char* name, const cublasLtMatmulAlgo_t& algo, cublasLtMatmulAlgoCapAttributes_t attr) { T value{}; size_t written = 0; if (cublasLtMatmulAlgoCapGetAttribute(&algo, attr, &value, sizeof(value), &written) == CUBLAS_STATUS_SUCCESS && written == sizeof(value)) caps[name] = value; } void put_cap_array(py::dict& caps, const char* name, const cublasLtMatmulAlgo_t& algo, cublasLtMatmulAlgoCapAttributes_t attr) { size_t bytes = 0; if (cublasLtMatmulAlgoCapGetAttribute(&algo, attr, nullptr, 0, &bytes) != CUBLAS_STATUS_SUCCESS || bytes == 0) return; std::vector values((bytes + sizeof(uint32_t) - 1) / sizeof(uint32_t)); size_t written = 0; if (cublasLtMatmulAlgoCapGetAttribute(&algo, attr, values.data(), bytes, &written) != CUBLAS_STATUS_SUCCESS) return; values.resize(written / sizeof(uint32_t)); caps[name] = values; } py::dict describe(const cublasLtMatmulAlgo_t& algo, const cublasLtMatmulHeuristicResult_t& checked, cublasStatus_t api_status) { py::dict result; for (const char* name : {"algorithm_id", "tile_id", "stages_id", "split_k", "reduction_scheme", "custom_option", "cta_swizzle", "inner_shape", "cluster_shape"}) result[name] = py::none(); put_config(result, "algorithm_id", algo, CUBLASLT_ALGO_CONFIG_ID); put_config(result, "tile_id", algo, CUBLASLT_ALGO_CONFIG_TILE_ID); put_config(result, "stages_id", algo, CUBLASLT_ALGO_CONFIG_STAGES_ID); // CUDA 13 documents SPLITK_NUM as int32_t. Preserve negative library // sentinel values instead of wrapping them into fictitious huge factors. put_config(result, "split_k", algo, CUBLASLT_ALGO_CONFIG_SPLITK_NUM); put_config(result, "reduction_scheme", algo, CUBLASLT_ALGO_CONFIG_REDUCTION_SCHEME); put_config(result, "custom_option", algo, CUBLASLT_ALGO_CONFIG_CUSTOM_OPTION); put_config(result, "cta_swizzle", algo, CUBLASLT_ALGO_CONFIG_CTA_SWIZZLING); #if CUDA_VERSION >= 12000 put_config(result, "inner_shape", algo, CUBLASLT_ALGO_CONFIG_INNER_SHAPE_ID); put_config(result, "cluster_shape", algo, CUBLASLT_ALGO_CONFIG_CLUSTER_SHAPE_ID); #endif result["required_workspace_bytes"] = checked.workspaceSize; result["waves"] = checked.wavesCount; result["state"] = static_cast(checked.state); result["api_status"] = static_cast(api_status); result["valid"] = api_status == CUBLAS_STATUS_SUCCESS && checked.state == CUBLAS_STATUS_SUCCESS; py::dict caps; put_cap_scalar(caps, "split_k_support", algo, CUBLASLT_ALGO_CAP_SPLITK_SUPPORT); put_cap_scalar(caps, "reduction_scheme_mask", algo, CUBLASLT_ALGO_CAP_REDUCTION_SCHEME_MASK); put_cap_scalar(caps, "cta_swizzle_support", algo, CUBLASLT_ALGO_CAP_CTA_SWIZZLING_SUPPORT); put_cap_scalar(caps, "custom_option_max", algo, CUBLASLT_ALGO_CAP_CUSTOM_OPTION_MAX); put_cap_scalar(caps, "strided_batch_support", algo, CUBLASLT_ALGO_CAP_STRIDED_BATCH_SUPPORT); put_cap_scalar(caps, "out_of_place_result_support", algo, CUBLASLT_ALGO_CAP_OUT_OF_PLACE_RESULT_SUPPORT); put_cap_array(caps, "tile_ids", algo, CUBLASLT_ALGO_CAP_TILE_IDS); put_cap_array(caps, "stages_ids", algo, CUBLASLT_ALGO_CAP_STAGES_IDS); caps["inner_cluster_shape_capability_note"] = "This CUDA 13 cublasLt.h exposes config IDs but no public capability " "attributes that enumerate inner/cluster shape IDs."; result["capabilities"] = caps; return result; } cublasLtMatmulAlgo_t init_algo(int algorithm_id) { cublasLtMatmulAlgo_t algo{}; LT_CHECK(cublasLtMatmulAlgoInit( get_handle(), CUBLAS_COMPUTE_32F, CUDA_R_32F, CUDA_R_4F_E2M1, CUDA_R_4F_E2M1, CUDA_R_16BF, CUDA_R_16BF, algorithm_id, &algo)); return algo; } template void maybe_set(cublasLtMatmulAlgo_t* algo, const py::dict& config, const char* key, cublasLtMatmulAlgoConfigAttributes_t attr) { if (!config.contains(key) || config[key].is_none()) return; const T value = config[key].cast(); LT_CHECK(cublasLtMatmulAlgoConfigSetAttribute(algo, attr, &value, sizeof(value))); } cublasLtMatmulAlgo_t configured_algo(const py::dict& config) { TORCH_CHECK(config.contains("algorithm_id"), "algorithm_id is required"); auto algo = init_algo(config["algorithm_id"].cast()); maybe_set(&algo, config, "tile_id", CUBLASLT_ALGO_CONFIG_TILE_ID); maybe_set(&algo, config, "stages_id", CUBLASLT_ALGO_CONFIG_STAGES_ID); maybe_set(&algo, config, "split_k", CUBLASLT_ALGO_CONFIG_SPLITK_NUM); maybe_set(&algo, config, "reduction_scheme", CUBLASLT_ALGO_CONFIG_REDUCTION_SCHEME); maybe_set(&algo, config, "custom_option", CUBLASLT_ALGO_CONFIG_CUSTOM_OPTION); maybe_set(&algo, config, "cta_swizzle", CUBLASLT_ALGO_CONFIG_CTA_SWIZZLING); #if CUDA_VERSION >= 12000 maybe_set(&algo, config, "inner_shape", CUBLASLT_ALGO_CONFIG_INNER_SHAPE_ID); maybe_set(&algo, config, "cluster_shape", CUBLASLT_ALGO_CONFIG_CLUSTER_SHAPE_ID); #endif return algo; } py::list enumerate(torch::Tensor activation_qdata, torch::Tensor activation_block_scale, torch::Tensor weight_qdata, torch::Tensor weight_block_scale, int64_t max_workspace, int requested_count) { c10::cuda::CUDAGuard guard(activation_qdata.device()); Problem problem(activation_qdata, activation_block_scale, weight_qdata, weight_block_scale); cublasLtMatmulPreference_t preference = nullptr; LT_CHECK(cublasLtMatmulPreferenceCreate(&preference)); LT_CHECK(cublasLtMatmulPreferenceSetAttribute( preference, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, &max_workspace, sizeof(max_workspace))); std::vector found(requested_count); int returned = 0; const auto status = cublasLtMatmulAlgoGetHeuristic( get_handle(), problem.operation, problem.a, problem.b, problem.c, problem.d, preference, requested_count, found.data(), &returned); cublasLtMatmulPreferenceDestroy(preference); LT_CHECK(status); py::list output; for (int i = 0; i < returned; ++i) output.append(describe(found[i].algo, found[i], CUBLAS_STATUS_SUCCESS)); return output; } py::dict check(torch::Tensor activation_qdata, torch::Tensor activation_block_scale, torch::Tensor weight_qdata, torch::Tensor weight_block_scale, py::dict config) { c10::cuda::CUDAGuard guard(activation_qdata.device()); Problem problem(activation_qdata, activation_block_scale, weight_qdata, weight_block_scale); auto algo = configured_algo(config); cublasLtMatmulHeuristicResult_t result{}; const auto status = cublasLtMatmulAlgoCheck( get_handle(), problem.operation, problem.a, problem.b, problem.c, problem.d, &algo, &result); return describe(algo, result, status); } void run(torch::Tensor activation_qdata, torch::Tensor activation_block_scale, torch::Tensor weight_qdata, torch::Tensor weight_block_scale, torch::Tensor alpha, torch::Tensor beta, torch::Tensor output, torch::Tensor workspace, py::dict config) { c10::cuda::CUDAGuard guard(activation_qdata.device()); require_cuda_contiguous(alpha, "alpha"); require_cuda_contiguous(beta, "beta"); require_cuda_contiguous(output, "output"); require_cuda_contiguous(workspace, "workspace"); require_same_device(alpha, activation_qdata, "alpha"); require_same_device(beta, activation_qdata, "beta"); require_same_device(output, activation_qdata, "output"); require_same_device(workspace, activation_qdata, "workspace"); TORCH_CHECK(alpha.scalar_type() == at::kFloat && alpha.numel() == 1, "alpha must be one device FP32 value"); TORCH_CHECK(beta.scalar_type() == at::kFloat && beta.numel() == 1, "beta must be one device FP32 value"); TORCH_CHECK(output.scalar_type() == at::kBFloat16 && output.dim() == 2, "output must be rank-two BF16"); TORCH_CHECK(workspace.scalar_type() == at::kByte, "workspace must use uint8 storage"); TORCH_CHECK(output.size(0) == activation_qdata.size(0) && output.size(1) == weight_qdata.size(0), "output must be [packed activation rows, weight rows]"); Problem problem(activation_qdata, activation_block_scale, weight_qdata, weight_block_scale); auto algo = configured_algo(config); cublasLtMatmulHeuristicResult_t checked{}; LT_CHECK(cublasLtMatmulAlgoCheck(get_handle(), problem.operation, problem.a, problem.b, problem.c, problem.d, &algo, &checked)); TORCH_CHECK(checked.state == CUBLAS_STATUS_SUCCESS, "selected algorithm failed AlgoCheck with state ", static_cast(checked.state)); TORCH_CHECK(checked.workspaceSize <= static_cast(workspace.numel()), "selected algorithm requires ", checked.workspaceSize, " workspace bytes but caller supplied ", workspace.numel()); const auto stream = at::cuda::getCurrentCUDAStream( activation_qdata.get_device()).stream(); void* workspace_ptr = workspace.numel() ? workspace.data_ptr() : nullptr; LT_CHECK(cublasLtMatmul( get_handle(), problem.operation, alpha.data_ptr(), weight_qdata.data_ptr(), problem.a, activation_qdata.data_ptr(), problem.b, beta.data_ptr(), output.data_ptr(), problem.c, output.data_ptr(), problem.d, &algo, workspace_ptr, workspace.numel(), stream)); } py::dict build_info() { py::dict result; result["cuda_version"] = CUDA_VERSION; result["cublas_version"] = CUBLAS_VERSION; result["stream_k_public_control"] = false; result["stream_k_note"] = "CUDA 13 cuBLASLt exposes no documented MatmulAlgoConfig attribute " "that directly selects Stream-K. Negative SPLITK_NUM values returned " "by heuristics are preserved as undocumented library sentinels, not " "claimed as public Stream-K control."; return result; } } // namespace PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) { module.def("enumerate", &enumerate); module.def("check", &check); module.def("run", &run); module.def("build_info", &build_info); }