diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml index a7b3e6b880..4b9d91ed56 100644 --- a/.github/workflows/publish.yml +++ b/.github/workflows/publish.yml @@ -41,7 +41,9 @@ jobs: # Using ubuntu-22.04 instead of 24.04 for more compatibility (glibc). Ideally we'd use the # manylinux docker image, but I haven't figured out how to install CUDA on manylinux. os: [ubuntu-22.04] - python-version: ["3.8", "3.9", "3.10", "3.11", "3.12", "3.13"] + # setup.py builds a single abi3 wheel pinned to the cp310 floor, which every torch + # version here (2.4-2.8) supports, so one Python build covers the whole matrix. + python-version: ["3.10"] torch-version: ["2.4.0", "2.5.1", "2.6.0", "2.7.1", "2.8.0"] cuda-version: ["12.9.1"] # We need separate wheels that either uses C++11 ABI (-D_GLIBCXX_USE_CXX11_ABI) or not. @@ -49,11 +51,6 @@ jobs: # Without this we get import error (undefined symbol: _ZN3c105ErrorC2ENS_14SourceLocationESs) # when building without C++11 ABI and using it on nvcr images. cxx11_abi: ["FALSE", "TRUE"] - exclude: - # see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix - # Pytorch < 2.5 does not support Python 3.13 - - torch-version: "2.4.0" - python-version: "3.13" uses: ./.github/workflows/_build.yml with: runs-on: ${{ matrix.os }} diff --git a/csrc/apis/attention.hpp b/csrc/apis/attention.hpp index 1abfd5c9b0..c9686112a3 100644 --- a/csrc/apis/attention.hpp +++ b/csrc/apis/attention.hpp @@ -15,6 +15,8 @@ #endif #include "layout.hpp" +#include +#include "../torch_library_utils.hpp" namespace deep_gemm::attention { @@ -463,41 +465,130 @@ static torch::Tensor fp8_paged_mqa_logits(const torch::Tensor& q, } #endif -static void register_apis(pybind11::module_& m) { +} // namespace deep_gemm::attention + +namespace deep_gemm::torch_registration { + +using namespace deep_gemm::torch_utils; + +#if DG_FP8_COMPATIBLE and DG_TENSORMAP_COMPATIBLE +static void fp8_gemm_nt_skip_head_mid( + const torch::Tensor& a, const torch::Tensor& sfa, + const torch::Tensor& b, const torch::Tensor& sfb, + const torch::Tensor& d, + const std::vector& head_splits, + const c10::optional>& recipe, + const std::string& compiled_dims, + const bool& disable_ue8m0_cast) { + attention::fp8_gemm_nt_skip_head_mid( + {a, sfa}, {b, sfb}, d, + list_to_tuple3(head_splits), + list_to_recipe3(recipe), + compiled_dims, disable_ue8m0_cast); +} + +static torch::Tensor fp8_fp4_mqa_logits( + const torch::Tensor& q, const c10::optional& q_sf, + const torch::Tensor& kv, const torch::Tensor& kv_sf, + const torch::Tensor& weights, + const torch::Tensor& cu_seq_len_k_start, + const torch::Tensor& cu_seq_len_k_end, + const bool& clean_logits, + const int64_t& max_seqlen_k, + at::ScalarType logits_dtype) { + return attention::fp8_fp4_mqa_logits( + std::make_tuple(q, q_sf), + std::make_tuple(kv, kv_sf), + weights, cu_seq_len_k_start, cu_seq_len_k_end, + clean_logits, static_cast(max_seqlen_k), + logits_dtype); +} + +static torch::Tensor get_paged_mqa_logits_metadata( + const torch::Tensor& context_lens, const int64_t& block_kv, + const int64_t& num_sms, const c10::optional& indices) { + return attention::get_paged_mqa_logits_metadata( + context_lens, static_cast(block_kv), + static_cast(num_sms), indices); +} + +static torch::Tensor fp8_fp4_paged_mqa_logits( + const torch::Tensor& q, const c10::optional& q_sf, + const torch::Tensor& kv_cache, + const torch::Tensor& weights, + const torch::Tensor& context_lens, + const torch::Tensor& block_table, + const torch::Tensor& schedule_meta, + const int64_t& max_context_len, + const bool& clean_logits, + at::ScalarType logits_dtype, + const c10::optional& indices) { + return attention::fp8_fp4_paged_mqa_logits( + std::make_tuple(q, q_sf), + kv_cache, weights, context_lens, block_table, schedule_meta, + static_cast(max_context_len), clean_logits, + logits_dtype, indices); +} + +static torch::Tensor fp8_mqa_logits( + const torch::Tensor& q, + const torch::Tensor& kv, const torch::Tensor& kv_sf, + const torch::Tensor& weights, + const torch::Tensor& cu_seq_len_k_start, + const torch::Tensor& cu_seq_len_k_end, + const bool& clean_logits, + const int64_t& max_seqlen_k) { + return attention::fp8_mqa_logits( + q, std::make_tuple(kv, kv_sf), weights, + cu_seq_len_k_start, cu_seq_len_k_end, + clean_logits, static_cast(max_seqlen_k)); +} + +static torch::Tensor fp8_paged_mqa_logits( + const torch::Tensor& q, + const torch::Tensor& kv_cache, + const torch::Tensor& weights, + const torch::Tensor& context_lens, + const torch::Tensor& block_table, + const torch::Tensor& schedule_meta, + const int64_t& max_context_len, + const bool& clean_logits, + const c10::optional& indices) { + return attention::fp8_paged_mqa_logits( + q, kv_cache, weights, + context_lens, block_table, schedule_meta, + static_cast(max_context_len), clean_logits, indices); +} +#endif + +} // namespace deep_gemm::torch_registration + +TORCH_LIBRARY_FRAGMENT(deep_gemm, m) { #if DG_FP8_COMPATIBLE and DG_TENSORMAP_COMPATIBLE - m.def("fp8_gemm_nt_skip_head_mid", &fp8_gemm_nt_skip_head_mid, - py::arg("a"), py::arg("b"), py::arg("d"), py::arg("head_splits"), - py::arg("recipe") = std::nullopt, - py::arg("compiled_dims") = "nk", - py::arg("disable_ue8m0_cast") = false); - m.def("fp8_fp4_mqa_logits", &fp8_fp4_mqa_logits, - py::arg("q"), py::arg("kv"), py::arg("weights"), - py::arg("cu_seq_len_k_start"), py::arg("cu_seq_len_k_end"), - py::arg("clean_logits") = true, - py::arg("max_seqlen_k") = 0, - py::arg("logits_dtype") = torch::kFloat32); - m.def("get_paged_mqa_logits_metadata", &get_paged_mqa_logits_metadata, - py::arg("context_lens"), py::arg("block_kv"), py::arg("num_sms"), - py::arg("indices") = std::nullopt); - m.def("fp8_fp4_paged_mqa_logits", &fp8_fp4_paged_mqa_logits, - py::arg("q"), py::arg("kv_cache"), py::arg("weights"), - py::arg("context_lens"), py::arg("block_table"), py::arg("schedule_meta"), - py::arg("max_context_len"), - py::arg("clean_logits") = false, - py::arg("logits_dtype") = torch::kFloat32, - py::arg("indices") = std::nullopt); - // Legacy API - m.def("fp8_mqa_logits", &fp8_mqa_logits, - py::arg("q"), py::arg("kv"), py::arg("weights"), - py::arg("cu_seq_len_k_start"), py::arg("cu_seq_len_k_end"), - py::arg("clean_logits") = true, - py::arg("max_seqlen_k") = 0); - m.def("fp8_paged_mqa_logits", &fp8_paged_mqa_logits, - py::arg("q"), py::arg("kv_cache"), py::arg("weights"), - py::arg("context_lens"), py::arg("block_table"), py::arg("schedule_meta"), - py::arg("max_context_len"), py::arg("clean_logits") = false, - py::arg("indices") = std::nullopt); + m.def( + "fp8_gemm_nt_skip_head_mid(Tensor a, Tensor sfa, Tensor b, Tensor sfb, Tensor(d!) d, int[3] head_splits, int[3]? recipe=None, str compiled_dims='nk', bool disable_ue8m0_cast=False) -> ()"); + m.def( + "fp8_fp4_mqa_logits(Tensor q, Tensor? q_sf, Tensor kv, Tensor kv_sf, Tensor weights, Tensor cu_seq_len_k_start, Tensor cu_seq_len_k_end, bool clean_logits=True, int max_seqlen_k=0, ScalarType logits_dtype=float) -> Tensor"); + m.def( + "get_paged_mqa_logits_metadata(Tensor context_lens, int block_kv, int num_sms, Tensor? indices=None) -> Tensor"); + m.def( + "fp8_fp4_paged_mqa_logits(Tensor q, Tensor? q_sf, Tensor kv_cache, Tensor weights, Tensor context_lens, Tensor block_table, Tensor schedule_meta, int max_context_len, bool clean_logits=False, ScalarType logits_dtype=float, Tensor? indices=None) -> Tensor"); + m.def( + "fp8_mqa_logits(Tensor q, Tensor kv, Tensor kv_sf, Tensor weights, Tensor cu_seq_len_k_start, Tensor cu_seq_len_k_end, bool clean_logits=True, int max_seqlen_k=0) -> Tensor"); + m.def( + "fp8_paged_mqa_logits(Tensor q, Tensor kv_cache, Tensor weights, Tensor context_lens, Tensor block_table, Tensor schedule_meta, int max_context_len, bool clean_logits=False, Tensor? indices=None) -> Tensor"); #endif } -} // namespace deep_gemm::attention +TORCH_LIBRARY_IMPL(deep_gemm, CUDA, m) { + using namespace deep_gemm::torch_registration; + +#if DG_FP8_COMPATIBLE and DG_TENSORMAP_COMPATIBLE + m.impl("fp8_gemm_nt_skip_head_mid", TORCH_FN(fp8_gemm_nt_skip_head_mid)); + m.impl("fp8_fp4_mqa_logits", TORCH_FN(fp8_fp4_mqa_logits)); + m.impl("get_paged_mqa_logits_metadata", TORCH_FN(get_paged_mqa_logits_metadata)); + m.impl("fp8_fp4_paged_mqa_logits", TORCH_FN(fp8_fp4_paged_mqa_logits)); + m.impl("fp8_mqa_logits", TORCH_FN(fp8_mqa_logits)); + m.impl("fp8_paged_mqa_logits", TORCH_FN(fp8_paged_mqa_logits)); +#endif +} diff --git a/csrc/apis/einsum.hpp b/csrc/apis/einsum.hpp index ff3ac590c0..d7a9cabe7f 100644 --- a/csrc/apis/einsum.hpp +++ b/csrc/apis/einsum.hpp @@ -1,7 +1,6 @@ #pragma once -#include -#include +#include #include "../utils/exception.hpp" #include "../utils/format.hpp" @@ -19,6 +18,8 @@ #include "../jit_kernels/impls/sm120_fp8_fp4_gemm_1d1d.hpp" #include "../jit_kernels/impls/smxx_cublaslt.hpp" #endif +#include +#include "../torch_library_utils.hpp" namespace deep_gemm::einsum { @@ -268,17 +269,45 @@ static void fp8_einsum(const std::string& expr, } #endif -static void register_apis(pybind11::module_& m) { +} // namespace deep_gemm::einsum + +namespace deep_gemm::torch_registration { + +using namespace deep_gemm::torch_utils; + #if DG_FP8_COMPATIBLE and DG_TENSORMAP_COMPATIBLE - m.def("einsum", &einsum, - py::arg("expr"), py::arg("a"), py::arg("b"), - py::arg("d"), py::arg("c") = std::nullopt, - py::arg("use_cublaslt") = false); - m.def("fp8_einsum", &fp8_einsum, - py::arg("expr"), py::arg("a"), py::arg("b"), - py::arg("d"), py::arg("c") = std::nullopt, - py::arg("recipe") = std::make_tuple(1, 128, 128)); +static void einsum(const std::string& expr, + const torch::Tensor& a, const torch::Tensor& b, + const torch::Tensor& d, const c10::optional& c, + const bool& use_cublaslt) { + einsum::einsum(expr, a, b, d, c, use_cublaslt); +} + +static void fp8_einsum(const std::string& expr, + const torch::Tensor& a, const torch::Tensor& sfa, + const torch::Tensor& b, const torch::Tensor& sfb, + const torch::Tensor& d, const c10::optional& c, + const std::vector& recipe) { + einsum::fp8_einsum(expr, {a, sfa}, {b, sfb}, d, c, list_to_tuple3(recipe)); +} +#endif + +} // namespace deep_gemm::torch_registration + +TORCH_LIBRARY_FRAGMENT(deep_gemm, m) { +#if DG_FP8_COMPATIBLE and DG_TENSORMAP_COMPATIBLE + m.def( + "einsum(str expr, Tensor a, Tensor b, Tensor(d!) d, Tensor? c=None, bool use_cublaslt=False) -> ()"); + m.def( + "fp8_einsum(str expr, Tensor a, Tensor sfa, Tensor b, Tensor sfb, Tensor(d!) d, Tensor? c=None, int[3] recipe) -> ()"); #endif } -} // namespace deep_gemm::einsum +TORCH_LIBRARY_IMPL(deep_gemm, CUDA, m) { + using namespace deep_gemm::torch_registration; + +#if DG_FP8_COMPATIBLE and DG_TENSORMAP_COMPATIBLE + m.impl("einsum", TORCH_FN(einsum)); + m.impl("fp8_einsum", TORCH_FN(fp8_einsum)); +#endif +} diff --git a/csrc/apis/gemm.hpp b/csrc/apis/gemm.hpp index 902ae699a9..a00954cf65 100644 --- a/csrc/apis/gemm.hpp +++ b/csrc/apis/gemm.hpp @@ -15,6 +15,8 @@ #include "../jit_kernels/impls/smxx_cublaslt.hpp" #include "layout.hpp" +#include +#include "../torch_library_utils.hpp" namespace deep_gemm::gemm { @@ -782,130 +784,314 @@ static void cublaslt_gemm_tt(const torch::Tensor& a, const torch::Tensor& b, cublaslt_gemm_nt(a.transpose(0, 1), b, d, c); } -static void register_apis(pybind11::module_& m) { +} // namespace deep_gemm::gemm + +namespace deep_gemm::torch_registration { + +using namespace deep_gemm::torch_utils; #if DG_FP8_COMPATIBLE and DG_TENSORMAP_COMPATIBLE - // FP8 FP4 GEMMs - m.def("fp8_fp4_gemm_nt", &fp8_fp4_gemm_nt, - py::arg("a"), py::arg("b"), py::arg("d"), - py::arg("c") = std::nullopt, py::arg("recipe") = std::nullopt, - py::arg("recipe_a") = std::nullopt, py::arg("recipe_b") = std::nullopt, - py::arg("compiled_dims") = "nk", - py::arg("disable_ue8m0_cast") = false); - m.def("fp8_fp4_gemm_nn", &fp8_fp4_gemm_nn, - py::arg("a"), py::arg("b"), py::arg("d"), - py::arg("c") = std::nullopt, py::arg("recipe") = std::nullopt, - py::arg("recipe_a") = std::nullopt, py::arg("recipe_b") = std::nullopt, - py::arg("compiled_dims") = "nk", - py::arg("disable_ue8m0_cast") = false); - m.def("fp8_fp4_gemm_tn", &fp8_fp4_gemm_tn, - py::arg("a"), py::arg("b"), py::arg("d"), - py::arg("c") = std::nullopt, py::arg("recipe") = std::nullopt, - py::arg("recipe_a") = std::nullopt, py::arg("recipe_b") = std::nullopt, - py::arg("compiled_dims") = "mn", - py::arg("disable_ue8m0_cast") = false); - m.def("fp8_fp4_gemm_tt", &fp8_fp4_gemm_tt, - py::arg("a"), py::arg("b"), py::arg("d"), - py::arg("c") = std::nullopt, py::arg("recipe") = std::nullopt, - py::arg("recipe_a") = std::nullopt, py::arg("recipe_b") = std::nullopt, - py::arg("compiled_dims") = "mn", - py::arg("disable_ue8m0_cast") = false); - m.def("m_grouped_fp8_fp4_gemm_nt_contiguous", &m_grouped_fp8_fp4_gemm_nt_contiguous, - py::arg("a"), py::arg("b"), py::arg("d"), py::arg("grouped_layout"), - py::arg("recipe") = std::nullopt, - py::arg("recipe_a") = std::nullopt, py::arg("recipe_b") = std::nullopt, - py::arg("compiled_dims") = "nk", - py::arg("disable_ue8m0_cast") = false, - py::arg("use_psum_layout") = false, - py::arg("ensure_zero_padding") = true, - py::arg("expected_m_for_psum_layout") = std::nullopt); - m.def("m_grouped_fp8_fp4_gemm_nn_contiguous", &m_grouped_fp8_fp4_gemm_nn_contiguous, - py::arg("a"), py::arg("b"), py::arg("d"), py::arg("grouped_layout"), - py::arg("recipe") = std::nullopt, - py::arg("recipe_a") = std::nullopt, py::arg("recipe_b") = std::nullopt, - py::arg("compiled_dims") = "nk", - py::arg("disable_ue8m0_cast") = false, - py::arg("use_psum_layout") = false, - py::arg("ensure_zero_padding") = true); - m.def("m_grouped_fp8_fp4_gemm_nt_masked", &m_grouped_fp8_fp4_gemm_nt_masked, - py::arg("a"), py::arg("b"), py::arg("d"), py::arg("masked_m"), - py::arg("expected_m"), py::arg("recipe") = std::nullopt, - py::arg("recipe_a") = std::nullopt, py::arg("recipe_b") = std::nullopt, - py::arg("compiled_dims") = "nk", py::arg("disable_ue8m0_cast") = false); - m.def("k_grouped_fp8_gemm_tn_contiguous", &k_grouped_fp8_gemm_tn_contiguous, - py::arg("a"), py::arg("b"), py::arg("d"), - py::arg("ks_cpu"), py::arg("grouped_layout"), - py::arg("c") = std::nullopt, - py::arg("recipe") = std::make_tuple(1, 1, 128), - py::arg("compiled_dims") = "mn", - py::arg("use_psum_layout") = false); - m.def("k_grouped_fp8_gemm_nt_contiguous", &k_grouped_fp8_gemm_nt_contiguous, - py::arg("a"), py::arg("b"), py::arg("d"), - py::arg("ks_cpu"), py::arg("grouped_layout"), - py::arg("c") = std::nullopt, - py::arg("recipe") = std::make_tuple(1, 1, 128), - py::arg("compiled_dims") = "mn", - py::arg("use_psum_layout") = false); - - // FP8 GEMM alias names - m.attr("fp8_gemm_nt") = m.attr("fp8_fp4_gemm_nt"); - m.attr("fp8_gemm_nn") = m.attr("fp8_fp4_gemm_nn"); - m.attr("fp8_gemm_tn") = m.attr("fp8_fp4_gemm_tn"); - m.attr("fp8_gemm_tt") = m.attr("fp8_fp4_gemm_tt"); - m.attr("m_grouped_fp8_gemm_nt_contiguous") = m.attr("m_grouped_fp8_fp4_gemm_nt_contiguous"); - m.attr("m_grouped_fp8_gemm_nn_contiguous") = m.attr("m_grouped_fp8_fp4_gemm_nn_contiguous"); - m.attr("m_grouped_fp8_gemm_nt_masked") = m.attr("m_grouped_fp8_fp4_gemm_nt_masked"); +static void fp8_fp4_gemm_nt( + const torch::Tensor& a, const torch::Tensor& sfa, + const torch::Tensor& b, const torch::Tensor& sfb, + const torch::Tensor& d, const c10::optional& c, + const c10::optional>& recipe, + const c10::optional>& recipe_a, + const c10::optional>& recipe_b, + const std::string& compiled_dims, const bool& disable_ue8m0_cast) { + gemm::fp8_fp4_gemm_nt({a, sfa}, {b, sfb}, d, c, + list_to_recipe3(recipe), list_to_recipe2(recipe_a), list_to_recipe2(recipe_b), + compiled_dims, disable_ue8m0_cast); +} + +static void fp8_fp4_gemm_nn( + const torch::Tensor& a, const torch::Tensor& sfa, + const torch::Tensor& b, const torch::Tensor& sfb, + const torch::Tensor& d, const c10::optional& c, + const c10::optional>& recipe, + const c10::optional>& recipe_a, + const c10::optional>& recipe_b, + const std::string& compiled_dims, const bool& disable_ue8m0_cast) { + gemm::fp8_fp4_gemm_nn({a, sfa}, {b, sfb}, d, c, + list_to_recipe3(recipe), list_to_recipe2(recipe_a), list_to_recipe2(recipe_b), + compiled_dims, disable_ue8m0_cast); +} + +static void fp8_fp4_gemm_tn( + const torch::Tensor& a, const torch::Tensor& sfa, + const torch::Tensor& b, const torch::Tensor& sfb, + const torch::Tensor& d, const c10::optional& c, + const c10::optional>& recipe, + const c10::optional>& recipe_a, + const c10::optional>& recipe_b, + const std::string& compiled_dims, const bool& disable_ue8m0_cast) { + gemm::fp8_fp4_gemm_tn({a, sfa}, {b, sfb}, d, c, + list_to_recipe3(recipe), list_to_recipe2(recipe_a), list_to_recipe2(recipe_b), + compiled_dims, disable_ue8m0_cast); +} + +static void fp8_fp4_gemm_tt( + const torch::Tensor& a, const torch::Tensor& sfa, + const torch::Tensor& b, const torch::Tensor& sfb, + const torch::Tensor& d, const c10::optional& c, + const c10::optional>& recipe, + const c10::optional>& recipe_a, + const c10::optional>& recipe_b, + const std::string& compiled_dims, const bool& disable_ue8m0_cast) { + gemm::fp8_fp4_gemm_tt({a, sfa}, {b, sfb}, d, c, + list_to_recipe3(recipe), list_to_recipe2(recipe_a), list_to_recipe2(recipe_b), + compiled_dims, disable_ue8m0_cast); +} + +static void m_grouped_fp8_fp4_gemm_nt_contiguous( + const torch::Tensor& a, const torch::Tensor& sfa, + const torch::Tensor& b, const torch::Tensor& sfb, + const torch::Tensor& d, const torch::Tensor& grouped_layout, + const c10::optional>& recipe, + const c10::optional>& recipe_a, + const c10::optional>& recipe_b, + const std::string& compiled_dims, const bool& disable_ue8m0_cast, + const bool& use_psum_layout, const bool& ensure_zero_padding, + const c10::optional& expected_m_for_psum_layout) { + gemm::m_grouped_fp8_fp4_gemm_nt_contiguous( + {a, sfa}, {b, sfb}, d, grouped_layout, + list_to_recipe3(recipe), list_to_recipe2(recipe_a), list_to_recipe2(recipe_b), + compiled_dims, disable_ue8m0_cast, use_psum_layout, ensure_zero_padding, + expected_m_for_psum_layout.has_value() + ? std::make_optional(static_cast(expected_m_for_psum_layout.value())) + : std::nullopt); +} + +static void m_grouped_fp8_fp4_gemm_nn_contiguous( + const torch::Tensor& a, const torch::Tensor& sfa, + const torch::Tensor& b, const torch::Tensor& sfb, + const torch::Tensor& d, const torch::Tensor& grouped_layout, + const c10::optional>& recipe, + const c10::optional>& recipe_a, + const c10::optional>& recipe_b, + const std::string& compiled_dims, const bool& disable_ue8m0_cast, + const bool& use_psum_layout, const bool& ensure_zero_padding) { + gemm::m_grouped_fp8_fp4_gemm_nn_contiguous( + {a, sfa}, {b, sfb}, d, grouped_layout, + list_to_recipe3(recipe), list_to_recipe2(recipe_a), list_to_recipe2(recipe_b), + compiled_dims, disable_ue8m0_cast, use_psum_layout, ensure_zero_padding); +} + +static void m_grouped_fp8_fp4_gemm_nt_masked( + const torch::Tensor& a, const torch::Tensor& sfa, + const torch::Tensor& b, const torch::Tensor& sfb, + const torch::Tensor& d, const torch::Tensor& masked_m, + const int64_t& expected_m, + const c10::optional>& recipe, + const c10::optional>& recipe_a, + const c10::optional>& recipe_b, + const std::string& compiled_dims, const bool& disable_ue8m0_cast) { + gemm::m_grouped_fp8_fp4_gemm_nt_masked( + {a, sfa}, {b, sfb}, d, masked_m, static_cast(expected_m), + list_to_recipe3(recipe), list_to_recipe2(recipe_a), list_to_recipe2(recipe_b), + compiled_dims, disable_ue8m0_cast); +} + +static void k_grouped_fp8_gemm_tn_contiguous( + const torch::Tensor& a, const torch::Tensor& sfa, + const torch::Tensor& b, const torch::Tensor& sfb, + const torch::Tensor& d, + const c10::optional>& ks_cpu, + const torch::Tensor& grouped_layout, + const c10::optional& c, + const std::vector& recipe, + const std::string& compiled_dims, const bool& use_psum_layout) { + gemm::k_grouped_fp8_gemm_tn_contiguous( + {a, sfa}, {b, sfb}, d, + list_to_optional_vector_int(ks_cpu), grouped_layout, c, + list_to_tuple3(recipe), compiled_dims, use_psum_layout); +} + +static void k_grouped_fp8_gemm_nt_contiguous( + const torch::Tensor& a, const torch::Tensor& sfa, + const torch::Tensor& b, const torch::Tensor& sfb, + const torch::Tensor& d, + const c10::optional>& ks_cpu, + const torch::Tensor& grouped_layout, + const c10::optional& c, + const std::vector& recipe, + const std::string& compiled_dims, const bool& use_psum_layout) { + gemm::k_grouped_fp8_gemm_nt_contiguous( + {a, sfa}, {b, sfb}, d, + list_to_optional_vector_int(ks_cpu), grouped_layout, c, + list_to_tuple3(recipe), compiled_dims, use_psum_layout); +} #endif #if DG_TENSORMAP_COMPATIBLE - // BF16 GEMMs - m.def("bf16_gemm_nt", &bf16_gemm_nt, - py::arg("a"), py::arg("b"), py::arg("d"), - py::arg("c") = std::nullopt, - py::arg("compiled_dims") = "nk"); - m.def("bf16_gemm_nn", &bf16_gemm_nn, - py::arg("a"), py::arg("b"), py::arg("d"), - py::arg("c") = std::nullopt, - py::arg("compiled_dims") = "nk"); - m.def("bf16_gemm_tn", &bf16_gemm_tn, - py::arg("a"), py::arg("b"), py::arg("d"), - py::arg("c") = std::nullopt, - py::arg("compiled_dims") = "mn"); - m.def("bf16_gemm_tt", &bf16_gemm_tt, - py::arg("a"), py::arg("b"), py::arg("d"), - py::arg("c") = std::nullopt, - py::arg("compiled_dims") = "mn"); - m.def("m_grouped_bf16_gemm_nt_contiguous", &m_grouped_bf16_gemm_nt_contiguous, - py::arg("a"), py::arg("b"), py::arg("d"), py::arg("grouped_layout"), - py::arg("compiled_dims") = "nk", - py::arg("use_psum_layout") = false, - py::arg("ensure_zero_padding") = true, - py::arg("expected_m_for_psum_layout") = std::nullopt); - m.def("m_grouped_bf16_gemm_nn_contiguous", &m_grouped_bf16_gemm_nn_contiguous, - py::arg("a"), py::arg("b"), py::arg("d"), py::arg("grouped_layout"), - py::arg("compiled_dims") = "nk", - py::arg("use_psum_layout") = false, - py::arg("ensure_zero_padding") = true); - m.def("m_grouped_bf16_gemm_nt_masked", &m_grouped_bf16_gemm_nt_masked, - py::arg("a"), py::arg("b"), py::arg("d"), py::arg("masked_m"), - py::arg("expected_m"), py::arg("compiled_dims") = "nk"); - m.def("k_grouped_bf16_gemm_tn_contiguous", &k_grouped_bf16_gemm_tn_contiguous, - py::arg("a"), py::arg("b"), py::arg("d"), - py::arg("ks_cpu"), py::arg("grouped_layout"), - py::arg("c") = std::nullopt, - py::arg("compiled_dims") = "mn", - py::arg("use_psum_layout") = false); +static void bf16_gemm_nt( + const torch::Tensor& a, const torch::Tensor& b, const torch::Tensor& d, + const c10::optional& c, const std::string& compiled_dims) { + gemm::bf16_gemm_nt(a, b, d, c, compiled_dims); +} + +static void bf16_gemm_nn( + const torch::Tensor& a, const torch::Tensor& b, const torch::Tensor& d, + const c10::optional& c, const std::string& compiled_dims) { + gemm::bf16_gemm_nn(a, b, d, c, compiled_dims); +} + +static void bf16_gemm_tn( + const torch::Tensor& a, const torch::Tensor& b, const torch::Tensor& d, + const c10::optional& c, const std::string& compiled_dims) { + gemm::bf16_gemm_tn(a, b, d, c, compiled_dims); +} + +static void bf16_gemm_tt( + const torch::Tensor& a, const torch::Tensor& b, const torch::Tensor& d, + const c10::optional& c, const std::string& compiled_dims) { + gemm::bf16_gemm_tt(a, b, d, c, compiled_dims); +} + +static void m_grouped_bf16_gemm_nt_contiguous( + const torch::Tensor& a, const torch::Tensor& b, const torch::Tensor& d, + const torch::Tensor& grouped_layout, const std::string& compiled_dims, + const bool& use_psum_layout, const bool& ensure_zero_padding, + const c10::optional& expected_m_for_psum_layout) { + gemm::m_grouped_bf16_gemm_nt_contiguous( + a, b, d, grouped_layout, compiled_dims, + use_psum_layout, ensure_zero_padding, + expected_m_for_psum_layout.has_value() + ? std::make_optional(static_cast(expected_m_for_psum_layout.value())) + : std::nullopt); +} + +static void m_grouped_bf16_gemm_nn_contiguous( + const torch::Tensor& a, const torch::Tensor& b, const torch::Tensor& d, + const torch::Tensor& grouped_layout, const std::string& compiled_dims, + const bool& use_psum_layout, const bool& ensure_zero_padding) { + gemm::m_grouped_bf16_gemm_nn_contiguous( + a, b, d, grouped_layout, compiled_dims, use_psum_layout, ensure_zero_padding); +} + +static void m_grouped_bf16_gemm_nt_masked( + const torch::Tensor& a, const torch::Tensor& b, const torch::Tensor& d, + const torch::Tensor& masked_m, const int64_t& expected_m, + const std::string& compiled_dims) { + gemm::m_grouped_bf16_gemm_nt_masked(a, b, d, masked_m, static_cast(expected_m), compiled_dims); +} + +static void k_grouped_bf16_gemm_tn_contiguous( + const torch::Tensor& a, const torch::Tensor& b, const torch::Tensor& d, + const c10::optional>& ks_cpu, + const torch::Tensor& grouped_layout, + const c10::optional& c, + const std::string& compiled_dims, const bool& use_psum_layout) { + gemm::k_grouped_bf16_gemm_tn_contiguous( + a, b, d, list_to_optional_vector_int(ks_cpu), + grouped_layout, c, compiled_dims, use_psum_layout); +} +#endif + +static void cublaslt_gemm_nt( + const torch::Tensor& a, const torch::Tensor& b, + const torch::Tensor& d, const c10::optional& c) { + gemm::cublaslt_gemm_nt(a, b, d, c); +} + +static void cublaslt_gemm_nn( + const torch::Tensor& a, const torch::Tensor& b, + const torch::Tensor& d, const c10::optional& c) { + gemm::cublaslt_gemm_nn(a, b, d, c); +} + +static void cublaslt_gemm_tn( + const torch::Tensor& a, const torch::Tensor& b, + const torch::Tensor& d, const c10::optional& c) { + gemm::cublaslt_gemm_tn(a, b, d, c); +} + +static void cublaslt_gemm_tt( + const torch::Tensor& a, const torch::Tensor& b, + const torch::Tensor& d, const c10::optional& c) { + gemm::cublaslt_gemm_tt(a, b, d, c); +} + +} // namespace deep_gemm::torch_registration + +TORCH_LIBRARY_FRAGMENT(deep_gemm, m) { +#if DG_FP8_COMPATIBLE and DG_TENSORMAP_COMPATIBLE + // GEMM — FP8/FP4 + m.def( + "fp8_fp4_gemm_nt(Tensor a, Tensor sfa, Tensor b, Tensor sfb, Tensor(d!) d, Tensor? c=None, int[3]? recipe=None, int[2]? recipe_a=None, int[2]? recipe_b=None, str compiled_dims='nk', bool disable_ue8m0_cast=False) -> ()"); + m.def( + "fp8_fp4_gemm_nn(Tensor a, Tensor sfa, Tensor b, Tensor sfb, Tensor(d!) d, Tensor? c=None, int[3]? recipe=None, int[2]? recipe_a=None, int[2]? recipe_b=None, str compiled_dims='nk', bool disable_ue8m0_cast=False) -> ()"); + m.def( + "fp8_fp4_gemm_tn(Tensor a, Tensor sfa, Tensor b, Tensor sfb, Tensor(d!) d, Tensor? c=None, int[3]? recipe=None, int[2]? recipe_a=None, int[2]? recipe_b=None, str compiled_dims='mn', bool disable_ue8m0_cast=False) -> ()"); + m.def( + "fp8_fp4_gemm_tt(Tensor a, Tensor sfa, Tensor b, Tensor sfb, Tensor(d!) d, Tensor? c=None, int[3]? recipe=None, int[2]? recipe_a=None, int[2]? recipe_b=None, str compiled_dims='mn', bool disable_ue8m0_cast=False) -> ()"); + m.def( + "m_grouped_fp8_fp4_gemm_nt_contiguous(Tensor a, Tensor sfa, Tensor b, Tensor sfb, Tensor(d!) d, Tensor grouped_layout, int[3]? recipe=None, int[2]? recipe_a=None, int[2]? recipe_b=None, str compiled_dims='nk', bool disable_ue8m0_cast=False, bool use_psum_layout=False, bool ensure_zero_padding=True, int? expected_m_for_psum_layout=None) -> ()"); + m.def( + "m_grouped_fp8_fp4_gemm_nn_contiguous(Tensor a, Tensor sfa, Tensor b, Tensor sfb, Tensor(d!) d, Tensor grouped_layout, int[3]? recipe=None, int[2]? recipe_a=None, int[2]? recipe_b=None, str compiled_dims='nk', bool disable_ue8m0_cast=False, bool use_psum_layout=False, bool ensure_zero_padding=True) -> ()"); + m.def( + "m_grouped_fp8_fp4_gemm_nt_masked(Tensor a, Tensor sfa, Tensor b, Tensor sfb, Tensor(d!) d, Tensor masked_m, int expected_m, int[3]? recipe=None, int[2]? recipe_a=None, int[2]? recipe_b=None, str compiled_dims='nk', bool disable_ue8m0_cast=False) -> ()"); + m.def( + "k_grouped_fp8_gemm_tn_contiguous(Tensor a, Tensor sfa, Tensor b, Tensor sfb, Tensor(d!) d, int[]? ks_cpu, Tensor grouped_layout, Tensor? c=None, int[3] recipe, str compiled_dims='mn', bool use_psum_layout=False) -> ()"); + m.def( + "k_grouped_fp8_gemm_nt_contiguous(Tensor a, Tensor sfa, Tensor b, Tensor sfb, Tensor(d!) d, int[]? ks_cpu, Tensor grouped_layout, Tensor? c=None, int[3] recipe, str compiled_dims='mn', bool use_psum_layout=False) -> ()"); #endif - // cuBLASLt GEMMs - m.def("cublaslt_gemm_nt", &cublaslt_gemm_nt, - py::arg("a"), py::arg("b"), py::arg("d"), py::arg("c") = std::nullopt); - m.def("cublaslt_gemm_nn", &cublaslt_gemm_nn, - py::arg("a"), py::arg("b"), py::arg("d"), py::arg("c") = std::nullopt); - m.def("cublaslt_gemm_tn", &cublaslt_gemm_tn, - py::arg("a"), py::arg("b"), py::arg("d"), py::arg("c") = std::nullopt); - m.def("cublaslt_gemm_tt", &cublaslt_gemm_tt, - py::arg("a"), py::arg("b"), py::arg("d"), py::arg("c") = std::nullopt); +#if DG_TENSORMAP_COMPATIBLE + // GEMM — BF16 + m.def( + "bf16_gemm_nt(Tensor a, Tensor b, Tensor(d!) d, Tensor? c=None, str compiled_dims='nk') -> ()"); + m.def( + "bf16_gemm_nn(Tensor a, Tensor b, Tensor(d!) d, Tensor? c=None, str compiled_dims='nk') -> ()"); + m.def( + "bf16_gemm_tn(Tensor a, Tensor b, Tensor(d!) d, Tensor? c=None, str compiled_dims='mn') -> ()"); + m.def( + "bf16_gemm_tt(Tensor a, Tensor b, Tensor(d!) d, Tensor? c=None, str compiled_dims='mn') -> ()"); + m.def( + "m_grouped_bf16_gemm_nt_contiguous(Tensor a, Tensor b, Tensor(d!) d, Tensor grouped_layout, str compiled_dims='nk', bool use_psum_layout=False, bool ensure_zero_padding=True, int? expected_m_for_psum_layout=None) -> ()"); + m.def( + "m_grouped_bf16_gemm_nn_contiguous(Tensor a, Tensor b, Tensor(d!) d, Tensor grouped_layout, str compiled_dims='nk', bool use_psum_layout=False, bool ensure_zero_padding=True) -> ()"); + m.def( + "m_grouped_bf16_gemm_nt_masked(Tensor a, Tensor b, Tensor(d!) d, Tensor masked_m, int expected_m, str compiled_dims='nk') -> ()"); + m.def( + "k_grouped_bf16_gemm_tn_contiguous(Tensor a, Tensor b, Tensor(d!) d, int[]? ks_cpu, Tensor grouped_layout, Tensor? c=None, str compiled_dims='mn', bool use_psum_layout=False) -> ()"); +#endif + + // GEMM — cuBLASLt + m.def("cublaslt_gemm_nt(Tensor a, Tensor b, Tensor(d!) d, Tensor? c=None) -> ()"); + m.def("cublaslt_gemm_nn(Tensor a, Tensor b, Tensor(d!) d, Tensor? c=None) -> ()"); + m.def("cublaslt_gemm_tn(Tensor a, Tensor b, Tensor(d!) d, Tensor? c=None) -> ()"); + m.def("cublaslt_gemm_tt(Tensor a, Tensor b, Tensor(d!) d, Tensor? c=None) -> ()"); } -} // namespace deep_gemm::gemm +TORCH_LIBRARY_IMPL(deep_gemm, CUDA, m) { + using namespace deep_gemm::torch_registration; + +#if DG_FP8_COMPATIBLE and DG_TENSORMAP_COMPATIBLE + m.impl("fp8_fp4_gemm_nt", TORCH_FN(fp8_fp4_gemm_nt)); + m.impl("fp8_fp4_gemm_nn", TORCH_FN(fp8_fp4_gemm_nn)); + m.impl("fp8_fp4_gemm_tn", TORCH_FN(fp8_fp4_gemm_tn)); + m.impl("fp8_fp4_gemm_tt", TORCH_FN(fp8_fp4_gemm_tt)); + m.impl("m_grouped_fp8_fp4_gemm_nt_contiguous", TORCH_FN(m_grouped_fp8_fp4_gemm_nt_contiguous)); + m.impl("m_grouped_fp8_fp4_gemm_nn_contiguous", TORCH_FN(m_grouped_fp8_fp4_gemm_nn_contiguous)); + m.impl("m_grouped_fp8_fp4_gemm_nt_masked", TORCH_FN(m_grouped_fp8_fp4_gemm_nt_masked)); + m.impl("k_grouped_fp8_gemm_tn_contiguous", TORCH_FN(k_grouped_fp8_gemm_tn_contiguous)); + m.impl("k_grouped_fp8_gemm_nt_contiguous", TORCH_FN(k_grouped_fp8_gemm_nt_contiguous)); +#endif + +#if DG_TENSORMAP_COMPATIBLE + m.impl("bf16_gemm_nt", TORCH_FN(bf16_gemm_nt)); + m.impl("bf16_gemm_nn", TORCH_FN(bf16_gemm_nn)); + m.impl("bf16_gemm_tn", TORCH_FN(bf16_gemm_tn)); + m.impl("bf16_gemm_tt", TORCH_FN(bf16_gemm_tt)); + m.impl("m_grouped_bf16_gemm_nt_contiguous", TORCH_FN(m_grouped_bf16_gemm_nt_contiguous)); + m.impl("m_grouped_bf16_gemm_nn_contiguous", TORCH_FN(m_grouped_bf16_gemm_nn_contiguous)); + m.impl("m_grouped_bf16_gemm_nt_masked", TORCH_FN(m_grouped_bf16_gemm_nt_masked)); + m.impl("k_grouped_bf16_gemm_tn_contiguous", TORCH_FN(k_grouped_bf16_gemm_tn_contiguous)); +#endif + + m.impl("cublaslt_gemm_nt", TORCH_FN(cublaslt_gemm_nt)); + m.impl("cublaslt_gemm_nn", TORCH_FN(cublaslt_gemm_nn)); + m.impl("cublaslt_gemm_tn", TORCH_FN(cublaslt_gemm_tn)); + m.impl("cublaslt_gemm_tt", TORCH_FN(cublaslt_gemm_tt)); +} diff --git a/csrc/apis/hyperconnection.hpp b/csrc/apis/hyperconnection.hpp index a695f5938e..8b37f65bb6 100644 --- a/csrc/apis/hyperconnection.hpp +++ b/csrc/apis/hyperconnection.hpp @@ -7,6 +7,7 @@ #include "../jit_kernels/impls/sm100_tf32_hc_prenorm_gemm.hpp" #include "../jit_kernels/impls/sm120_tf32_hc_prenorm_gemm.hpp" #endif +#include namespace deep_gemm::hyperconnection { @@ -59,15 +60,37 @@ static void tf32_hc_prenorm_gemm(const torch::Tensor& a, DG_HOST_UNREACHABLE("Unsupported architecture"); } } +#endif + +} // namespace deep_gemm::hyperconnection + +namespace deep_gemm::torch_registration { +#if DG_FP8_COMPATIBLE and DG_TENSORMAP_COMPATIBLE +static void tf32_hc_prenorm_gemm(const torch::Tensor& a, const torch::Tensor& b, + const torch::Tensor& d, const torch::Tensor& sqr_sum, + const c10::optional& num_splits) { + hyperconnection::tf32_hc_prenorm_gemm( + a, b, d, sqr_sum, + num_splits.has_value() + ? std::make_optional(static_cast(num_splits.value())) + : std::nullopt); +} #endif -static void register_apis(pybind11::module_& m) { +} // namespace deep_gemm::torch_registration + +TORCH_LIBRARY_FRAGMENT(deep_gemm, m) { #if DG_FP8_COMPATIBLE and DG_TENSORMAP_COMPATIBLE - m.def("tf32_hc_prenorm_gemm", &tf32_hc_prenorm_gemm, - py::arg("a"), py::arg("b"), py::arg("d"), py::arg("sqr_sum"), - py::arg("num_splits") = std::nullopt); + m.def( + "tf32_hc_prenorm_gemm(Tensor a, Tensor b, Tensor(d!) d, Tensor(sqr_sum!) sqr_sum, int? num_splits=None) -> ()"); #endif } -} // namespace deep_gemm::hyperconnection +TORCH_LIBRARY_IMPL(deep_gemm, CUDA, m) { + using namespace deep_gemm::torch_registration; + +#if DG_FP8_COMPATIBLE and DG_TENSORMAP_COMPATIBLE + m.impl("tf32_hc_prenorm_gemm", TORCH_FN(tf32_hc_prenorm_gemm)); +#endif +} diff --git a/csrc/apis/layout.hpp b/csrc/apis/layout.hpp index 5f5c870a62..10bceabf93 100644 --- a/csrc/apis/layout.hpp +++ b/csrc/apis/layout.hpp @@ -7,6 +7,8 @@ #if DG_TENSORMAP_COMPATIBLE #include "../jit_kernels/impls/smxx_layout.hpp" #endif +#include +#include "../torch_library_utils.hpp" namespace deep_gemm::layout { @@ -137,34 +139,104 @@ static torch::Tensor transform_k_grouped_sf_into_required_layout(const torch::Te #endif -static void register_apis(pybind11::module_& m) { +} // namespace deep_gemm::layout + +namespace deep_gemm::torch_registration { + +using namespace deep_gemm::torch_utils; + +#if DG_TENSORMAP_COMPATIBLE +static torch::Tensor transform_sf_into_required_layout( + const torch::Tensor& sf, const int64_t& mn, const int64_t& k, + const std::vector& recipe, + const c10::optional& num_groups, + const c10::optional& is_sfa, + const bool& disable_ue8m0_cast, + const c10::optional& psum_layout) { + return layout::transform_sf_into_required_layout( + sf, static_cast(mn), static_cast(k), + list_to_recipe_variant(recipe), + num_groups.has_value() ? std::make_optional(static_cast(num_groups.value())) : std::nullopt, + is_sfa, + disable_ue8m0_cast, + psum_layout); +} + +static torch::Tensor get_mn_major_tma_aligned_tensor(const torch::Tensor& sf) { + return ::deep_gemm::get_mn_major_tma_aligned_tensor(sf); +} + +static torch::Tensor get_mn_major_tma_aligned_packed_ue8m0_tensor( + const torch::Tensor& sf, const c10::optional& psum_layout) { + return ::deep_gemm::get_mn_major_tma_aligned_packed_ue8m0_tensor(sf, psum_layout); +} + +static torch::Tensor get_k_grouped_mn_major_tma_aligned_packed_ue8m0_tensor( + const torch::Tensor& sf, const torch::Tensor& grouped_layout, + const c10::optional>& ks_cpu, + const int64_t& gran_k, const int64_t& k_alignment, + const bool& use_psum_layout) { + return ::deep_gemm::get_k_grouped_mn_major_tma_aligned_packed_ue8m0_tensor( + sf, grouped_layout, + list_to_optional_vector_int(ks_cpu), + static_cast(gran_k), static_cast(k_alignment), + use_psum_layout); +} +#endif + +static int64_t get_tma_aligned_size(const int64_t& x, const int64_t& element_size) { + return ::deep_gemm::get_tma_aligned_size(static_cast(x), static_cast(element_size)); +} + +static void set_mk_alignment_for_contiguous_layout(const int64_t& new_value) { + heuristics_runtime->set_mk_alignment_for_contiguous_layout(static_cast(new_value)); +} + +static int64_t get_mk_alignment_for_contiguous_layout() { + return heuristics_runtime->get_mk_alignment_for_contiguous_layout(); +} + +static int64_t get_theoretical_mk_alignment_for_contiguous_layout( + const c10::optional& expected_m, + const c10::optional& num_groups) { + return HeuristicsRuntime::get_theoretical_mk_alignment_for_contiguous_layout( + expected_m.has_value() ? std::make_optional(static_cast(expected_m.value())) : std::nullopt, + num_groups.has_value() ? std::make_optional(static_cast(num_groups.value())) : std::nullopt); +} + +} // namespace deep_gemm::torch_registration + +TORCH_LIBRARY_FRAGMENT(deep_gemm, m) { #if DG_TENSORMAP_COMPATIBLE - m.def("transform_sf_into_required_layout", &transform_sf_into_required_layout, - py::arg("sf"), py::arg("mn"), py::arg("k"), py::arg("recipe"), - py::arg("num_groups") = std::nullopt, - py::arg("is_sfa") = std::nullopt, - py::arg("disable_ue8m0_cast") = false, - py::arg("psum_layout") = std::nullopt); - - m.def("get_tma_aligned_size", &get_tma_aligned_size); - m.def("get_mn_major_tma_aligned_tensor", &get_mn_major_tma_aligned_tensor); - m.def("get_mn_major_tma_aligned_packed_ue8m0_tensor", &get_mn_major_tma_aligned_packed_ue8m0_tensor, - py::arg("sf"), py::arg("psum_layout") = std::nullopt); - m.def("get_k_grouped_mn_major_tma_aligned_packed_ue8m0_tensor", &get_k_grouped_mn_major_tma_aligned_packed_ue8m0_tensor, - py::arg("sf"), py::arg("grouped_layout"), py::arg("ks_cpu"), py::arg("gran_k"), py::arg("k_alignment"), - py::arg("use_psum_layout") = false); + m.def( + "transform_sf_into_required_layout(Tensor sf, int mn, int k, int[] recipe, int? num_groups=None, bool? is_sfa=None, bool disable_ue8m0_cast=False, Tensor? psum_layout=None) -> Tensor"); + m.def("get_tma_aligned_size(int x, int element_size) -> int", TORCH_FN(deep_gemm::torch_registration::get_tma_aligned_size)); + m.def("get_mn_major_tma_aligned_tensor(Tensor sf) -> Tensor"); + m.def( + "get_mn_major_tma_aligned_packed_ue8m0_tensor(Tensor sf, Tensor? psum_layout=None) -> Tensor"); + m.def( + "get_k_grouped_mn_major_tma_aligned_packed_ue8m0_tensor(Tensor sf, Tensor grouped_layout, int[]? ks_cpu, int gran_k, int k_alignment, bool use_psum_layout=False) -> Tensor"); #endif - m.def("set_mk_alignment_for_contiguous_layout", [&](const int& new_value) { - heuristics_runtime->set_mk_alignment_for_contiguous_layout(new_value); - }); - m.def("get_mk_alignment_for_contiguous_layout", [&]() { - return heuristics_runtime->get_mk_alignment_for_contiguous_layout(); - }); - m.def("get_theoretical_mk_alignment_for_contiguous_layout", [&](const std::optional& expected_m, - const std::optional& num_groups) { - return heuristics_runtime->get_theoretical_mk_alignment_for_contiguous_layout(expected_m, num_groups); - }, py::arg("expected_m") = std::nullopt, py::arg("num_groups") = std::nullopt); + m.def("set_mk_alignment_for_contiguous_layout(int new_value) -> ()", + TORCH_FN(deep_gemm::torch_registration::set_mk_alignment_for_contiguous_layout)); + m.def("get_mk_alignment_for_contiguous_layout() -> int", + TORCH_FN(deep_gemm::torch_registration::get_mk_alignment_for_contiguous_layout)); + m.def("get_theoretical_mk_alignment_for_contiguous_layout(int? expected_m=None, int? num_groups=None) -> int", + TORCH_FN(deep_gemm::torch_registration::get_theoretical_mk_alignment_for_contiguous_layout)); } -} // namespace deep_gemm::layout +TORCH_LIBRARY_IMPL(deep_gemm, CUDA, m) { + using namespace deep_gemm::torch_registration; + +#if DG_TENSORMAP_COMPATIBLE + m.impl("transform_sf_into_required_layout", + TORCH_FN(transform_sf_into_required_layout)); + m.impl("get_mn_major_tma_aligned_tensor", + TORCH_FN(get_mn_major_tma_aligned_tensor)); + m.impl("get_mn_major_tma_aligned_packed_ue8m0_tensor", + TORCH_FN(get_mn_major_tma_aligned_packed_ue8m0_tensor)); + m.impl("get_k_grouped_mn_major_tma_aligned_packed_ue8m0_tensor", + TORCH_FN(get_k_grouped_mn_major_tma_aligned_packed_ue8m0_tensor)); +#endif +} diff --git a/csrc/apis/mega.hpp b/csrc/apis/mega.hpp index b4ab0b3f0b..1bcafde4d7 100644 --- a/csrc/apis/mega.hpp +++ b/csrc/apis/mega.hpp @@ -1,11 +1,12 @@ #pragma once -#include #include -#include +#include +#include #include #include +#include "../utils/math.hpp" #if DG_TENSORMAP_COMPATIBLE #include "../jit/compiler.hpp" @@ -13,6 +14,8 @@ #include "../jit/device_runtime.hpp" #include "../jit_kernels/impls/sm100_bf16_mega_moe.hpp" #include "../jit_kernels/impls/sm100_fp8_fp4_mega_moe.hpp" +#include +#include "../torch_library_utils.hpp" namespace deep_gemm::mega { @@ -31,10 +34,71 @@ static int get_block_m_for_mega_moe( return block_m; } -static std::tuple(const torch::Tensor&)>> -get_symm_buffer_size_for_mega_moe( +struct SymmBufferLayoutInfo { + int64_t num_bytes = 0; + int64_t input_token_base = 0; + int64_t input_sf_base = 0; + int64_t input_topk_idx_base = 0; + int64_t input_topk_weights_base = 0; + int64_t shared_l1_sf_base = 0; + int64_t shared_l2_token_base = 0; + int64_t shared_l2_sf_base = 0; + int64_t l1_token_base = 0; + int64_t l1_sf_base = 0; + int64_t l2_token_base = 0; + int64_t l2_sf_base = 0; + bool with_sf = false; + int num_max_tokens_per_rank = 0; + int num_topk = 0; + int hidden = 0; + int intermediate_hidden = 0; + int num_shared_experts = 0; + int shared_intermediate_hidden = 0; + int num_ring_tokens = 0; + int num_sf_ring_tokens = 0; + + // Flatten into a plain `int[]` so it can cross the TORCH_LIBRARY boundary + std::vector to_int_list() const { + return { + num_bytes, input_token_base, input_sf_base, input_topk_idx_base, input_topk_weights_base, + shared_l1_sf_base, shared_l2_token_base, shared_l2_sf_base, + l1_token_base, l1_sf_base, l2_token_base, l2_sf_base, + static_cast(with_sf), num_max_tokens_per_rank, num_topk, + hidden, intermediate_hidden, num_shared_experts, shared_intermediate_hidden, + num_ring_tokens, num_sf_ring_tokens, + }; + } + + static SymmBufferLayoutInfo from_int_list(const std::vector& values) { + DG_HOST_ASSERT(static_cast(values.size()) == 21); + SymmBufferLayoutInfo info; + info.num_bytes = values[0]; + info.input_token_base = values[1]; + info.input_sf_base = values[2]; + info.input_topk_idx_base = values[3]; + info.input_topk_weights_base = values[4]; + info.shared_l1_sf_base = values[5]; + info.shared_l2_token_base = values[6]; + info.shared_l2_sf_base = values[7]; + info.l1_token_base = values[8]; + info.l1_sf_base = values[9]; + info.l2_token_base = values[10]; + info.l2_sf_base = values[11]; + // `with_sf` is a bool, encoded as 0/1 since the list is all `int64_t`. + info.with_sf = values[12] != 0; + info.num_max_tokens_per_rank = static_cast(values[13]); + info.num_topk = static_cast(values[14]); + info.hidden = static_cast(values[15]); + info.intermediate_hidden = static_cast(values[16]); + info.num_shared_experts = static_cast(values[17]); + info.shared_intermediate_hidden = static_cast(values[18]); + info.num_ring_tokens = static_cast(values[19]); + info.num_sf_ring_tokens = static_cast(values[20]); + return info; + } +}; + +static SymmBufferLayoutInfo build_symm_buffer_layout( const int& num_ranks, const int& num_experts, const int& num_max_tokens_per_rank, const int& num_topk, const int& hidden, const int& intermediate_hidden, @@ -93,65 +157,104 @@ get_symm_buffer_size_for_mega_moe( DG_HOST_ASSERT(num_sf_ring_tokens % 4 == 0); } - // Slice function: creates tensor views from the raw buffer. + SymmBufferLayoutInfo layout_info; + layout_info.num_bytes = mega_buffer.get_num_bytes(); + layout_info.input_token_base = reinterpret_cast(mega_buffer.input_token_buffer.base); + layout_info.input_sf_base = reinterpret_cast(mega_buffer.input_sf_buffer.base); + layout_info.input_topk_idx_base = reinterpret_cast(mega_buffer.input_topk_idx_buffer.base); + layout_info.input_topk_weights_base = reinterpret_cast(mega_buffer.input_topk_weights_buffer.base); + layout_info.shared_l1_sf_base = reinterpret_cast(mega_buffer.shared_l1_sf_buffer.base); + layout_info.shared_l2_token_base = reinterpret_cast(mega_buffer.shared_l2_token_buffer.base); + layout_info.shared_l2_sf_base = reinterpret_cast(mega_buffer.shared_l2_sf_buffer.base); + layout_info.l1_token_base = reinterpret_cast(mega_buffer.l1_token_buffer.base); + layout_info.l1_sf_base = reinterpret_cast(mega_buffer.l1_sf_buffer.base); + layout_info.l2_token_base = reinterpret_cast(mega_buffer.l2_token_buffer.base); + layout_info.l2_sf_base = reinterpret_cast(mega_buffer.l2_sf_buffer.base); + layout_info.with_sf = with_sf; + layout_info.num_max_tokens_per_rank = num_max_tokens_per_rank; + layout_info.num_topk = num_topk; + layout_info.hidden = hidden; + layout_info.intermediate_hidden = intermediate_hidden; + layout_info.num_shared_experts = num_shared_experts; + layout_info.shared_intermediate_hidden = shared_intermediate_hidden; + layout_info.num_ring_tokens = num_ring_tokens; + layout_info.num_sf_ring_tokens = num_sf_ring_tokens; + return layout_info; +} + +static std::tuple> get_symm_buffer_size_for_mega_moe( + const int& num_ranks, const int& num_experts, + const int& num_max_tokens_per_rank, const int& num_topk, + const int& hidden, const int& intermediate_hidden, + const std::string& mma_type, const std::string& activation, + const int& num_shared_experts = 0) { + const auto layout_info = build_symm_buffer_layout( + num_ranks, num_experts, num_max_tokens_per_rank, num_topk, + hidden, intermediate_hidden, mma_type, activation, num_shared_experts); + return std::make_tuple(layout_info.num_bytes, layout_info.to_int_list()); +} + +using SymmBufferSlice = std::tuple; + +static SymmBufferSlice slice_symm_buffer_from_layout( + const torch::Tensor& buffer, const SymmBufferLayoutInfo& layout_info) { // NOTES: `x_sf` is K-major, while `l1_acts_sf` and `l2_acts_sf` are M-major - auto slice_input_buffers = [=](const torch::Tensor& buffer) { - auto x = torch::from_blob( - math::advance_ptr(buffer.data_ptr(), reinterpret_cast(mega_buffer.input_token_buffer.base)), - {num_max_tokens_per_rank, hidden}, - torch::TensorOptions().dtype(with_sf ? torch::kFloat8_e4m3fn : torch::kBFloat16).device(buffer.device())); - auto x_sf = with_sf ? torch::from_blob( - math::advance_ptr(buffer.data_ptr(), reinterpret_cast(mega_buffer.input_sf_buffer.base)), - {num_max_tokens_per_rank, hidden / 128}, - torch::TensorOptions().dtype(torch::kInt).device(buffer.device())) : torch::Tensor(); - auto topk_idx = torch::from_blob( - math::advance_ptr(buffer.data_ptr(), reinterpret_cast(mega_buffer.input_topk_idx_buffer.base)), - {num_max_tokens_per_rank, num_topk}, - torch::TensorOptions().dtype(torch::kInt64).device(buffer.device())); - auto topk_weights = torch::from_blob( - math::advance_ptr(buffer.data_ptr(), reinterpret_cast(mega_buffer.input_topk_weights_buffer.base)), - {num_max_tokens_per_rank, num_topk}, - torch::TensorOptions().dtype(torch::kFloat32).device(buffer.device())); - - auto shared_l1_acts = x; - auto shared_l1_acts_sf = (with_sf and num_shared_experts > 0) ? torch::from_blob( - math::advance_ptr(buffer.data_ptr(), reinterpret_cast(mega_buffer.shared_l1_sf_buffer.base)), - {layout::get_num_max_shared_sf_tokens(num_max_tokens_per_rank), hidden / 128}, - {1, layout::get_num_max_shared_sf_tokens(num_max_tokens_per_rank)}, - torch::TensorOptions().dtype(torch::kInt).device(buffer.device())) : torch::Tensor(); - auto shared_l2_acts = num_shared_experts > 0 ? torch::from_blob( - math::advance_ptr(buffer.data_ptr(), reinterpret_cast(mega_buffer.shared_l2_token_buffer.base)), - {num_max_tokens_per_rank, shared_intermediate_hidden}, - torch::TensorOptions().dtype(with_sf ? torch::kFloat8_e4m3fn : torch::kBFloat16).device(buffer.device())) : torch::Tensor(); - auto shared_l2_acts_sf = (with_sf and num_shared_experts > 0) ? torch::from_blob( - math::advance_ptr(buffer.data_ptr(), reinterpret_cast(mega_buffer.shared_l2_sf_buffer.base)), - {layout::get_num_max_shared_sf_tokens(num_max_tokens_per_rank), shared_intermediate_hidden / 128}, - {1, layout::get_num_max_shared_sf_tokens(num_max_tokens_per_rank)}, - torch::TensorOptions().dtype(torch::kInt).device(buffer.device())) : torch::Tensor(); - - auto l1_acts = torch::from_blob( - math::advance_ptr(buffer.data_ptr(), reinterpret_cast(mega_buffer.l1_token_buffer.base)), - {num_ring_tokens, hidden}, - torch::TensorOptions().dtype(with_sf ? torch::kFloat8_e4m3fn : torch::kBFloat16).device(buffer.device())); - auto l1_acts_sf = with_sf ? torch::from_blob( - math::advance_ptr(buffer.data_ptr(), reinterpret_cast(mega_buffer.l1_sf_buffer.base)), - {num_sf_ring_tokens, hidden / 128}, - {1, num_sf_ring_tokens}, - torch::TensorOptions().dtype(torch::kInt).device(buffer.device())) : torch::Tensor(); - auto l2_acts = torch::from_blob( - math::advance_ptr(buffer.data_ptr(), reinterpret_cast(mega_buffer.l2_token_buffer.base)), - {num_ring_tokens, intermediate_hidden}, - torch::TensorOptions().dtype(with_sf ? torch::kFloat8_e4m3fn : torch::kBFloat16).device(buffer.device())); - auto l2_acts_sf = with_sf ? torch::from_blob( - math::advance_ptr(buffer.data_ptr(), reinterpret_cast(mega_buffer.l2_sf_buffer.base)), - {num_sf_ring_tokens, intermediate_hidden / 128}, - {1, num_sf_ring_tokens}, - torch::TensorOptions().dtype(torch::kInt).device(buffer.device())) : torch::Tensor(); - return std::make_tuple(x, x_sf, topk_idx, topk_weights, - shared_l1_acts, shared_l1_acts_sf, shared_l2_acts, shared_l2_acts_sf, - l1_acts, l1_acts_sf, l2_acts, l2_acts_sf); - }; - return {mega_buffer.get_num_bytes(), slice_input_buffers}; + auto x = torch::from_blob( + math::advance_ptr(buffer.data_ptr(), layout_info.input_token_base), + {layout_info.num_max_tokens_per_rank, layout_info.hidden}, + torch::TensorOptions().dtype(layout_info.with_sf ? torch::kFloat8_e4m3fn : torch::kBFloat16).device(buffer.device())); + auto x_sf = layout_info.with_sf ? torch::from_blob( + math::advance_ptr(buffer.data_ptr(), layout_info.input_sf_base), + {layout_info.num_max_tokens_per_rank, layout_info.hidden / 128}, + torch::TensorOptions().dtype(torch::kInt).device(buffer.device())) : torch::Tensor(); + auto topk_idx = torch::from_blob( + math::advance_ptr(buffer.data_ptr(), layout_info.input_topk_idx_base), + {layout_info.num_max_tokens_per_rank, layout_info.num_topk}, + torch::TensorOptions().dtype(torch::kInt64).device(buffer.device())); + auto topk_weights = torch::from_blob( + math::advance_ptr(buffer.data_ptr(), layout_info.input_topk_weights_base), + {layout_info.num_max_tokens_per_rank, layout_info.num_topk}, + torch::TensorOptions().dtype(torch::kFloat32).device(buffer.device())); + + auto shared_l1_acts = x; + auto shared_l1_acts_sf = (layout_info.with_sf and layout_info.num_shared_experts > 0) ? torch::from_blob( + math::advance_ptr(buffer.data_ptr(), layout_info.shared_l1_sf_base), + {layout::get_num_max_shared_sf_tokens(layout_info.num_max_tokens_per_rank), layout_info.hidden / 128}, + {1, layout::get_num_max_shared_sf_tokens(layout_info.num_max_tokens_per_rank)}, + torch::TensorOptions().dtype(torch::kInt).device(buffer.device())) : torch::Tensor(); + auto shared_l2_acts = layout_info.num_shared_experts > 0 ? torch::from_blob( + math::advance_ptr(buffer.data_ptr(), layout_info.shared_l2_token_base), + {layout_info.num_max_tokens_per_rank, layout_info.shared_intermediate_hidden}, + torch::TensorOptions().dtype(layout_info.with_sf ? torch::kFloat8_e4m3fn : torch::kBFloat16).device(buffer.device())) : torch::Tensor(); + auto shared_l2_acts_sf = (layout_info.with_sf and layout_info.num_shared_experts > 0) ? torch::from_blob( + math::advance_ptr(buffer.data_ptr(), layout_info.shared_l2_sf_base), + {layout::get_num_max_shared_sf_tokens(layout_info.num_max_tokens_per_rank), layout_info.shared_intermediate_hidden / 128}, + {1, layout::get_num_max_shared_sf_tokens(layout_info.num_max_tokens_per_rank)}, + torch::TensorOptions().dtype(torch::kInt).device(buffer.device())) : torch::Tensor(); + + auto l1_acts = torch::from_blob( + math::advance_ptr(buffer.data_ptr(), layout_info.l1_token_base), + {layout_info.num_ring_tokens, layout_info.hidden}, + torch::TensorOptions().dtype(layout_info.with_sf ? torch::kFloat8_e4m3fn : torch::kBFloat16).device(buffer.device())); + auto l1_acts_sf = layout_info.with_sf ? torch::from_blob( + math::advance_ptr(buffer.data_ptr(), layout_info.l1_sf_base), + {layout_info.num_sf_ring_tokens, layout_info.hidden / 128}, + {1, layout_info.num_sf_ring_tokens}, + torch::TensorOptions().dtype(torch::kInt).device(buffer.device())) : torch::Tensor(); + auto l2_acts = torch::from_blob( + math::advance_ptr(buffer.data_ptr(), layout_info.l2_token_base), + {layout_info.num_ring_tokens, layout_info.intermediate_hidden}, + torch::TensorOptions().dtype(layout_info.with_sf ? torch::kFloat8_e4m3fn : torch::kBFloat16).device(buffer.device())); + auto l2_acts_sf = layout_info.with_sf ? torch::from_blob( + math::advance_ptr(buffer.data_ptr(), layout_info.l2_sf_base), + {layout_info.num_sf_ring_tokens, layout_info.intermediate_hidden / 128}, + {1, layout_info.num_sf_ring_tokens}, + torch::TensorOptions().dtype(torch::kInt).device(buffer.device())) : torch::Tensor(); + return std::make_tuple(x, x_sf, topk_idx, topk_weights, + shared_l1_acts, shared_l1_acts_sf, shared_l2_acts, shared_l2_acts_sf, + l1_acts, l1_acts_sf, l2_acts, l2_acts_sf); } static void fp8_fp4_mega_moe( @@ -239,22 +342,22 @@ static void fp8_fp4_mega_moe( DG_HOST_ASSERT(cumulative_local_expert_recv_stats->is_contiguous()); } - // Check buffer bytes + // Check buffer bytes and slice views from one shared layout plan. const auto num_ranks = static_cast(sym_buffer_ptrs.size()); const auto num_experts_ = num_experts_per_rank * num_ranks; - const auto [num_required_bytes, slice] = get_symm_buffer_size_for_mega_moe( + const auto layout_info = build_symm_buffer_layout( num_ranks, num_experts, num_max_tokens_per_rank, num_topk, hidden, intermediate_hidden, "fp8xfp4", activation, num_shared_experts ); - DG_HOST_ASSERT(sym_buffer.nbytes() >= static_cast(num_required_bytes)); + DG_HOST_ASSERT(sym_buffer.nbytes() >= static_cast(layout_info.num_bytes)); DG_HOST_ASSERT(num_experts == num_experts_); - // Already registered tensors const auto [x, x_sf, topk_idx, topk_weights, shared_l1_acts, shared_l1_acts_sf, shared_l2_acts, shared_l2_acts_sf, - l1_acts, l1_acts_sf, l2_acts, l2_acts_sf] = slice(sym_buffer); + l1_acts, l1_acts_sf, l2_acts, l2_acts_sf] = + slice_symm_buffer_from_layout(sym_buffer, layout_info); // Dispatch into different architectures if (arch_major == 10) { @@ -351,22 +454,22 @@ static void bf16_mega_moe( DG_HOST_ASSERT(cumulative_local_expert_recv_stats->is_contiguous()); } - // Check buffer bytes + // Check buffer bytes and slice views from one shared layout plan. const auto num_ranks = static_cast(sym_buffer_ptrs.size()); const auto num_experts_ = num_experts_per_rank * num_ranks; - const auto [num_required_bytes, slice] = get_symm_buffer_size_for_mega_moe( + const auto layout_info = build_symm_buffer_layout( num_ranks, num_experts, num_max_tokens_per_rank, num_topk, hidden, intermediate_hidden, "bf16xbf16", activation, num_shared_experts ); - DG_HOST_ASSERT(sym_buffer.nbytes() >= static_cast(num_required_bytes)); + DG_HOST_ASSERT(sym_buffer.nbytes() >= static_cast(layout_info.num_bytes)); DG_HOST_ASSERT(num_experts == num_experts_); - // Already registered tensors const auto [x, _x_sf, topk_idx, topk_weights, shared_l1_acts, _shared_l1_acts_sf, shared_l2_acts, _shared_l2_acts_sf, - l1_acts, _l1_acts_sf, l2_acts, _l2_acts_sf] = slice(sym_buffer); + l1_acts, _l1_acts_sf, l2_acts, _l2_acts_sf] = + slice_symm_buffer_from_layout(sym_buffer, layout_info); // Dispatch into different architectures if (arch_major == 10) { @@ -393,14 +496,154 @@ static void bf16_mega_moe( sym_buffer.zero_(); } -static void register_apis(pybind11::module_& m) { +} // namespace deep_gemm::mega + +namespace deep_gemm::torch_registration { + +using namespace deep_gemm::torch_utils; + +static int64_t get_token_alignment_for_mega_moe() { + return static_cast(mega::get_token_alignment_for_mega_moe()); +} + +static int64_t get_block_m_for_mega_moe( + const int64_t& num_ranks, const int64_t& num_experts, + const int64_t& num_max_tokens_per_rank, const int64_t& num_tokens, + const int64_t& num_topk, const std::string& mma_type) { + return static_cast(mega::get_block_m_for_mega_moe( + static_cast(num_ranks), static_cast(num_experts), + static_cast(num_max_tokens_per_rank), static_cast(num_tokens), + static_cast(num_topk), mma_type)); +} + +static std::tuple> get_symm_buffer_size_for_mega_moe( + const int64_t& num_ranks, const int64_t& num_experts, + const int64_t& num_max_tokens_per_rank, const int64_t& num_topk, + const int64_t& hidden, const int64_t& intermediate_hidden, + const std::string& mma_type, const std::string& activation, + const int64_t& num_shared_experts) { + return mega::get_symm_buffer_size_for_mega_moe( + static_cast(num_ranks), static_cast(num_experts), + static_cast(num_max_tokens_per_rank), static_cast(num_topk), + static_cast(hidden), static_cast(intermediate_hidden), + mma_type, activation, static_cast(num_shared_experts)); +} + +static mega::SymmBufferSlice _slice_symm_buffer_for_mega_moe( + const torch::Tensor& buffer, + const std::vector& layout_info) { + return mega::slice_symm_buffer_from_layout( + buffer, mega::SymmBufferLayoutInfo::from_int_list(layout_info)); +} + +static void fp8_fp4_mega_moe( + const torch::Tensor& y, + const torch::Tensor& l1_weights, const torch::Tensor& l1_weights_sf, + const torch::Tensor& l2_weights, const torch::Tensor& l2_weights_sf, + const c10::optional& shared_l1_weights, + const c10::optional& shared_l1_weights_sf, + const c10::optional& shared_l2_weights, + const c10::optional& shared_l2_weights_sf, + const c10::optional& cumulative_local_expert_recv_stats, + const torch::Tensor& sym_buffer, + const std::vector& sym_buffer_ptrs, + const int64_t& rank_idx, + const int64_t& num_max_tokens_per_rank, + const int64_t& num_experts, const int64_t& num_topk, + const std::vector& recipe, + const std::string& activation, + const c10::optional& activation_clamp, + const bool& fast_math) { + std::optional> shared_l1_opt = std::nullopt; + std::optional> shared_l2_opt = std::nullopt; + if (shared_l1_weights.has_value()) { + DG_HOST_ASSERT(shared_l1_weights_sf.has_value() and shared_l2_weights.has_value() and shared_l2_weights_sf.has_value()); + shared_l1_opt = std::make_tuple(shared_l1_weights.value(), shared_l1_weights_sf.value()); + shared_l2_opt = std::make_tuple(shared_l2_weights.value(), shared_l2_weights_sf.value()); + } else { + DG_HOST_ASSERT(not shared_l1_weights_sf.has_value() and not shared_l2_weights.has_value() and not shared_l2_weights_sf.has_value()); + } + + mega::fp8_fp4_mega_moe( + y, + std::make_tuple(l1_weights, l1_weights_sf), + std::make_tuple(l2_weights, l2_weights_sf), + shared_l1_opt, + shared_l2_opt, + cumulative_local_expert_recv_stats, + sym_buffer, + sym_buffer_ptrs, + static_cast(rank_idx), + static_cast(num_max_tokens_per_rank), + static_cast(num_experts), static_cast(num_topk), + list_to_tuple3(recipe), + activation, + activation_clamp.has_value() + ? std::make_optional(static_cast(activation_clamp.value())) + : std::nullopt, + fast_math); +} + +static void bf16_mega_moe( + const torch::Tensor& y, + const torch::Tensor& l1_weights, + const torch::Tensor& l2_weights, + const c10::optional& shared_l1_weights, + const c10::optional& shared_l2_weights, + const c10::optional& cumulative_local_expert_recv_stats, + const torch::Tensor& sym_buffer, + const std::vector& sym_buffer_ptrs, + const int64_t& rank_idx, + const int64_t& num_max_tokens_per_rank, + const int64_t& num_experts, const int64_t& num_topk, + const std::string& activation, + const c10::optional& activation_clamp, + const bool& fast_math) { + mega::bf16_mega_moe( + y, l1_weights, l2_weights, + shared_l1_weights, + shared_l2_weights, + cumulative_local_expert_recv_stats, + sym_buffer, + sym_buffer_ptrs, + static_cast(rank_idx), + static_cast(num_max_tokens_per_rank), + static_cast(num_experts), static_cast(num_topk), + activation, + activation_clamp.has_value() + ? std::make_optional(static_cast(activation_clamp.value())) + : std::nullopt, + fast_math); +} + +} // namespace deep_gemm::torch_registration + +TORCH_LIBRARY_FRAGMENT(deep_gemm, m) { #if DG_TENSORMAP_COMPATIBLE - m.def("get_token_alignment_for_mega_moe", &get_token_alignment_for_mega_moe); - m.def("get_block_m_for_mega_moe", &get_block_m_for_mega_moe); - m.def("get_symm_buffer_size_for_mega_moe", &get_symm_buffer_size_for_mega_moe); - m.def("fp8_fp4_mega_moe", &fp8_fp4_mega_moe); - m.def("bf16_mega_moe", &bf16_mega_moe); + m.def( + "get_token_alignment_for_mega_moe() -> int", + TORCH_FN(deep_gemm::torch_registration::get_token_alignment_for_mega_moe)); + m.def( + "get_block_m_for_mega_moe(int num_ranks, int num_experts, int num_max_tokens_per_rank, int num_tokens, int num_topk, str mma_type) -> int", + TORCH_FN(deep_gemm::torch_registration::get_block_m_for_mega_moe)); + m.def( + "get_symm_buffer_size_for_mega_moe(int num_ranks, int num_experts, int num_max_tokens_per_rank, int num_topk, int hidden, int intermediate_hidden, str mma_type, str activation, int num_shared_experts=0) -> (int, int[])", + TORCH_FN(deep_gemm::torch_registration::get_symm_buffer_size_for_mega_moe)); + m.def( + "_slice_symm_buffer_for_mega_moe(Tensor buffer, int[] layout_info) -> (Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor)"); + m.def( + "fp8_fp4_mega_moe(Tensor(y!) y, Tensor l1_weights, Tensor l1_weights_sf, Tensor l2_weights, Tensor l2_weights_sf, Tensor? shared_l1_weights, Tensor? shared_l1_weights_sf, Tensor? shared_l2_weights, Tensor? shared_l2_weights_sf, Tensor(cumulative_local_expert_recv_stats!)? cumulative_local_expert_recv_stats, Tensor(sym_buffer!) sym_buffer, int[] sym_buffer_ptrs, int rank_idx, int num_max_tokens_per_rank, int num_experts, int num_topk, int[3] recipe, str activation, float? activation_clamp, bool fast_math) -> ()"); + m.def( + "bf16_mega_moe(Tensor(y!) y, Tensor l1_weights, Tensor l2_weights, Tensor? shared_l1_weights, Tensor? shared_l2_weights, Tensor(cumulative_local_expert_recv_stats!)? cumulative_local_expert_recv_stats, Tensor(sym_buffer!) sym_buffer, int[] sym_buffer_ptrs, int rank_idx, int num_max_tokens_per_rank, int num_experts, int num_topk, str activation, float? activation_clamp, bool fast_math) -> ()"); #endif } -} // namespace deep_gemm::mega +TORCH_LIBRARY_IMPL(deep_gemm, CUDA, m) { + using namespace deep_gemm::torch_registration; + +#if DG_TENSORMAP_COMPATIBLE + m.impl("_slice_symm_buffer_for_mega_moe", TORCH_FN(_slice_symm_buffer_for_mega_moe)); + m.impl("fp8_fp4_mega_moe", TORCH_FN(fp8_fp4_mega_moe)); + m.impl("bf16_mega_moe", TORCH_FN(bf16_mega_moe)); +#endif +} diff --git a/csrc/apis/runtime.hpp b/csrc/apis/runtime.hpp index 58fef941b7..c97bd3f6b4 100644 --- a/csrc/apis/runtime.hpp +++ b/csrc/apis/runtime.hpp @@ -2,50 +2,73 @@ #if DG_TENSORMAP_COMPATIBLE #include "../jit/compiler.hpp" +#include "../jit/kernel_runtime.hpp" #endif #include "../jit/device_runtime.hpp" #include "../jit_kernels/heuristics/runtime.hpp" -namespace deep_gemm::runtime { - -static void register_apis(pybind11::module_& m) { - m.def("set_num_sms", [&](const int& new_num_sms) { - device_runtime->set_num_sms(new_num_sms); - }); - m.def("get_num_sms", [&]() { - return device_runtime->get_num_sms(); - }); - m.def("set_tc_util", [&](const int& new_tc_util) { - device_runtime->set_tc_util(new_tc_util); - }); - m.def("get_tc_util", [&]() { - return device_runtime->get_tc_util(); - }); - m.def("set_pdl", [&](const bool& new_enable_pdl) { - device_runtime->set_pdl(new_enable_pdl); - }); - m.def("get_pdl", [&]() { - return device_runtime->get_pdl(); - }); - m.def("set_ignore_compile_dims", [&](const bool& new_value) { - heuristics_runtime->set_ignore_compile_dims(new_value); - }); - m.def("set_block_size_multiple_of", [&](const std::variant>& new_value) { - if (std::holds_alternative(new_value)) { - auto x = std::get(new_value); - heuristics_runtime->set_block_size_multiple_of(x, x); - } else { - auto [x, y] = std::get>(new_value); - heuristics_runtime->set_block_size_multiple_of(x, y); - } - }); - m.def("init", [&](const std::string& library_root_path, const std::string& cuda_home_path_by_python) { +#include + +namespace deep_gemm::torch_registration { + +static void set_num_sms(const int64_t& new_num_sms) { + device_runtime->set_num_sms(static_cast(new_num_sms)); +} + +static int64_t get_num_sms() { + return device_runtime->get_num_sms(); +} + +static void set_tc_util(const int64_t& new_tc_util) { + device_runtime->set_tc_util(static_cast(new_tc_util)); +} + +static int64_t get_tc_util() { + return device_runtime->get_tc_util(); +} + +static void set_pdl(const bool& new_enable_pdl) { + device_runtime->set_pdl(new_enable_pdl); +} + +static bool get_pdl() { + return device_runtime->get_pdl(); +} + +static void set_ignore_compile_dims(const bool& new_value) { + heuristics_runtime->set_ignore_compile_dims(new_value); +} + +static void set_block_size_multiple_of(const std::vector& value) { + if (value.size() == 1) { + const int v = static_cast(value[0]); + heuristics_runtime->set_block_size_multiple_of(v, v); + } else { + DG_HOST_ASSERT(value.size() == 2); + heuristics_runtime->set_block_size_multiple_of( + static_cast(value[0]), static_cast(value[1])); + } +} + +static void init(const std::string& library_root_path, + const std::string& cuda_home_path_by_python) { #if DG_TENSORMAP_COMPATIBLE Compiler::prepare_init(library_root_path, cuda_home_path_by_python); KernelRuntime::prepare_init(cuda_home_path_by_python); IncludeParser::prepare_init(library_root_path); #endif - }); } -} // namespace deep_gemm::runtime +} // namespace deep_gemm::torch_registration + +TORCH_LIBRARY_FRAGMENT(deep_gemm, m) { + m.def("set_num_sms(int new_num_sms) -> ()", TORCH_FN(deep_gemm::torch_registration::set_num_sms)); + m.def("get_num_sms() -> int", TORCH_FN(deep_gemm::torch_registration::get_num_sms)); + m.def("set_tc_util(int new_tc_util) -> ()", TORCH_FN(deep_gemm::torch_registration::set_tc_util)); + m.def("get_tc_util() -> int", TORCH_FN(deep_gemm::torch_registration::get_tc_util)); + m.def("set_pdl(bool new_enable_pdl) -> ()", TORCH_FN(deep_gemm::torch_registration::set_pdl)); + m.def("get_pdl() -> bool", TORCH_FN(deep_gemm::torch_registration::get_pdl)); + m.def("set_ignore_compile_dims(bool new_value) -> ()", TORCH_FN(deep_gemm::torch_registration::set_ignore_compile_dims)); + m.def("set_block_size_multiple_of(int[] value) -> ()", TORCH_FN(deep_gemm::torch_registration::set_block_size_multiple_of)); + m.def("init(str library_root_path, str cuda_home_path_by_python) -> ()", TORCH_FN(deep_gemm::torch_registration::init)); +} diff --git a/csrc/jit/device_runtime.hpp b/csrc/jit/device_runtime.hpp index d433558bfc..31b08b7a8f 100644 --- a/csrc/jit/device_runtime.hpp +++ b/csrc/jit/device_runtime.hpp @@ -4,6 +4,8 @@ #include #include +#include + #include "../utils/exception.hpp" #include "../utils/lazy_init.hpp" diff --git a/csrc/jit_kernels/impls/runtime_utils.hpp b/csrc/jit_kernels/impls/runtime_utils.hpp index 6a1e0684f6..6b8dcb6866 100644 --- a/csrc/jit_kernels/impls/runtime_utils.hpp +++ b/csrc/jit_kernels/impls/runtime_utils.hpp @@ -1,7 +1,7 @@ #pragma once #include -#include +#include #include "../heuristics/sm90.hpp" #include "../../jit/handle.hpp" diff --git a/csrc/jit_kernels/impls/sm100_bf16_gemm.hpp b/csrc/jit_kernels/impls/sm100_bf16_gemm.hpp index f9b2f361cb..9129b7bcc8 100644 --- a/csrc/jit_kernels/impls/sm100_bf16_gemm.hpp +++ b/csrc/jit_kernels/impls/sm100_bf16_gemm.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/device_runtime.hpp" diff --git a/csrc/jit_kernels/impls/sm100_bf16_mega_moe.hpp b/csrc/jit_kernels/impls/sm100_bf16_mega_moe.hpp index 273874cfd9..47ad292bc2 100644 --- a/csrc/jit_kernels/impls/sm100_bf16_mega_moe.hpp +++ b/csrc/jit_kernels/impls/sm100_bf16_mega_moe.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/kernel_runtime.hpp" diff --git a/csrc/jit_kernels/impls/sm100_bmk_bnk_mn.hpp b/csrc/jit_kernels/impls/sm100_bmk_bnk_mn.hpp index 65c9d501c2..292bd3903e 100644 --- a/csrc/jit_kernels/impls/sm100_bmk_bnk_mn.hpp +++ b/csrc/jit_kernels/impls/sm100_bmk_bnk_mn.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/device_runtime.hpp" diff --git a/csrc/jit_kernels/impls/sm100_fp8_fp4_gemm_1d1d.hpp b/csrc/jit_kernels/impls/sm100_fp8_fp4_gemm_1d1d.hpp index 9e4ba58b60..95f855298c 100644 --- a/csrc/jit_kernels/impls/sm100_fp8_fp4_gemm_1d1d.hpp +++ b/csrc/jit_kernels/impls/sm100_fp8_fp4_gemm_1d1d.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/device_runtime.hpp" diff --git a/csrc/jit_kernels/impls/sm100_fp8_fp4_mega_moe.hpp b/csrc/jit_kernels/impls/sm100_fp8_fp4_mega_moe.hpp index 9c4be4b08d..8251b48fb4 100644 --- a/csrc/jit_kernels/impls/sm100_fp8_fp4_mega_moe.hpp +++ b/csrc/jit_kernels/impls/sm100_fp8_fp4_mega_moe.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/kernel_runtime.hpp" diff --git a/csrc/jit_kernels/impls/sm100_tf32_hc_prenorm_gemm.hpp b/csrc/jit_kernels/impls/sm100_tf32_hc_prenorm_gemm.hpp index 0071e2c57f..4ee309e67c 100644 --- a/csrc/jit_kernels/impls/sm100_tf32_hc_prenorm_gemm.hpp +++ b/csrc/jit_kernels/impls/sm100_tf32_hc_prenorm_gemm.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/device_runtime.hpp" diff --git a/csrc/jit_kernels/impls/sm120_bf16_gemm.hpp b/csrc/jit_kernels/impls/sm120_bf16_gemm.hpp index e272caf346..94be80b79b 100644 --- a/csrc/jit_kernels/impls/sm120_bf16_gemm.hpp +++ b/csrc/jit_kernels/impls/sm120_bf16_gemm.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/device_runtime.hpp" diff --git a/csrc/jit_kernels/impls/sm120_bmk_bnk_mn.hpp b/csrc/jit_kernels/impls/sm120_bmk_bnk_mn.hpp index 55365324c3..1472b26de9 100644 --- a/csrc/jit_kernels/impls/sm120_bmk_bnk_mn.hpp +++ b/csrc/jit_kernels/impls/sm120_bmk_bnk_mn.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/device_runtime.hpp" diff --git a/csrc/jit_kernels/impls/sm120_fp8_fp4_gemm_1d1d.hpp b/csrc/jit_kernels/impls/sm120_fp8_fp4_gemm_1d1d.hpp index c8fab5839e..47176295e5 100644 --- a/csrc/jit_kernels/impls/sm120_fp8_fp4_gemm_1d1d.hpp +++ b/csrc/jit_kernels/impls/sm120_fp8_fp4_gemm_1d1d.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/device_runtime.hpp" diff --git a/csrc/jit_kernels/impls/sm120_tf32_hc_prenorm_gemm.hpp b/csrc/jit_kernels/impls/sm120_tf32_hc_prenorm_gemm.hpp index 3067baca5a..f375ca75d6 100644 --- a/csrc/jit_kernels/impls/sm120_tf32_hc_prenorm_gemm.hpp +++ b/csrc/jit_kernels/impls/sm120_tf32_hc_prenorm_gemm.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/device_runtime.hpp" diff --git a/csrc/jit_kernels/impls/sm90_bf16_gemm.hpp b/csrc/jit_kernels/impls/sm90_bf16_gemm.hpp index 24edd46562..130dcc220b 100644 --- a/csrc/jit_kernels/impls/sm90_bf16_gemm.hpp +++ b/csrc/jit_kernels/impls/sm90_bf16_gemm.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/kernel_runtime.hpp" diff --git a/csrc/jit_kernels/impls/sm90_bmk_bnk_mn.hpp b/csrc/jit_kernels/impls/sm90_bmk_bnk_mn.hpp index 473677b70c..7aed97598b 100644 --- a/csrc/jit_kernels/impls/sm90_bmk_bnk_mn.hpp +++ b/csrc/jit_kernels/impls/sm90_bmk_bnk_mn.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/device_runtime.hpp" diff --git a/csrc/jit_kernels/impls/sm90_fp8_gemm_1d1d.hpp b/csrc/jit_kernels/impls/sm90_fp8_gemm_1d1d.hpp index 120b91faf8..1350f32f9a 100644 --- a/csrc/jit_kernels/impls/sm90_fp8_gemm_1d1d.hpp +++ b/csrc/jit_kernels/impls/sm90_fp8_gemm_1d1d.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/device_runtime.hpp" diff --git a/csrc/jit_kernels/impls/sm90_fp8_gemm_1d2d.hpp b/csrc/jit_kernels/impls/sm90_fp8_gemm_1d2d.hpp index 892edaee7d..daeece0326 100644 --- a/csrc/jit_kernels/impls/sm90_fp8_gemm_1d2d.hpp +++ b/csrc/jit_kernels/impls/sm90_fp8_gemm_1d2d.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/device_runtime.hpp" diff --git a/csrc/jit_kernels/impls/sm90_tf32_hc_prenorm_gemm.hpp b/csrc/jit_kernels/impls/sm90_tf32_hc_prenorm_gemm.hpp index c17d1b554e..7fbedf813f 100644 --- a/csrc/jit_kernels/impls/sm90_tf32_hc_prenorm_gemm.hpp +++ b/csrc/jit_kernels/impls/sm90_tf32_hc_prenorm_gemm.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/compiler.hpp" #include "../../jit/device_runtime.hpp" diff --git a/csrc/jit_kernels/impls/smxx_layout.hpp b/csrc/jit_kernels/impls/smxx_layout.hpp index 82de55e198..23dfb19c4c 100644 --- a/csrc/jit_kernels/impls/smxx_layout.hpp +++ b/csrc/jit_kernels/impls/smxx_layout.hpp @@ -1,6 +1,6 @@ #pragma once -#include +#include #include "../../jit/kernel_runtime.hpp" #include "../../jit/compiler.hpp" diff --git a/csrc/python_api.cpp b/csrc/python_api.cpp index a966afe1ed..260b545125 100644 --- a/csrc/python_api.cpp +++ b/csrc/python_api.cpp @@ -1,5 +1,5 @@ -#include -#include +#include +#include "utils/registration.h" #include "apis/attention.hpp" #include "apis/einsum.hpp" @@ -9,20 +9,4 @@ #include "apis/mega.hpp" #include "apis/runtime.hpp" -#ifndef TORCH_EXTENSION_NAME -#define TORCH_EXTENSION_NAME _C -#endif - -// ReSharper disable once CppParameterMayBeConstPtrOrRef -PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { - m.doc() = "DeepGEMM C++ library"; - - // TODO: make SM80 incompatible issues raise errors - deep_gemm::attention::register_apis(m); - deep_gemm::einsum::register_apis(m); - deep_gemm::hyperconnection::register_apis(m); - deep_gemm::gemm::register_apis(m); - deep_gemm::layout::register_apis(m); - deep_gemm::mega::register_apis(m); - deep_gemm::runtime::register_apis(m); -} +REGISTER_EXTENSION(_C_extension) diff --git a/csrc/torch_library_utils.hpp b/csrc/torch_library_utils.hpp new file mode 100644 index 0000000000..2d80758ba3 --- /dev/null +++ b/csrc/torch_library_utils.hpp @@ -0,0 +1,63 @@ +#pragma once + +#include +#include +#include +#include + +#include "utils/exception.hpp" + +namespace deep_gemm::torch_utils { + +inline std::optional> list_to_recipe3( + const c10::optional>& recipe) { + if (not recipe.has_value() or recipe->empty()) { + return std::nullopt; + } + DG_HOST_ASSERT(recipe->size() == 3); + return std::make_tuple(static_cast((*recipe)[0]), + static_cast((*recipe)[1]), + static_cast((*recipe)[2])); +} + +inline std::optional> list_to_recipe2( + const c10::optional>& recipe) { + if (not recipe.has_value() or recipe->empty()) { + return std::nullopt; + } + DG_HOST_ASSERT(recipe->size() == 2); + return std::make_tuple(static_cast((*recipe)[0]), static_cast((*recipe)[1])); +} + +inline std::variant, std::tuple> list_to_recipe_variant( + const std::vector& recipe) { + DG_HOST_ASSERT(recipe.size() == 2 or recipe.size() == 3); + if (recipe.size() == 2) { + return std::make_tuple(static_cast(recipe[0]), static_cast(recipe[1])); + } + return std::make_tuple(static_cast(recipe[0]), + static_cast(recipe[1]), + static_cast(recipe[2])); +} + +inline std::tuple list_to_tuple3(const std::vector& values) { + DG_HOST_ASSERT(values.size() == 3); + return std::make_tuple(static_cast(values[0]), + static_cast(values[1]), + static_cast(values[2])); +} + +inline std::optional> list_to_optional_vector_int( + const c10::optional>& values) { + if (not values.has_value()) { + return std::nullopt; + } + std::vector out; + out.reserve(values->size()); + for (const auto value : *values) { + out.push_back(static_cast(value)); + } + return out; +} + +} // namespace deep_gemm::torch_utils diff --git a/csrc/utils/layout.hpp b/csrc/utils/layout.hpp index 07a81c4e37..ae030571d3 100644 --- a/csrc/utils/layout.hpp +++ b/csrc/utils/layout.hpp @@ -1,7 +1,7 @@ #pragma once #include -#include +#include #include "math.hpp" #include "exception.hpp" diff --git a/csrc/utils/math.hpp b/csrc/utils/math.hpp index 0aa28eb400..81ecd427c5 100644 --- a/csrc/utils/math.hpp +++ b/csrc/utils/math.hpp @@ -1,7 +1,7 @@ // TODO: merge this file with `math.cuh` (the device part) #pragma once -#include +#include #include "exception.hpp" diff --git a/csrc/utils/registration.h b/csrc/utils/registration.h new file mode 100644 index 0000000000..f625822732 --- /dev/null +++ b/csrc/utils/registration.h @@ -0,0 +1,17 @@ +#pragma once + +#include + +#define _CONCAT(A, B) A##B +#define CONCAT(A, B) _CONCAT(A, B) + +#define _STRINGIFY(A) #A +#define STRINGIFY(A) _STRINGIFY(A) + +// Empty PyInit so the .so is importable; ops still register via TORCH_LIBRARY. +#define REGISTER_EXTENSION(NAME) \ + PyMODINIT_FUNC CONCAT(PyInit_, NAME)() { \ + static struct PyModuleDef module = {PyModuleDef_HEAD_INIT, \ + STRINGIFY(NAME), nullptr, 0, nullptr}; \ + return PyModule_Create(&module); \ + } diff --git a/deep_gemm/_C.py b/deep_gemm/_C.py new file mode 100644 index 0000000000..fb8960bb86 --- /dev/null +++ b/deep_gemm/_C.py @@ -0,0 +1,357 @@ +import torch +from pathlib import Path + + +def _load_extension(): + so_files = list(Path(__file__).parent.glob('_C_extension*.so')) + assert len(so_files) == 1, ( + f'Expected one _C_extension*.so file, found {len(so_files)}: {so_files}' + ) + torch.ops.load_library(str(so_files[0])) + + +_load_extension() +_torch_ops = torch.ops.deep_gemm + + +def _bind_guarded_ops(*names): + """Bind ops when all are registered (matches one C++ #if guard group).""" + present = [name for name in names if hasattr(_torch_ops, name)] + if not present: + return + assert len(present) == len(names), ( + f'Guard group mismatch: {sorted(set(names) - set(present))} missing while ' + f'{present} are registered — the C++ #if guards for these ops have diverged.' + ) + globals().update({name: getattr(_torch_ops, name) for name in names}) + + +init = _torch_ops.init +set_num_sms = _torch_ops.set_num_sms +get_num_sms = _torch_ops.get_num_sms +set_tc_util = _torch_ops.set_tc_util +get_tc_util = _torch_ops.get_tc_util +set_pdl = _torch_ops.set_pdl +get_pdl = _torch_ops.get_pdl +set_ignore_compile_dims = _torch_ops.set_ignore_compile_dims +get_mk_alignment_for_contiguous_layout = _torch_ops.get_mk_alignment_for_contiguous_layout +get_theoretical_mk_alignment_for_contiguous_layout = _torch_ops.get_theoretical_mk_alignment_for_contiguous_layout +cublaslt_gemm_nt = _torch_ops.cublaslt_gemm_nt +cublaslt_gemm_nn = _torch_ops.cublaslt_gemm_nn +cublaslt_gemm_tn = _torch_ops.cublaslt_gemm_tn +cublaslt_gemm_tt = _torch_ops.cublaslt_gemm_tt + + +def set_block_size_multiple_of(value): + if isinstance(value, int): + return _torch_ops.set_block_size_multiple_of([value]) + return _torch_ops.set_block_size_multiple_of(list(value)) + + +def set_mk_alignment_for_contiguous_layout(value): + return _torch_ops.set_mk_alignment_for_contiguous_layout(value) + + +def _unpack_ab_pair(a, b): + return a[0], a[1], b[0], b[1] + + +def _unpack_q(q): + if isinstance(q, tuple): + q_fp = q[0] + q_sf = q[1] if len(q) > 1 else None + else: + q_fp, q_sf = q, None + return q_fp, q_sf + + +def _unpack_kv(kv): + return kv[0], kv[1] + + +def _as_int_list(value): + """Reject a bare scalar instead of letting int[N] silently broadcast it into a list.""" + return None if value is None else list(value) + + +def _register_deep_gemm_kernels(): + """Export DeepGEMM kernels only when C++ ops are registered.""" + def fp8_fp4_gemm_nt(a, b, d, c=None, recipe=None, recipe_a=None, recipe_b=None, + compiled_dims='nk', disable_ue8m0_cast=False): + a_tensor, sfa, b_tensor, sfb = _unpack_ab_pair(a, b) + return _torch_ops.fp8_fp4_gemm_nt( + a_tensor, sfa, b_tensor, sfb, d, c, _as_int_list(recipe), _as_int_list(recipe_a), _as_int_list(recipe_b), + compiled_dims, disable_ue8m0_cast, + ) + + def fp8_fp4_gemm_nn(a, b, d, c=None, recipe=None, recipe_a=None, recipe_b=None, + compiled_dims='nk', disable_ue8m0_cast=False): + a_tensor, sfa, b_tensor, sfb = _unpack_ab_pair(a, b) + return _torch_ops.fp8_fp4_gemm_nn( + a_tensor, sfa, b_tensor, sfb, d, c, _as_int_list(recipe), _as_int_list(recipe_a), _as_int_list(recipe_b), + compiled_dims, disable_ue8m0_cast, + ) + + def fp8_fp4_gemm_tn(a, b, d, c=None, recipe=None, recipe_a=None, recipe_b=None, + compiled_dims='mn', disable_ue8m0_cast=False): + a_tensor, sfa, b_tensor, sfb = _unpack_ab_pair(a, b) + return _torch_ops.fp8_fp4_gemm_tn( + a_tensor, sfa, b_tensor, sfb, d, c, _as_int_list(recipe), _as_int_list(recipe_a), _as_int_list(recipe_b), + compiled_dims, disable_ue8m0_cast, + ) + + def fp8_fp4_gemm_tt(a, b, d, c=None, recipe=None, recipe_a=None, recipe_b=None, + compiled_dims='mn', disable_ue8m0_cast=False): + a_tensor, sfa, b_tensor, sfb = _unpack_ab_pair(a, b) + return _torch_ops.fp8_fp4_gemm_tt( + a_tensor, sfa, b_tensor, sfb, d, c, _as_int_list(recipe), _as_int_list(recipe_a), _as_int_list(recipe_b), + compiled_dims, disable_ue8m0_cast, + ) + + def m_grouped_fp8_fp4_gemm_nt_contiguous(a, b, d, grouped_layout, recipe=None, recipe_a=None, recipe_b=None, + compiled_dims='nk', disable_ue8m0_cast=False, use_psum_layout=False, + ensure_zero_padding=True, expected_m_for_psum_layout=None): + a_tensor, sfa, b_tensor, sfb = _unpack_ab_pair(a, b) + return _torch_ops.m_grouped_fp8_fp4_gemm_nt_contiguous( + a_tensor, sfa, b_tensor, sfb, d, grouped_layout, _as_int_list(recipe), _as_int_list(recipe_a), _as_int_list(recipe_b), + compiled_dims, disable_ue8m0_cast, use_psum_layout, ensure_zero_padding, + expected_m_for_psum_layout, + ) + + def m_grouped_fp8_fp4_gemm_nn_contiguous(a, b, d, grouped_layout, recipe=None, recipe_a=None, recipe_b=None, + compiled_dims='nk', disable_ue8m0_cast=False, use_psum_layout=False, + ensure_zero_padding=True): + a_tensor, sfa, b_tensor, sfb = _unpack_ab_pair(a, b) + return _torch_ops.m_grouped_fp8_fp4_gemm_nn_contiguous( + a_tensor, sfa, b_tensor, sfb, d, grouped_layout, _as_int_list(recipe), _as_int_list(recipe_a), _as_int_list(recipe_b), + compiled_dims, disable_ue8m0_cast, use_psum_layout, ensure_zero_padding, + ) + + def m_grouped_fp8_fp4_gemm_nt_masked(a, b, d, masked_m, expected_m, recipe=None, recipe_a=None, recipe_b=None, + compiled_dims='nk', disable_ue8m0_cast=False): + a_tensor, sfa, b_tensor, sfb = _unpack_ab_pair(a, b) + return _torch_ops.m_grouped_fp8_fp4_gemm_nt_masked( + a_tensor, sfa, b_tensor, sfb, d, masked_m, expected_m, _as_int_list(recipe), _as_int_list(recipe_a), _as_int_list(recipe_b), + compiled_dims, disable_ue8m0_cast, + ) + + def k_grouped_fp8_gemm_tn_contiguous(a, b, d, ks_cpu, grouped_layout, c=None, recipe=(1, 1, 128), + compiled_dims='mn', use_psum_layout=False): + a_tensor, sfa, b_tensor, sfb = _unpack_ab_pair(a, b) + return _torch_ops.k_grouped_fp8_gemm_tn_contiguous( + a_tensor, sfa, b_tensor, sfb, d, ks_cpu, grouped_layout, c, list(recipe), + compiled_dims, use_psum_layout, + ) + + def k_grouped_fp8_gemm_nt_contiguous(a, b, d, ks_cpu, grouped_layout, c=None, recipe=(1, 1, 128), + compiled_dims='mn', use_psum_layout=False): + a_tensor, sfa, b_tensor, sfb = _unpack_ab_pair(a, b) + return _torch_ops.k_grouped_fp8_gemm_nt_contiguous( + a_tensor, sfa, b_tensor, sfb, d, ks_cpu, grouped_layout, c, list(recipe), + compiled_dims, use_psum_layout, + ) + + def fp8_gemm_nt_skip_head_mid(a, b, d, head_splits, recipe=None, compiled_dims='nk', disable_ue8m0_cast=False): + a_tensor, sfa, b_tensor, sfb = _unpack_ab_pair(a, b) + return _torch_ops.fp8_gemm_nt_skip_head_mid( + a_tensor, sfa, b_tensor, sfb, d, list(head_splits), _as_int_list(recipe), compiled_dims, disable_ue8m0_cast, + ) + + def fp8_einsum(expr, a, b, d, c=None, recipe=(1, 128, 128)): + return _torch_ops.fp8_einsum(expr, a[0], a[1], b[0], b[1], d, c, list(recipe)) + + def fp8_fp4_mqa_logits(q, kv, weights, cu_seq_len_k_start, cu_seq_len_k_end, clean_logits=True, + max_seqlen_k=0, logits_dtype=torch.float32): + q_fp, q_sf = _unpack_q(q) + kv_fp, kv_sf = _unpack_kv(kv) + return _torch_ops.fp8_fp4_mqa_logits( + q_fp, q_sf, kv_fp, kv_sf, weights, cu_seq_len_k_start, cu_seq_len_k_end, + clean_logits, max_seqlen_k, logits_dtype, + ) + + def fp8_fp4_paged_mqa_logits(q, kv_cache, weights, context_lens, block_table, schedule_meta, max_context_len, + clean_logits=False, logits_dtype=torch.float32, indices=None): + q_fp, q_sf = _unpack_q(q) + return _torch_ops.fp8_fp4_paged_mqa_logits( + q_fp, q_sf, kv_cache, weights, context_lens, block_table, schedule_meta, max_context_len, + clean_logits, logits_dtype, indices, + ) + + def fp8_mqa_logits(q, kv, weights, cu_seq_len_k_start, cu_seq_len_k_end, clean_logits=True, max_seqlen_k=0): + kv_fp, kv_sf = _unpack_kv(kv) + return _torch_ops.fp8_mqa_logits(q, kv_fp, kv_sf, weights, cu_seq_len_k_start, cu_seq_len_k_end, clean_logits, max_seqlen_k) + + def fp8_paged_mqa_logits(q, kv_cache, weights, context_lens, block_table, schedule_meta, max_context_len, + clean_logits=False, indices=None): + return _torch_ops.fp8_paged_mqa_logits( + q, kv_cache, weights, context_lens, block_table, schedule_meta, max_context_len, clean_logits, indices, + ) + + globals().update({ + 'fp8_fp4_gemm_nt': fp8_fp4_gemm_nt, + 'fp8_fp4_gemm_nn': fp8_fp4_gemm_nn, + 'fp8_fp4_gemm_tn': fp8_fp4_gemm_tn, + 'fp8_fp4_gemm_tt': fp8_fp4_gemm_tt, + 'fp8_gemm_nt': fp8_fp4_gemm_nt, + 'fp8_gemm_nn': fp8_fp4_gemm_nn, + 'fp8_gemm_tn': fp8_fp4_gemm_tn, + 'fp8_gemm_tt': fp8_fp4_gemm_tt, + 'm_grouped_fp8_fp4_gemm_nt_contiguous': m_grouped_fp8_fp4_gemm_nt_contiguous, + 'm_grouped_fp8_fp4_gemm_nn_contiguous': m_grouped_fp8_fp4_gemm_nn_contiguous, + 'm_grouped_fp8_fp4_gemm_nt_masked': m_grouped_fp8_fp4_gemm_nt_masked, + 'm_grouped_fp8_gemm_nt_contiguous': m_grouped_fp8_fp4_gemm_nt_contiguous, + 'm_grouped_fp8_gemm_nn_contiguous': m_grouped_fp8_fp4_gemm_nn_contiguous, + 'm_grouped_fp8_gemm_nt_masked': m_grouped_fp8_fp4_gemm_nt_masked, + 'k_grouped_fp8_gemm_tn_contiguous': k_grouped_fp8_gemm_tn_contiguous, + 'k_grouped_fp8_gemm_nt_contiguous': k_grouped_fp8_gemm_nt_contiguous, + 'fp8_gemm_nt_skip_head_mid': fp8_gemm_nt_skip_head_mid, + 'fp8_einsum': fp8_einsum, + 'fp8_fp4_mqa_logits': fp8_fp4_mqa_logits, + 'fp8_fp4_paged_mqa_logits': fp8_fp4_paged_mqa_logits, + 'fp8_mqa_logits': fp8_mqa_logits, + 'fp8_paged_mqa_logits': fp8_paged_mqa_logits, + }) + + # DG_TENSORMAP_COMPATIBLE — gemm.hpp (BF16 impl conditional) + _bind_guarded_ops( + 'bf16_gemm_nt', + 'bf16_gemm_nn', + 'bf16_gemm_tn', + 'bf16_gemm_tt', + 'm_grouped_bf16_gemm_nt_contiguous', + 'm_grouped_bf16_gemm_nn_contiguous', + 'm_grouped_bf16_gemm_nt_masked', + 'k_grouped_bf16_gemm_tn_contiguous', + ) + + # DG_FP8_COMPATIBLE and DG_TENSORMAP_COMPATIBLE — grouped only because these three + # happen to share the same guard today; if any one's guard changes, split it out. + _bind_guarded_ops( + 'einsum', # einsum.hpp + 'tf32_hc_prenorm_gemm', # hyperconnection.hpp + 'get_paged_mqa_logits_metadata', # attention.hpp + ) + + # DG_TENSORMAP_COMPATIBLE — layout.hpp (schema and impl conditional) + _bind_guarded_ops( + 'transform_sf_into_required_layout', + 'get_tma_aligned_size', + 'get_mn_major_tma_aligned_tensor', + 'get_mn_major_tma_aligned_packed_ue8m0_tensor', + 'get_k_grouped_mn_major_tma_aligned_packed_ue8m0_tensor', + ) + + +_register_deep_gemm_kernels() + + +def get_symm_buffer_size_for_mega_moe(*args, **kwargs): + return _torch_ops.get_symm_buffer_size_for_mega_moe(*args, **kwargs) + + +def _slice_symm_buffer_for_mega_moe(buffer, *args, **kwargs): + return _torch_ops._slice_symm_buffer_for_mega_moe(buffer, *args, **kwargs) + + +# DG_TENSORMAP_COMPATIBLE — mega.hpp (C++ impl conditional; matches legacy pybind export guard) +_bind_guarded_ops( + 'get_token_alignment_for_mega_moe', + 'get_block_m_for_mega_moe', +) + + +def fp8_fp4_mega_moe(y, l1_weights, l2_weights, shared_l1_weights, shared_l2_weights, + cumulative_local_expert_recv_stats, sym_buffer, + sym_buffer_ptrs, rank_idx, num_max_tokens_per_rank, num_experts, num_topk, recipe, + activation, activation_clamp, fast_math): + shared_l1_w = shared_l1_sf = shared_l2_w = shared_l2_sf = None + if shared_l1_weights is not None: + shared_l1_w, shared_l1_sf = shared_l1_weights + shared_l2_w, shared_l2_sf = shared_l2_weights + return _torch_ops.fp8_fp4_mega_moe( + y, l1_weights[0], l1_weights[1], l2_weights[0], l2_weights[1], + shared_l1_w, shared_l1_sf, shared_l2_w, shared_l2_sf, + cumulative_local_expert_recv_stats, sym_buffer, list(sym_buffer_ptrs), rank_idx, + num_max_tokens_per_rank, num_experts, num_topk, list(recipe), activation, + activation_clamp, fast_math, + ) + + +def bf16_mega_moe(y, l1_weights, l2_weights, shared_l1_weights, shared_l2_weights, + cumulative_local_expert_recv_stats, sym_buffer, + sym_buffer_ptrs, rank_idx, num_max_tokens_per_rank, num_experts, num_topk, + activation, activation_clamp, fast_math): + return _torch_ops.bf16_mega_moe( + y, l1_weights, l2_weights, shared_l1_weights, shared_l2_weights, + cumulative_local_expert_recv_stats, sym_buffer, + list(sym_buffer_ptrs), rank_idx, num_max_tokens_per_rank, num_experts, num_topk, + activation, activation_clamp, fast_math, + ) + + +_UNCONDITIONAL_API = ( + # Runtime + 'init', + 'set_num_sms', 'get_num_sms', + 'set_tc_util', 'get_tc_util', + 'set_pdl', 'get_pdl', + 'set_ignore_compile_dims', + 'set_block_size_multiple_of', + 'set_mk_alignment_for_contiguous_layout', + 'get_mk_alignment_for_contiguous_layout', + 'get_theoretical_mk_alignment_for_contiguous_layout', + # cuBLASLt GEMMs + 'cublaslt_gemm_nt', 'cublaslt_gemm_nn', + 'cublaslt_gemm_tn', 'cublaslt_gemm_tt', + # Mega MoE (imported via deep_gemm.mega; always defined, fails at call if unregistered) + 'get_symm_buffer_size_for_mega_moe', + '_slice_symm_buffer_for_mega_moe', + 'fp8_fp4_mega_moe', + 'bf16_mega_moe', +) + +_DEEP_GEMM_API = ( + # FP8/FP4 GEMMs + 'fp8_fp4_gemm_nt', 'fp8_fp4_gemm_nn', + 'fp8_fp4_gemm_tn', 'fp8_fp4_gemm_tt', + 'fp8_gemm_nt', 'fp8_gemm_nn', + 'fp8_gemm_tn', 'fp8_gemm_tt', + 'fp8_gemm_nt_skip_head_mid', + 'm_grouped_fp8_fp4_gemm_nt_contiguous', + 'm_grouped_fp8_fp4_gemm_nn_contiguous', + 'm_grouped_fp8_fp4_gemm_nt_masked', + 'm_grouped_fp8_gemm_nt_contiguous', + 'm_grouped_fp8_gemm_nn_contiguous', + 'm_grouped_fp8_gemm_nt_masked', + 'k_grouped_fp8_gemm_tn_contiguous', + 'k_grouped_fp8_gemm_nt_contiguous', + # BF16 GEMMs (guarded) + 'bf16_gemm_nt', 'bf16_gemm_nn', + 'bf16_gemm_tn', 'bf16_gemm_tt', + 'm_grouped_bf16_gemm_nt_contiguous', + 'm_grouped_bf16_gemm_nn_contiguous', + 'm_grouped_bf16_gemm_nt_masked', + 'k_grouped_bf16_gemm_tn_contiguous', + # Einsum + 'einsum', + 'fp8_einsum', + # Attention + 'fp8_fp4_mqa_logits', + 'get_paged_mqa_logits_metadata', + 'fp8_fp4_paged_mqa_logits', + 'fp8_mqa_logits', + 'fp8_paged_mqa_logits', + # Hyperconnection (guarded) + 'tf32_hc_prenorm_gemm', + # Layout (guarded) + 'transform_sf_into_required_layout', + 'get_tma_aligned_size', + 'get_mn_major_tma_aligned_tensor', + 'get_mn_major_tma_aligned_packed_ue8m0_tensor', + 'get_k_grouped_mn_major_tma_aligned_packed_ue8m0_tensor', + # Mega helpers (guarded) + 'get_token_alignment_for_mega_moe', + 'get_block_m_for_mega_moe', +) + +__all__ = list(_UNCONDITIONAL_API) + [name for name in _DEEP_GEMM_API if name in globals()] diff --git a/deep_gemm/__init__.py b/deep_gemm/__init__.py index f5e52b2d12..0ebe611908 100644 --- a/deep_gemm/__init__.py +++ b/deep_gemm/__init__.py @@ -1,6 +1,5 @@ import os import subprocess -import torch # Set some default environment provided at setup try: diff --git a/deep_gemm/include/deep_gemm/layout/mqa_logits.cuh b/deep_gemm/include/deep_gemm/layout/mqa_logits.cuh index 14485b54c1..3358eda73e 100644 --- a/deep_gemm/include/deep_gemm/layout/mqa_logits.cuh +++ b/deep_gemm/include/deep_gemm/layout/mqa_logits.cuh @@ -1,5 +1,6 @@ #pragma once +#include #include #include diff --git a/deep_gemm/mega/__init__.py b/deep_gemm/mega/__init__.py index 2314564e22..1bb2417018 100644 --- a/deep_gemm/mega/__init__.py +++ b/deep_gemm/mega/__init__.py @@ -31,13 +31,14 @@ def __init__(self, group: dist.ProcessGroup, self.hidden = hidden self.intermediate_hidden = intermediate_hidden - # Allocate a symmetric buffer - num_bytes, slice_input_buffers = _C.get_symm_buffer_size_for_mega_moe( + # Allocate a symmetric buffer. The layout is computed once here and reused for + # slicing below. + num_bytes, layout_info = _C.get_symm_buffer_size_for_mega_moe( group.size(), num_experts, num_max_tokens_per_rank, num_topk, hidden, intermediate_hidden, mma_type, activation, - num_shared_experts + num_shared_experts, ) allocator = torch if group.size() == 1 else symm_mem self.buffer = allocator.empty(num_bytes, dtype=torch.int8, device='cuda') @@ -56,7 +57,7 @@ def __init__(self, group: dist.ProcessGroup, self.shared_l1_acts, self.shared_l1_acts_sf, self.shared_l2_acts, self.shared_l2_acts_sf, self.l1_acts, self.l1_acts_sf, - self.l2_acts, self.l2_acts_sf) = slice_input_buffers(self.buffer) + self.l2_acts, self.l2_acts_sf) = _C._slice_symm_buffer_for_mega_moe(self.buffer, layout_info) def destroy(self): self.handle = None diff --git a/scripts/generate_pyi.py b/scripts/generate_pyi.py index df7490d410..1f1dcbe9a0 100644 --- a/scripts/generate_pyi.py +++ b/scripts/generate_pyi.py @@ -1,108 +1,11 @@ +"""Generate deep_gemm/_C.pyi from TORCH_LIBRARY schemas and deep_gemm/_C.py wrappers.""" +import ast import re from pathlib import Path - -def build_cpp_function_index(root_path): - func_index = {} - extensions = {'.cpp', '.cc', '.cxx', '.c', '.hpp', '.h'} - - pattern = re.compile( - r'([\w:\s*<&>,\[\]\(\)]+?)' - r'\s+' - r'([a-zA-Z_][a-zA-Z0-9_:]*)' - r'\s*\(', - ) - - for file_path in Path(root_path).rglob('*'): - if file_path.suffix.lower() not in extensions: - continue - if not file_path.is_file(): - continue - - try: - with open(file_path, 'r', encoding='utf-8', errors='ignore') as f: - content = f.read() - except Exception as e: - print(f'Failed to read file {file_path}: {e}') - continue - - # Remove the compile directives and comments - lines = content.split('\n') - clean_lines = [line for line in lines if not line.strip().startswith(('#', '//'))] - content = '\n'.join(clean_lines) - - for match in pattern.finditer(content): - return_type_part = match.group(1).strip() - full_func_name = match.group(2).strip() - - if not return_type_part or not re.match(r'^[a-zA-Z_]', return_type_part): - continue - - first_token = return_type_part.split()[0] - if first_token in {'return', 'if', 'else', 'for', 'while', 'switch', 'case', 'throw', 'catch', 'auto'}: - continue - - # Extract base name - if '::' in full_func_name: - base_name = full_func_name.split('::')[-1] - else: - base_name = full_func_name - - if not re.match(r'^[a-zA-Z_][a-zA-Z0-9_]*$', base_name): - continue - - # Find matching ')' - paren_start = match.end() - 1 - paren_count = 0 - pos = paren_start - while pos < len(content): - ch = content[pos] - if ch == '(': - paren_count += 1 - elif ch == ')': - paren_count -= 1 - if paren_count == 0: - break - elif paren_count < 0: - pos = -1 - break - pos += 1 - else: - continue - - if pos == -1: - continue - - # Check context before match: should be at statement boundary - match_start = match.start() - context_before = content[max(0, match_start - 50):match_start] - if context_before and re.search(r'[a-zA-Z0-9_]$', context_before.rstrip()): - continue - - # Check for definition or header declaration - is_header = file_path.suffix.lower() in {'.h', '.hpp', '.cuh'} - after_paren = content[pos+1:pos+500] - has_brace = '{' in after_paren - has_semicolon = ';' in after_paren.split('{')[0] - - if has_brace or (is_header and has_semicolon): - sig_start = match.start(1) - full_signature = content[sig_start:pos+1].strip() - if base_name not in func_index: - func_index[base_name] = full_signature - - return func_index - - class BracketTracker: - """ - Tracks nesting levels of various brackets in C++ code: - - () → paren - - [] → bracket - - {} → brace - - <> → angle (treated as template brackets only at top level) - Provides is_top_level() to check if currently outside all brackets. - """ + """Track () [] {} <> nesting for top-level comma/default splitting.""" + def __init__(self): self.paren = 0 # () self.bracket = 0 # [] @@ -110,9 +13,6 @@ def __init__(self): self.angle = 0 # <> def update(self, char: str): - """ - Update internal counters based on the given character. - """ if char == '(': self.paren += 1 elif char == ')': @@ -133,32 +33,376 @@ def update(self, char: str): self.angle -= 1 def _in_top_level_of_other_brackets(self): - """ - Check if not inside parentheses, square brackets, or braces (for correct template bracket recognition). - """ return self.paren == 0 and self.bracket == 0 and self.brace == 0 def is_top_level(self): - """ - Check if completely at top level (all bracket counters are zero). - """ - return (self.paren == 0 and - self.bracket == 0 and - self.brace == 0 and - self.angle == 0) - - -def extract_m_def_statements(root_path): - """ - Scan all c files under root_path and extract all m.def(...) statements. - """ - results = [] - extensions = {'.hpp', '.cpp', '.h', '.cc'} + return self.paren == 0 and self.bracket == 0 and self.brace == 0 and self.angle == 0 + + +def split_top_level_commas(value: str) -> list[str]: + """Split on commas not nested inside brackets.""" + parts = [] + current = [] + tracker = BracketTracker() + for ch in value: + if ch in '()[]{}<>': + tracker.update(ch) + if ch == ',' and tracker.is_top_level(): + parts.append(''.join(current).strip()) + current = [] + else: + current.append(ch) + if current: + parts.append(''.join(current).strip()) + return parts + + +def find_top_level_equals(value: str) -> int: + """Return index of top-level '=', or -1.""" + tracker = BracketTracker() + for i, ch in enumerate(value): + if ch in '()[]{}<>': + tracker.update(ch) + elif ch == '=' and tracker.is_top_level(): + return i + return -1 + + +def schema_type_to_python(type_str: str) -> str: + """Map a TORCH schema type to a Python annotation.""" + type_str = type_str.strip() + optional = type_str.endswith('?') + if optional: + type_str = type_str[:-1].strip() + + if type_str.startswith('Tensor'): + py_type = 'torch.Tensor' + elif type_str == 'int': + py_type = 'int' + elif type_str == 'bool': + py_type = 'bool' + elif type_str == 'float': + py_type = 'float' + elif type_str == 'str': + py_type = 'str' + elif type_str == 'int[]': + py_type = 'list[int]' + elif re.match(r'^int\[\d+\]$', type_str): + # Fixed-size int[N] maps directly to a same-arity tuple. + n = int(re.match(r'^int\[(\d+)\]$', type_str).group(1)) + py_type = f"tuple[{', '.join(['int'] * n)}]" + elif type_str == 'ScalarType': + py_type = 'torch.dtype' + else: + print(f'Warning: unrecognized schema type {type_str!r}, using Any') + py_type = 'Any' + + if optional: + return f'Optional[{py_type}]' + return py_type - # Regex: match m.def( ... ), supports multi-line - pattern = re.compile(r'm\.def\s*\(') - for file_path in Path(root_path).rglob('*'): +def schema_return_to_python(return_str: str) -> str: + """Map a TORCH schema return type to a Python annotation.""" + return_str = return_str.strip() + if return_str == '()': + return 'None' + if return_str in {'int', 'bool', 'float', 'str', 'Tensor'}: + return { + 'int': 'int', + 'bool': 'bool', + 'float': 'float', + 'str': 'str', + 'Tensor': 'torch.Tensor', + }[return_str] + if return_str.startswith('(') and return_str.endswith(')'): + inner = return_str[1:-1].strip() + if not inner: + return 'tuple[()]' + parts = split_top_level_commas(inner) + py_parts = [schema_return_to_python(part) for part in parts] + return f'tuple[{", ".join(py_parts)}]' + print(f'Warning: unrecognized schema return type {return_str!r}, using Any') + return 'Any' + + +_SCALAR_TYPE_DEFAULTS = { + 'float': 'torch.float32', + 'float32': 'torch.float32', + 'double': 'torch.float64', + 'float64': 'torch.float64', + 'half': 'torch.float16', + 'float16': 'torch.float16', + 'bfloat16': 'torch.bfloat16', + 'byte': 'torch.uint8', + 'char': 'torch.int8', + 'short': 'torch.int16', + 'int': 'torch.int32', + 'long': 'torch.int64', +} + + +def schema_default_to_python(default_str: str) -> str: + """Convert a TORCH schema default literal to a Python expression string.""" + default_str = default_str.strip() + if default_str in {'None', 'True', 'False'}: + return default_str + if (default_str.startswith("'") and default_str.endswith("'")) or ( + default_str.startswith('"') and default_str.endswith('"')): + return default_str + if default_str in _SCALAR_TYPE_DEFAULTS: + return _SCALAR_TYPE_DEFAULTS[default_str] + if re.match(r'^[+-]?\d+$', default_str): + return default_str + if re.match(r'^[+-]?\d*\.\d+([eE][+-]?\d+)?$', default_str): + return default_str + print(f'Warning: unrecognized schema default {default_str!r}, using None') + return 'None' + + +def parse_schema_arg(arg_str: str) -> dict: + """Parse one TORCH schema argument such as 'Tensor? c=None'.""" + arg_str = arg_str.strip() + if not arg_str: + raise ValueError('empty schema argument') + + default = None + eq_pos = find_top_level_equals(arg_str) + if eq_pos != -1: + default = schema_default_to_python(arg_str[eq_pos + 1:].strip()) + arg_str = arg_str[:eq_pos].strip() + + match = re.match(r'^(.+?)\s+([a-zA-Z_][a-zA-Z0-9_]*)$', arg_str) + if not match: + raise ValueError(f'could not parse schema argument: {arg_str!r}') + return { + 'name': match.group(2), + 'py_type': schema_type_to_python(match.group(1)), + 'default': default, + } + + +def parse_torch_schema(schema: str) -> dict: + """Parse a TORCH schema into name, parameters, and return type.""" + arrow = schema.rfind(' -> ') + if arrow == -1: + raise ValueError(f'schema missing return type: {schema!r}') + + signature = schema[:arrow].strip() + return_type = schema_return_to_python(schema[arrow + 4:].strip()) + + open_paren = signature.find('(') + if open_paren == -1: + raise ValueError(f'schema missing argument list: {schema!r}') + + name = signature[:open_paren].strip() + paren_depth = 0 + close_paren = -1 + for i in range(open_paren, len(signature)): + if signature[i] == '(': + paren_depth += 1 + elif signature[i] == ')': + paren_depth -= 1 + if paren_depth == 0: + close_paren = i + break + if close_paren == -1: + raise ValueError(f'unclosed argument list in schema: {schema!r}') + + args_blob = signature[open_paren + 1:close_paren].strip() + parameters = [] + if args_blob: + for arg in split_top_level_commas(args_blob): + parameters.append(parse_schema_arg(arg)) + + return { + 'python_function_name': name, + 'parameters': parameters, + 'return_type': return_type, + } + + +def _merge_named_pairs(parameters: list[dict], pairs: tuple[tuple[str, str], ...]) -> list[dict]: + """Replace (tensor, scale_factor) arg pairs with one tuple-typed parameter.""" + by_name = {param['name']: param for param in parameters} + sf_of = dict(pairs) + drop = set(sf_of.values()) + + out = [] + for param in parameters: + if param['name'] in drop: + continue + sf_name = sf_of.get(param['name']) + if sf_name is None: + out.append(dict(param)) + continue + base_optional = param['py_type'] == 'Optional[torch.Tensor]' + sf_optional = by_name[sf_name]['py_type'] == 'Optional[torch.Tensor]' + if base_optional and sf_optional: + py_type = 'Optional[tuple[torch.Tensor, torch.Tensor]]' + elif sf_optional: + py_type = 'tuple[torch.Tensor, Optional[torch.Tensor]]' + else: + py_type = 'tuple[torch.Tensor, torch.Tensor]' + out.append({ + 'name': param['name'], + 'py_type': py_type, + 'default': None, + }) + return out + + +def _is_tensor_schema_param(param: dict) -> bool: + py_type = param['py_type'] + return py_type in {'torch.Tensor', 'Optional[torch.Tensor]'} + + +def _is_tensor_scale_factor_pair(base_name: str, sf_name: str) -> bool: + if base_name == 'a' and sf_name == 'sfa': + return True + if base_name == 'b' and sf_name == 'sfb': + return True + return sf_name == f'{base_name}_sf' + + +def detect_tensor_sf_pairs(parameters: list[dict]) -> list[tuple[str, str]]: + """Detect consecutive (tensor, scale_factor) arg pairs.""" + pairs = [] + i = 0 + while i < len(parameters) - 1: + left, right = parameters[i], parameters[i + 1] + if ( + _is_tensor_schema_param(left) + and _is_tensor_schema_param(right) + and _is_tensor_scale_factor_pair(left['name'], right['name']) + ): + pairs.append((left['name'], right['name'])) + i += 2 + else: + i += 1 + return pairs + + +def _maybe_widen_int_list_value_param(parameters: list[dict]) -> None: + """Single int[] value param in a Python wrapper usually accepts int | list[int].""" + if len(parameters) == 1 and parameters[0]['name'] == 'value': + if parameters[0]['py_type'] == 'list[int]': + parameters[0]['py_type'] = 'int | list[int]' + + +def _promote_transform_sf_recipe_type(op_name: str, parameters: list[dict]) -> None: + """transform_sf_into_required_layout's recipe is a std::variant, which int[N] can't express, so promote it here.""" + if op_name != 'transform_sf_into_required_layout': + return + recipe = next((p for p in parameters if p['name'] == 'recipe'), None) + if recipe is not None and recipe['py_type'] == 'list[int]': + recipe['py_type'] = 'tuple[int, int] | tuple[int, int, int]' + + +def adjust_for_c_py_wrapper( + name: str, + parameters: list[dict], + wrapper_defaults: dict[str, dict[str, str]] | None = None, +) -> list[dict]: + """Adjust flat schema params to match deep_gemm._C Python wrappers.""" + pairs = detect_tensor_sf_pairs(parameters) + if pairs: + parameters = _merge_named_pairs(parameters, tuple(pairs)) + + _maybe_widen_int_list_value_param(parameters) + _promote_transform_sf_recipe_type(name, parameters) + + return parameters + + +def format_ast_default(node: ast.AST) -> str: + """Convert an AST default value node to a Python expression string.""" + if isinstance(node, ast.Constant): + if node.value is None: + return 'None' + if isinstance(node.value, bool): + return 'True' if node.value else 'False' + if isinstance(node.value, str): + return f'"{node.value}"' + if isinstance(node.value, (int, float)): + return repr(node.value) + if isinstance(node, ast.Tuple): + elts = ', '.join(format_ast_default(element) for element in node.elts) + return f'({elts})' + if isinstance(node, ast.List): + elts = ', '.join(format_ast_default(element) for element in node.elts) + return f'[{elts}]' + if isinstance(node, ast.Name): + return node.id + if isinstance(node, ast.Attribute): + return ast.unparse(node) + if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.USub): + return f'-{format_ast_default(node.operand)}' + return ast.unparse(node) + + +def extract_function_defaults(func_def: ast.FunctionDef) -> dict[str, str]: + """Extract {param_name: default_expr} from a Python function definition.""" + defaults: dict[str, str] = {} + args = func_def.args + pos_args = args.args + if args.defaults: + first_default_idx = len(pos_args) - len(args.defaults) + for idx, default_node in enumerate(args.defaults): + defaults[pos_args[first_default_idx + idx].arg] = format_ast_default(default_node) + for arg, default_node in zip(args.kwonlyargs, args.kw_defaults): + if default_node is not None: + defaults[arg.arg] = format_ast_default(default_node) + return defaults + + +def _is_globals_update_call(node: ast.Call) -> bool: + if not isinstance(node.func, ast.Attribute) or node.func.attr != 'update': + return False + base = node.func.value + if isinstance(base, ast.Name): + return base.id == 'globals' + if isinstance(base, ast.Call) and isinstance(base.func, ast.Name): + return base.func.id == 'globals' + return False + + +def parse_c_py_metadata(c_py_path: Path) -> dict[str, dict[str, str]]: + """Parse wrapper defaults from deep_gemm/_C.py without importing it.""" + source = c_py_path.read_text(encoding='utf-8') + module = ast.parse(source, filename=str(c_py_path)) + + func_defaults: dict[str, dict[str, str]] = {} + + for node in ast.walk(module): + if isinstance(node, ast.FunctionDef): + func_defaults[node.name] = extract_function_defaults(node) + + for node in ast.walk(module): + if not isinstance(node, ast.Call): + continue + if not _is_globals_update_call(node): + continue + if not node.args or not isinstance(node.args[0], ast.Dict): + continue + alias_dict = node.args[0] + for key_node, value_node in zip(alias_dict.keys, alias_dict.values): + if not isinstance(key_node, ast.Constant) or not isinstance(key_node.value, str): + continue + alias_name = key_node.value + if isinstance(value_node, ast.Name): + if value_node.id in func_defaults: + func_defaults[alias_name] = func_defaults[value_node.id] + + return func_defaults + + +def extract_m_def_statements(root_path) -> list[str]: + """Scan C++ sources under root_path for m.def(...) registrations.""" + statements = [] + extensions = {'.hpp', '.cpp', '.h', '.cc'} + + for file_path in sorted(Path(root_path).rglob('*')): if file_path.suffix.lower() not in extensions: continue if not file_path.is_file(): @@ -177,15 +421,11 @@ def extract_m_def_statements(root_path): while i < len(lines): line = lines[i] if 'm.def(' in line: - # Found a potential starting line - start_i = i - # Check if it's a comment stripped = line.lstrip() if stripped.startswith('//') or stripped.startswith('/*'): i += 1 continue - # Try to match the complete m.def(...) call paren_count = 0 j = i found_start = False @@ -201,7 +441,6 @@ def extract_m_def_statements(root_path): if found_start: paren_count -= 1 if paren_count == 0: - # Found complete statement full_stmt = ''.join(lines[i:j+1]).rstrip() m_def_list.append(full_stmt) i = j @@ -209,28 +448,16 @@ def extract_m_def_statements(root_path): if paren_count <= 0 and found_start: break j += 1 - else: - pass i += 1 if m_def_list: - results.append({ - 'file': str(file_path), - 'm_def_statements': m_def_list - }) + statements.extend(m_def_list) - return results + return statements def parse_m_def_statement(m_def_str): - result = { - 'python_function_name': None, - 'num_args': 0, - 'default_args': {}, - 'is_lambda': False, - } - - # Extract top-level arguments + """Parse a TORCH_LIBRARY m.def(...) statement.""" start = m_def_str.find('m.def(') if start == -1: raise ValueError(f'[{m_def_str}] Could not find m.def start position') @@ -252,634 +479,108 @@ def parse_m_def_statement(m_def_str): raise ValueError(f'[{m_def_str}] m.def parentheses not closed') args_content = m_def_str[content_start:content_end] + args_list = split_top_level_commas(args_content) + if not args_list: + raise ValueError(f'[{m_def_str}] m.def has no arguments') - # Split arguments using BracketTracker - args_list = [] - current = [] - tracker = BracketTracker() - - for ch in args_content: - if ch in '()[]{}<>': - tracker.update(ch) - if ch == ',' and tracker.is_top_level(): - args_list.append(''.join(current).strip()) - current = [] - else: - current.append(ch) - - if current: - args_list.append(''.join(current).strip()) - - if len(args_list) < 2: - raise ValueError(f'[{m_def_str}] m.def has insufficient arguments') - - # Extract Python function name first = args_list[0].strip() str_match = re.match(r'^"([^"\\]*(?:\\.[^"\\]*)*)"', first) - if str_match: - result['python_function_name'] = str_match.group(1) - else: + if not str_match: raise ValueError(f'[{m_def_str}] m.def first argument should be a string literal') - cpp_func_part = args_list[1].strip() - if cpp_func_part.startswith('&'): - cpp_func_part = cpp_func_part[1:].strip() + return parse_torch_schema(str_match.group(1)) - if cpp_func_part.startswith('['): - result['is_lambda'] = True - result['cpp_function_name'] = None - else: - if '::' in cpp_func_part: - cpp_func_name = cpp_func_part.split('::')[-1] - else: - cpp_func_name = cpp_func_part - match = re.match(r'^([a-zA-Z_][a-zA-Z0-9_]*)', cpp_func_name) - if match: - result['cpp_function_name'] = match.group(1) - else: - result['cpp_function_name'] = cpp_func_name +def apply_wrapper_defaults(name: str, parameters: list[dict], wrapper_defaults: dict[str, dict[str, str]]) -> list[dict]: + """Overlay public API defaults from deep_gemm/_C.py onto schema-derived parameters.""" + by_name = wrapper_defaults.get(name, {}) + if not by_name: + return parameters - # Parse py::arg arguments - py_args = args_list[2:] - result['num_args'] = len(py_args) + out = [] + for param in parameters: + param = dict(param) + if param['name'] in by_name: + param['default'] = by_name[param['name']] + out.append(param) + return out - for idx, arg_expr in enumerate(py_args): - expr = arg_expr.strip() - # Find top-level '=' - eq_pos = -1 - p_depth = b_depth = br_depth = angle_depth = 0 - i = 0 - while i < len(expr): - ch = expr[i] - if ch == '(': - p_depth += 1 - elif ch == ')': - p_depth -= 1 - elif ch == '[': - b_depth += 1 - elif ch == ']': - b_depth -= 1 - elif ch == '{': - br_depth += 1 - elif ch == '}': - br_depth -= 1 - elif ch == '<' and p_depth == 0 and b_depth == 0 and br_depth == 0: - angle_depth += 1 - elif ch == '>' and angle_depth > 0 and p_depth == 0 and b_depth == 0 and br_depth == 0: - angle_depth -= 1 - elif ch == '=' and all(d == 0 for d in [p_depth, b_depth, br_depth, angle_depth]): - eq_pos = i - break - i += 1 - if eq_pos != -1: - default_val = expr[eq_pos + 1:].strip() - if not default_val: - raise ValueError(f'[{expr}] Default value is empty (arg {idx})') - result['default_args'][idx] = default_val - - return result - - -def extract_cpp_signature_from_content(cpp_func_name, content): - """ - Search for the C++ function signature of cpp_func_name in the given file content. - """ - if not cpp_func_name: - return None - - # Build regex: match function starting with cpp_func_name (after word boundary) - # Note: function name may be preceded by return type (with templates, namespaces, etc.), followed by '(' - pattern = re.compile( - r'^\s*' # leading whitespace - r'([\w:\s*<&>,\[\]\(\)]+?)' # return type (non-greedy, allows templates, pointers, etc.) - r'\s+' # at least one space - r'\b' + re.escape(cpp_func_name) + r'\b' # function name (word boundary) - r'\s*\(', # optional whitespace + start of param list - re.MULTILINE +def generate_pyi_function(parsed, wrapper_defaults=None): + """Generate a typed .pyi stub for one registered op.""" + py_name = parsed['python_function_name'] + parameters = adjust_for_c_py_wrapper( + py_name, + parsed['parameters'], + wrapper_defaults=wrapper_defaults, ) + if wrapper_defaults: + parameters = apply_wrapper_defaults(py_name, parameters, wrapper_defaults) + return_type = parsed['return_type'] - for match in pattern.finditer(content): - # Find '(' position after function name - paren_start = match.end() - 1 - if content[paren_start] != '(': - paren_start = content.find('(', match.end(0) - 1) - if paren_start == -1: - continue - - # From '(', match to corresponding ')' - paren_count = 0 - pos = paren_start - while pos < len(content): - ch = content[pos] - if ch == '(': - paren_count += 1 - elif ch == ')': - paren_count -= 1 - if paren_count == 0: - start_sig = match.start(1) - full_signature = content[start_sig:pos+1].strip() - return full_signature - pos += 1 - - return None - - -def parse_mdef_and_attach_cpp_signatures(item, func_index): - """ - Enhance item by parsing m.def and extracting C++ function signature from global index - """ - statements_with_parsed_signatures = [] - - for stmt in item['m_def_statements']: - parsed = parse_m_def_statement(stmt,) - cpp_func_name = parsed.get('cpp_function_name') - - cpp_sig = None - if cpp_func_name and cpp_func_name in func_index: - cpp_sig = func_index[cpp_func_name] - else: - if not parsed['is_lambda']: - print(f'Warning: C++ function "{cpp_func_name}" not found in any .cpp file') - - parsed['cpp_signature'] = cpp_sig - statements_with_parsed_signatures.append({ - 'raw': stmt, - 'parsed': parsed - }) - - return { - 'm_def_statements': statements_with_parsed_signatures - } - - -def parse_cpp_signature(cpp_sig): - """ - Parse a C++ function signature and extract return type, parameter types, and names. - """ - if not cpp_sig or not cpp_sig.strip(): - return None - - # Find function name: last identifier before '(' - paren_pos = cpp_sig.find('(') - if paren_pos == -1: - return None - - before_paren = cpp_sig[:paren_pos].strip() - if not before_paren: - return None - - # Function name is the last word in before_paren (may include templates like func) - tokens = before_paren.split() - if len(tokens) < 2: - return None - - # Heuristic: function name is usually the last token (may include <>) - func_name_part = tokens[-1] - return_type = ' '.join(tokens[:-1]).strip() - - # Now extract parameter list content - param_list_str = cpp_sig[paren_pos+1:cpp_sig.rfind(')')].strip() - parameters = [] - - if param_list_str and param_list_str != 'void': # 'void' means no parameters - # Split parameters (handle commas not inside templates/brackets) - param_decls = split_cpp_parameters(param_list_str) - for decl in param_decls: - decl = decl.strip() - if not decl: - continue - # Try to split type and name from right to left - param_info = parse_parameter_declaration(decl) - if param_info: - parameters.append(param_info) - - return { - 'return_type': return_type, - 'parameters': parameters, - 'num_parameters': len(parameters) - } - - -def split_cpp_parameters(param_str: str): - """ - Split a C++ parameter list string by top-level commas, - e.g., 'int a, std::vector b' → ['int a', 'std::vector b'] - """ - if not param_str.strip() or param_str == 'void': - return [] - params = [] - current = [] - tracker = BracketTracker() - - for ch in param_str: - if ch in '()[]{}<>': - tracker.update(ch) - if ch == ',' and tracker.is_top_level(): - param = ''.join(current).strip() - if param: # Only add non-empty parameters - params.append(param) - current = [] - else: - current.append(ch) - - if current: - final_param = ''.join(current).strip() - if final_param: # Only add non-empty parameters - params.append(final_param) - return params - - -def parse_parameter_declaration(decl: str): - """ - Parse a single parameter declaration, e.g., 'const std::string& name' → {'type': 'const std::string&', 'name': 'name'} - Improved version that better handles template types. - """ - decl = decl.strip() - if not decl: - return None - - # Remove possible default value (starting from top-level '=') - tracker = BracketTracker() - eq_pos = -1 - for i, ch in enumerate(decl): - if ch in '()[]{}<>': - tracker.update(ch) - elif ch == '=' and tracker.is_top_level(): - eq_pos = i - break - - if eq_pos != -1: - decl = decl[:eq_pos].strip() - - # Now decl is 'type name' or just 'type' - # Instead of simple splitting, we'll use a more robust approach - # to find the parameter name - - # First, let's handle the case where there's no explicit parameter name - # (this sometimes happens in function declarations) - if not re.search(r'[a-zA-Z_][a-zA-Z0-9_]*$', decl): - # No parameter name found, just return the type - return { - 'type': decl, - 'name': None - } - - # Use bracket tracking to find where the type ends and name begins - tracker = BracketTracker() - name_start = -1 - - # Scan from the end to find the start of the parameter name - # We look for the first identifier that's outside all brackets - i = len(decl) - 1 - while i >= 0: - ch = decl[i] - - if ch in '()[]{}<>': - tracker.update(ch) - - # If we're at top level and find an identifier character - if tracker.is_top_level() and re.match(r'[a-zA-Z0-9_]', ch): - # Track back to find the start of this identifier - name_start = i - while name_start > 0 and re.match(r'[a-zA-Z0-9_]', decl[name_start - 1]): - name_start -= 1 - - # Check if this might be part of a type keyword (like 'int', 'bool', etc.) - potential_name = decl[name_start:i+1] - type_keywords = {'int', 'long', 'short', 'char', 'bool', 'float', 'double', - 'void', 'auto', 'const', 'static', 'volatile', 'mutable', - 'unsigned', 'signed'} - - # If it's not a type keyword and looks like a parameter name, use it - if (potential_name not in type_keywords and - re.match(r'^[a-zA-Z_][a-zA-Z0-9_]*$', potential_name)): - break - - i -= 1 - - if name_start != -1 and i >= 0: - param_name = decl[name_start:i+1] - param_type = decl[:name_start].strip() - - # Clean up the type - remove trailing &, * and whitespace - param_type = param_type.rstrip('&* \t') - - return { - 'type': param_type, - 'name': param_name - } - - # Fallback: if we can't find a clear parameter name, just return the type - return { - 'type': decl, - 'name': None - } - - -def extract_cpp_signature_details(item): - """ - For each m.def entry in item, parse cpp_signature to extract return type and parameter details. - """ - statements_with_parsed_signatures = [] - for stmt_info in item['m_def_statements']: - parsed = stmt_info['parsed'] - cpp_sig = parsed.get('cpp_signature') - - cpp_params_info = None - if cpp_sig: - try: - cpp_params_info = parse_cpp_signature(cpp_sig) - except Exception as e: - print(f'Failed to parse C++ signature: {e}') - - parsed['cpp_parsed_signature'] = cpp_params_info - statements_with_parsed_signatures.append({ - 'raw': stmt_info['raw'], - 'parsed': parsed - }) - - return { - 'm_def_statements': statements_with_parsed_signatures - } - - -def cpp_type_to_python_type(cpp_type: str) -> str: - if not cpp_type: - return 'Any' - - original = cpp_type.strip() - if not original: - return 'Any' - - # Remove C++ specifiers that don't affect Python type - cleaned = re.sub(r'\b(static|inline|constexpr|thread_local|extern|mutable|const|volatile|endif)\b', '', original) - cleaned = cleaned.replace('&', '').replace('*', '').strip() - cleaned = re.sub(r'\s+', ' ', cleaned).strip() - - # Handle void - if cleaned == 'void': - return 'None' - - # Handle template types — ORDER MATTERS! Must come before internal type checks. - - # std::pair - if cleaned.startswith('std::pair<'): - inner = cleaned[10:-1].strip() # len('std::pair<') == 10 - args = split_template_args(inner) - if len(args) == 2: - t1 = cpp_type_to_python_type(args[0]) - t2 = cpp_type_to_python_type(args[1]) - return f'tuple[{t1}, {t2}]' - else: - print(f'Warning: std::pair with unexpected number of args: {cleaned}') - return 'Any' - - # std::tuple - if cleaned.startswith('std::tuple<'): - inner = cleaned[11:-1].strip() # len('std::tuple<') == 11 - args = split_template_args(inner) - py_types = [cpp_type_to_python_type(arg) for arg in args] - return f"tuple[{', '.join(py_types)}]" - - # std::vector - if cleaned.startswith('std::vector<'): - inner = cleaned[12:-1].strip() # len('std::vector<') == 12 - args = split_template_args(inner) - if len(args) == 1: - inner_py = cpp_type_to_python_type(args[0]) - return f'list[{inner_py}]' - else: - print(f'Warning: std::vector with unexpected args: {cleaned}') - return 'Any' - - # std::optional - if cleaned.startswith('std::optional<'): - inner = cleaned[14:-1].strip() # len('std::optional<') == 14 - args = split_template_args(inner) - if len(args) == 1: - inner_py = cpp_type_to_python_type(args[0]) - return f'Optional[{inner_py}]' - else: - print(f'Warning: std::optional with unexpected args: {cleaned}') - return 'Any' - - # std::string - if re.search(r'\bstd::string\b', original): - return 'str' - - # C-style strings: char*, const char*, char[], etc. - if re.search(r'\b(?:const\s+)?char\s*[\*\[]', original): - return 'str' - - # Boolean - if re.search(r'\bbool\b', cleaned): - return 'bool' - - # Integer types (including fixed-width and common aliases) - if re.search(r'\b(int|long|short|size_t|ssize_t|ptrdiff_t|' - r'int8_t|int16_t|int32_t|int64_t|' - r'uint8_t|uint16_t|uint32_t|uint64_t)\b', cleaned): - return 'int' - - # Floating-point - if re.search(r'\b(float|double|long\s+double)\b', cleaned): - return 'float' - - # torch::Tensor - if re.search(r'\btorch::Tensor\b', original): - return 'torch.Tensor' - - # Unrecognized type - print(f'Warning: Unrecognized C++ type: {original}') - return 'Any' - - -def split_template_args(template_args: str): - """ - Split template arguments, e.g., 'int, std::vector' → ['int', 'std::vector'] - """ - if not template_args.strip(): - return [] - args = [] - current = [] - tracker = BracketTracker() - - for ch in template_args: - if ch in '()[]{}<>': - tracker.update(ch) - if ch == ',' and tracker.is_top_level(): - args.append(''.join(current).strip()) - current = [] + param_lines = [] + for param in parameters: + name = param['name'] + if param['default'] is not None: + param_lines.append(f' {name}: {param["py_type"]} = {param["default"]}') else: - current.append(ch) - - if current: - args.append(''.join(current).strip()) - return args - - -def cpp_default_to_python_default(cpp_default: str): - """ - Convert C++ default value string to valid Python expression string. - """ - if not cpp_default: - return 'None' - - s = cpp_default.strip() - - # Handle string literals: 'bf16' → 'bf16' - # Match: starts and ends with unescaped double quotes - string_match = re.match(r'^"([^"\\]*(?:\\.[^"\\]*)*)"$', s) - if string_match: - return s - - # Handle boolean literals - if s == 'false': - return 'False' - if s == 'true': - return 'True' - - # Handle null-like values: nullptr, nullopt, NULL, etc. - if s in ('nullptr', 'NULL') or 'nullopt' in s: - return 'None' - - # Handle std::tuple({128, 128}) → (128, 128) - tuple_match = re.match(r'std::tuple\s*<[^>]*>\s*\(\s*({.*?})\s*\)', s) - if tuple_match: - inner = tuple_match.group(1) # {128, 128} - inner_py = inner.replace('{', '(').replace('}', ')') - return inner_py - - # Handle std::make_tuple(1, 2, 3) → (1, 2, 3) - make_tuple_match = re.match(r'std::make_tuple\s*\(\s*(.*?)\s*\)', s) - if make_tuple_match: - inner = make_tuple_match.group(1) - # Ensure it's a valid tuple even with one element: add comma if needed? - # But in C++ default args, it's usually multi-element, so we assume valid. - return f'({inner})' - - # Handle std::vector({1,2,3}) → [1, 2, 3] - vector_match = re.match(r'std::vector\s*<[^>]*>\s*\(\s*({.*?})\s*\)', s) - if vector_match: - inner = vector_match.group(1) - inner_py = inner.replace('{', '[').replace('}', ']') - return inner_py - - # Handle numeric literals: integers and floats - if re.match(r'^[+-]?\d+$', s): # integer - return s - if re.match(r'^[+-]?\d*\.\d+([eE][+-]?\d+)?$', s): # float - return s - - # Fallback: unrecognized → warn and return None - print(f'Warning: Unrecognized default value: {s}') - return 'None' + param_lines.append(f' {name}: {param["py_type"]}') + if param_lines: + params_block = ',\n'.join(param_lines) + return f'def {py_name}(\n{params_block}\n) -> {return_type}: ...' + return f'def {py_name}() -> {return_type}: ...' -def generate_pyi_function(item_entry): - parsed = item_entry['parsed'] - py_name = parsed['python_function_name'] - - if parsed.get('is_lambda'): - return f'def {py_name}(*args, **kwargs) -> Any: ...' - - sig_info = parsed.get('cpp_parsed_signature') - default_args = parsed.get('default_args', {}) - - if not sig_info: - return f'def {py_name}(*args, **kwargs) -> Any: ...' - return_type = cpp_type_to_python_type(sig_info['return_type']) - params = sig_info['parameters'] - num_params = len(params) +def generate_pyi_file_content( + parsed_ops, + module_name: str = 'my_module', + wrapper_defaults=None, +): + """Assemble the full .pyi file from parsed TORCH ops.""" + decls = [] - # Build parameter list - param_lines = [] - for i in range(num_params): - param_info = params[i] if i < len(params) else {'type': 'Any', 'name': f'arg{i}'} - param_type = cpp_type_to_python_type(param_info['type']) - param_name = param_info['name'] or f'arg{i}' - - # Replace invalid Python identifiers (e.g., keywords) - if param_name in {'def', 'class', 'from', 'import', 'None', 'True', 'False'}: - param_name = f'{param_name}_' - - # Check for default value - if i in default_args: - cpp_default = default_args[i] - py_default = cpp_default_to_python_default(cpp_default) - param_str = f' {param_name}: {param_type} = {py_default}' - else: - param_str = f' {param_name}: {param_type}' + for parsed in parsed_ops: + name = parsed['python_function_name'] + try: + decl = generate_pyi_function(parsed, wrapper_defaults=wrapper_defaults) + except Exception as e: + decl = f'# ERROR: failed to generate stub for {name}: {e}' + decls.append(decl) - param_lines.append(param_str) + lines = [ + f'# Stubs for module: {module_name}', + '', + 'from typing import Any, Optional', + 'import torch', + '', + '', + ] - if param_lines: - params_block = ',\n'.join(param_lines) - func_def = f'def {py_name}(\n{params_block}\n) -> {return_type}: ...' - else: - func_def = f'def {py_name}() -> {return_type}: ...' - - return func_def - - -def generate_pyi_file_content(enhanced_results, module_name: str = 'my_module'): - function_decls = [] - has_optional = False - has_torch = False - has_numpy = False - - for item in enhanced_results: - for stmt in item['m_def_statements']: - try: - decl = generate_pyi_function(stmt) - function_decls.append(decl) - - if 'Optional[' in decl: - has_optional = True - if 'torch.Tensor' in decl: - has_torch = True - if 'numpy.ndarray' in decl or 'py::array' in str(stmt): - has_numpy = True - except Exception as e: - func_name = stmt['parsed'].get('python_function_name', 'unknown') - function_decls.append(f'# ERROR: failed to generate stub for {func_name}: {e}') - - imports = ['from typing import Any'] - if has_optional: - imports[0] += ', Optional' - - if has_torch: - imports.append('import torch') - if has_numpy: - imports.append('import numpy') - - lines = [f'# Stubs for module: {module_name}', ''] - lines.extend(imports) - lines.append('') - lines.append('') - - for decl in function_decls: - lines.append(decl) - lines.append('') - lines.append('') + for decl in decls: + lines.extend([decl, '', '']) return '\n'.join(lines) -def generate_pyi_file(name, root, output_dir='.'): - func_index = build_cpp_function_index(root) - results = extract_m_def_statements(root) +def generate_pyi_file(name, root, output_dir='.', c_py_path=None): + """Generate stubs/.pyi from csrc/ schemas and optional _C.py wrapper defaults.""" + m_def_statements = extract_m_def_statements(root) + parsed_ops = [parse_m_def_statement(stmt) for stmt in m_def_statements] - cpp_results = [] - for item in results: - enhanced_item = parse_mdef_and_attach_cpp_signatures(item, func_index) - cpp_item = extract_cpp_signature_details(enhanced_item) - cpp_results.append(cpp_item) + wrapper_defaults = {} + if c_py_path is not None: + c_py_path = Path(c_py_path) + if c_py_path.is_file(): + wrapper_defaults = parse_c_py_metadata(c_py_path) + else: + print(f'Warning: wrapper file not found: {c_py_path}') - pyi_content = generate_pyi_file_content(cpp_results, module_name=name) + pyi_content = generate_pyi_file_content( + parsed_ops, + module_name=name, + wrapper_defaults=wrapper_defaults, + ) output_path = Path(output_dir) / f'{name}.pyi' output_path.parent.mkdir(parents=True, exist_ok=True) diff --git a/setup.py b/setup.py index c4d74ae929..50a5fd4202 100644 --- a/setup.py +++ b/setup.py @@ -26,7 +26,8 @@ # Compiler flags cxx_flags = ['-std=c++17', '-O3', '-fPIC', '-Wno-psabi', '-Wno-deprecated-declarations', - f'-D_GLIBCXX_USE_CXX11_ABI={int(torch.compiled_with_cxx11_abi())}'] + f'-D_GLIBCXX_USE_CXX11_ABI={int(torch.compiled_with_cxx11_abi())}', + '-DPy_LIMITED_API=0x030a0000'] if DG_JIT_USE_RUNTIME_API: cxx_flags.append('-DDG_JIT_USE_RUNTIME_API') @@ -103,12 +104,13 @@ def get_ext_modules(): if DG_SKIP_CUDA_BUILD: return [] - return [CUDAExtension(name='deep_gemm._C', + return [CUDAExtension(name='deep_gemm._C_extension', sources=sources, include_dirs=build_include_dirs, libraries=build_libraries, library_dirs=build_library_dirs, - extra_compile_args=cxx_flags)] + extra_compile_args=cxx_flags, + py_limited_api=True)] class CustomBuildPy(build_py): @@ -126,7 +128,12 @@ def run(self): build_py.run(self) def generate_pyi_file(self): - generate_pyi_file(name='_C', root='./csrc', output_dir='./stubs') + generate_pyi_file( + name='_C', + root='./csrc', + output_dir='./stubs', + c_py_path='./deep_gemm/_C.py', + ) pyi_source = os.path.join(current_dir, 'stubs', '_C.pyi') pyi_target = os.path.join(self.build_lib, 'deep_gemm', '_C.pyi') @@ -207,6 +214,7 @@ def run(self): }, ext_modules=get_ext_modules(), zip_safe=False, + options={'bdist_wheel': {'py_limited_api': 'cp310'}}, cmdclass={ 'build_py': CustomBuildPy, 'bdist_wheel': CachedWheelsCommand,