[Public release 26/04] Introducing Mega MoE, FP4 Indexer and other features/fixes (#304)
* Merge with private repo * Update README * Update README * Update README * Add PyTorch requirements * Fix sync scopes for MQA logits (#256) * Update README
This commit is contained in:
+68
-65
@@ -6,7 +6,7 @@
|
||||
#include "../jit_kernels/impls/sm90_fp8_gemm_1d1d.hpp"
|
||||
#include "../jit_kernels/impls/sm90_fp8_gemm_1d2d.hpp"
|
||||
#include "../jit_kernels/impls/sm90_bf16_gemm.hpp"
|
||||
#include "../jit_kernels/impls/sm100_fp8_gemm_1d1d.hpp"
|
||||
#include "../jit_kernels/impls/sm100_fp8_fp4_gemm_1d1d.hpp"
|
||||
#include "../jit_kernels/impls/sm100_bf16_gemm.hpp"
|
||||
#endif
|
||||
|
||||
@@ -23,7 +23,7 @@ static bool early_return(const int& m, const int &n, const int& k,
|
||||
return true;
|
||||
|
||||
// Checks
|
||||
const bool& is_cd_same = c.has_value() and c->data_ptr() == d.data_ptr();
|
||||
const bool is_cd_same = c.has_value() and c->data_ptr() == d.data_ptr();
|
||||
if (is_cd_same)
|
||||
DG_HOST_ASSERT(c->sizes() == d.sizes() and c->strides() == d.strides());
|
||||
if (c.has_value()) {
|
||||
@@ -57,8 +57,8 @@ static void fp8_fp4_gemm_nt(const std::pair<torch::Tensor, torch::Tensor>& a,
|
||||
const std::string& compiled_dims,
|
||||
const bool& disable_ue8m0_cast) {
|
||||
// Shape must be `[M, K] @ [N, K].T`
|
||||
const auto& major_a = get_major_type_ab(a.first);
|
||||
const auto& major_b = get_major_type_ab(b.first);
|
||||
const auto major_a = get_major_type_ab(a.first);
|
||||
const auto major_b = get_major_type_ab(b.first);
|
||||
if (fp8_requires_k_major()) {
|
||||
DG_HOST_ASSERT(major_a == cute::UMMA::Major::K);
|
||||
DG_HOST_ASSERT(major_b == cute::UMMA::Major::K);
|
||||
@@ -89,7 +89,7 @@ static void fp8_fp4_gemm_nt(const std::pair<torch::Tensor, torch::Tensor>& a,
|
||||
if (gran_n == 1) {
|
||||
sm90_fp8_gemm_1d1d(a.first, sfa, b.first, sfb, c, d, m, n, k, major_a, major_b, compiled_dims);
|
||||
} else {
|
||||
const auto& major_sfb = get_major_type_ab(sfb);
|
||||
const auto major_sfb = get_major_type_ab(sfb);
|
||||
sm90_fp8_gemm_1d2d(a.first, sfa, b.first, sfb, c, d, m, n, k, major_a, major_b, major_sfb, compiled_dims);
|
||||
}
|
||||
} else if (arch_major == 10 and sfa.scalar_type() == torch::kInt) {
|
||||
@@ -152,8 +152,8 @@ static void m_grouped_fp8_fp4_gemm_nt_contiguous(const std::pair<torch::Tensor,
|
||||
const bool& use_psum_layout,
|
||||
const std::optional<int>& expected_m_for_psum_layout) {
|
||||
// Shape must be `[M, K] @ [G, N, K].mT`
|
||||
const auto& major_a = get_major_type_ab(a.first);
|
||||
const auto& major_b = get_major_type_ab(b.first);
|
||||
const auto major_a = get_major_type_ab(a.first);
|
||||
const auto major_b = get_major_type_ab(b.first);
|
||||
DG_HOST_ASSERT(major_a == cute::UMMA::Major::K);
|
||||
if (fp8_requires_k_major())
|
||||
DG_HOST_ASSERT(major_b == cute::UMMA::Major::K);
|
||||
@@ -171,10 +171,10 @@ static void m_grouped_fp8_fp4_gemm_nt_contiguous(const std::pair<torch::Tensor,
|
||||
|
||||
// Layout checks
|
||||
if (use_psum_layout) {
|
||||
const auto& [num_groups_] = get_shape<1>(grouped_layout);
|
||||
const auto [num_groups_] = get_shape<1>(grouped_layout);
|
||||
DG_HOST_ASSERT(num_groups == num_groups_);
|
||||
} else {
|
||||
const auto& [m__] = get_shape<1>(grouped_layout);
|
||||
const auto [m__] = get_shape<1>(grouped_layout);
|
||||
DG_HOST_ASSERT(m == m__);
|
||||
DG_HOST_ASSERT(not expected_m_for_psum_layout.has_value());
|
||||
}
|
||||
@@ -192,10 +192,10 @@ static void m_grouped_fp8_fp4_gemm_nt_contiguous(const std::pair<torch::Tensor,
|
||||
|
||||
// Dispatch implementation
|
||||
if (arch_major == 9 and sfa.scalar_type() == torch::kFloat) {
|
||||
const auto& major_sfb = get_major_type_ab(sfb);
|
||||
DG_HOST_ASSERT(not use_psum_layout);
|
||||
const auto major_sfb = get_major_type_ab(sfb);
|
||||
sm90_m_grouped_fp8_gemm_contiguous_1d2d(a.first, sfa, b.first, sfb, d, grouped_layout,
|
||||
num_groups, m, n, k, major_a, major_b, major_sfb, compiled_dims);
|
||||
num_groups, m, n, k, major_a, major_b, major_sfb,
|
||||
compiled_dims, use_psum_layout, expected_m_for_psum_layout);
|
||||
} else if (arch_major == 10 and sfa.scalar_type() == torch::kInt) {
|
||||
sm100_m_grouped_fp8_fp4_gemm_contiguous_1d1d(a.first, sfa, b.first, sfb, d, grouped_layout,
|
||||
num_groups, m, n, k, gran_k_a, gran_k_b, major_a, major_b,
|
||||
@@ -230,8 +230,8 @@ static void m_grouped_fp8_fp4_gemm_nt_masked(const std::pair<torch::Tensor, torc
|
||||
const std::string& compiled_dims,
|
||||
const bool& disable_ue8m0_cast) {
|
||||
// Shape must be `[G, M, K] @ [G, N, K].mT`
|
||||
const auto& major_a = get_major_type_ab(a.first);
|
||||
const auto& major_b = get_major_type_ab(b.first);
|
||||
const auto major_a = get_major_type_ab(a.first);
|
||||
const auto major_b = get_major_type_ab(b.first);
|
||||
DG_HOST_ASSERT(major_a == cute::UMMA::Major::K and major_b == cute::UMMA::Major::K);
|
||||
DG_HOST_ASSERT(masked_m.is_contiguous());
|
||||
|
||||
@@ -256,7 +256,7 @@ static void m_grouped_fp8_fp4_gemm_nt_masked(const std::pair<torch::Tensor, torc
|
||||
|
||||
// Dispatch implementation
|
||||
if (arch_major == 9 and sfa.scalar_type() == torch::kFloat) {
|
||||
const auto& major_sfb = get_major_type_ab(sfb);
|
||||
const auto major_sfb = get_major_type_ab(sfb);
|
||||
sm90_m_grouped_fp8_gemm_masked_1d2d(a.first, sfa, b.first, sfb, d, masked_m,
|
||||
num_groups, m, n, k, expected_m, major_a, major_b, major_sfb, compiled_dims);
|
||||
} else if (arch_major == 10 and sfa.scalar_type() == torch::kInt) {
|
||||
@@ -277,12 +277,15 @@ static void k_grouped_fp8_gemm_tn_contiguous(const std::pair<torch::Tensor, torc
|
||||
const std::tuple<int, int, int>& recipe,
|
||||
const std::string& compiled_dims) {
|
||||
// Must be 1D1D kernel
|
||||
DG_HOST_ASSERT(recipe == std::make_tuple(1, 1, 128));
|
||||
DG_HOST_ASSERT(std::get<0>(recipe) == 1 and std::get<1>(recipe) == 1);
|
||||
|
||||
const int gran_k = std::get<2>(recipe);
|
||||
DG_HOST_ASSERT(gran_k == 32 or gran_k == 128);
|
||||
|
||||
// Shape checks
|
||||
const auto& [num_groups, m, n] = get_shape<3>(d);
|
||||
const auto& [sum_k_ , m_] = get_shape<2>(a.first);
|
||||
const auto& [sum_k__, n_] = get_shape<2>(b.first);
|
||||
const auto [num_groups, m, n] = get_shape<3>(d);
|
||||
const auto [sum_k_ , m_] = get_shape<2>(a.first);
|
||||
const auto [sum_k__, n_] = get_shape<2>(b.first);
|
||||
const int sum_k = std::accumulate(ks.begin(), ks.end(), 0);
|
||||
DG_HOST_ASSERT(m == m_ and n == n_ and sum_k == sum_k_ and sum_k == sum_k__);
|
||||
|
||||
@@ -297,13 +300,13 @@ static void k_grouped_fp8_gemm_tn_contiguous(const std::pair<torch::Tensor, torc
|
||||
return;
|
||||
|
||||
// Transform SF with padding
|
||||
const auto& sfa = layout::transform_k_grouped_sf_into_required_layout(a.second, ks, ks_tensor, recipe);
|
||||
const auto& sfb = layout::transform_k_grouped_sf_into_required_layout(b.second, ks, ks_tensor, recipe);
|
||||
const auto sfa = layout::transform_k_grouped_sf_into_required_layout(a.second, ks, ks_tensor, recipe);
|
||||
const auto sfb = layout::transform_k_grouped_sf_into_required_layout(b.second, ks, ks_tensor, recipe);
|
||||
|
||||
// Dispatch implementation
|
||||
const auto& arch_major = device_runtime->get_arch_major();
|
||||
const auto arch_major = device_runtime->get_arch_major();
|
||||
if (arch_major == 10) {
|
||||
sm100_k_grouped_fp8_gemm_1d1d(a.first, sfa, b.first, sfb, c, d, m, n, ks, ks_tensor,
|
||||
sm100_k_grouped_fp8_gemm_1d1d(a.first, sfa, b.first, sfb, c, d, m, n, ks, ks_tensor, gran_k,
|
||||
cute::UMMA::Major::MN, cute::UMMA::Major::MN, compiled_dims);
|
||||
} else {
|
||||
DG_HOST_UNREACHABLE("Unsupported architecture");
|
||||
@@ -322,9 +325,9 @@ static void k_grouped_fp8_gemm_nt_contiguous(const std::pair<torch::Tensor, torc
|
||||
DG_HOST_ASSERT(recipe == std::make_tuple(1, 1, 128));
|
||||
|
||||
// Shape checks
|
||||
const auto& [num_groups, m, n] = get_shape<3>(d);
|
||||
const auto& sum_mk = a.first.numel();
|
||||
const auto& sum_nk = b.first.numel();
|
||||
const auto [num_groups, m, n] = get_shape<3>(d);
|
||||
const auto sum_mk = a.first.numel();
|
||||
const auto sum_nk = b.first.numel();
|
||||
const int sum_k = std::accumulate(ks.begin(), ks.end(), 0);
|
||||
DG_HOST_ASSERT(sum_mk == static_cast<int64_t>(sum_k) * m);
|
||||
DG_HOST_ASSERT(sum_nk == static_cast<int64_t>(sum_k) * n);
|
||||
@@ -340,17 +343,17 @@ static void k_grouped_fp8_gemm_nt_contiguous(const std::pair<torch::Tensor, torc
|
||||
return;
|
||||
|
||||
// Transform SF with padding
|
||||
const auto& sfa = layout::transform_k_grouped_sf_into_required_layout(a.second, ks, ks_tensor, recipe);
|
||||
const auto& sfb = layout::transform_k_grouped_sf_into_required_layout(b.second, ks, ks_tensor, recipe);
|
||||
const auto sfa = layout::transform_k_grouped_sf_into_required_layout(a.second, ks, ks_tensor, recipe);
|
||||
const auto sfb = layout::transform_k_grouped_sf_into_required_layout(b.second, ks, ks_tensor, recipe);
|
||||
|
||||
// Allocate tensormap buffer
|
||||
// `4` means the double buffering for both A and B operands (2 * 2)
|
||||
const auto& num_sms = device_runtime->get_num_sms();
|
||||
const auto& tensor_map_buffer = torch::empty({num_sms * 4 * static_cast<int>(sizeof(CUtensorMap))},
|
||||
a.first.options().dtype(torch::kByte));
|
||||
const auto num_sms = device_runtime->get_num_sms();
|
||||
const auto tensor_map_buffer = torch::empty({num_sms * 4 * static_cast<int>(sizeof(CUtensorMap))},
|
||||
a.first.options().dtype(torch::kByte));
|
||||
|
||||
// Dispatch implementation
|
||||
const auto& arch_major = device_runtime->get_arch_major();
|
||||
const auto arch_major = device_runtime->get_arch_major();
|
||||
if (arch_major == 9) {
|
||||
sm90_k_grouped_fp8_gemm_1d1d(a.first, sfa, b.first, sfb, c, d, m, n, ks, ks_tensor, tensor_map_buffer,
|
||||
cute::UMMA::Major::K, cute::UMMA::Major::K, compiled_dims);
|
||||
@@ -367,16 +370,16 @@ static void bf16_gemm_nt(const torch::Tensor& a,
|
||||
const std::optional<torch::Tensor>& c,
|
||||
const std::string& compiled_dims) {
|
||||
// Shape must be `[M, K] @ [N, K].T`
|
||||
const auto& major_a = get_major_type_ab(a);
|
||||
const auto& major_b = get_major_type_ab(b);
|
||||
const auto major_a = get_major_type_ab(a);
|
||||
const auto major_b = get_major_type_ab(b);
|
||||
|
||||
// C/D must be N-major
|
||||
check_major_type_cd(d);
|
||||
|
||||
// Type and shape checks
|
||||
const auto& [m , k ] = get_shape<2>(a);
|
||||
const auto& [n , k_] = get_shape<2>(b);
|
||||
const auto& [m_, n_] = get_shape<2>(d);
|
||||
const auto [m , k ] = get_shape<2>(a);
|
||||
const auto [n , k_] = get_shape<2>(b);
|
||||
const auto [m_, n_] = get_shape<2>(d);
|
||||
DG_HOST_ASSERT(m == m_ and n == n_ and k == k_);
|
||||
DG_HOST_ASSERT(a.scalar_type() == torch::kBFloat16);
|
||||
DG_HOST_ASSERT(b.scalar_type() == torch::kBFloat16);
|
||||
@@ -387,7 +390,7 @@ static void bf16_gemm_nt(const torch::Tensor& a,
|
||||
return;
|
||||
|
||||
// Dispatch into different implements
|
||||
const auto& arch_major = device_runtime->get_arch_major();
|
||||
const auto arch_major = device_runtime->get_arch_major();
|
||||
if (arch_major == 9) {
|
||||
sm90_bf16_gemm(a, b, c, d, m, n, k, major_a, major_b, compiled_dims);
|
||||
} else if (arch_major == 10) {
|
||||
@@ -427,15 +430,15 @@ static void m_grouped_bf16_gemm_nt_contiguous(const torch::Tensor& a, const torc
|
||||
const bool& use_psum_layout,
|
||||
const std::optional<int>& expected_m_for_psum_layout) {
|
||||
// Shape must be `[M, K] @ [G, N, K].mT`
|
||||
const auto& major_a = get_major_type_ab(a);
|
||||
const auto& major_b = get_major_type_ab(b);
|
||||
const auto major_a = get_major_type_ab(a);
|
||||
const auto major_b = get_major_type_ab(b);
|
||||
DG_HOST_ASSERT(major_a == cute::UMMA::Major::K);
|
||||
DG_HOST_ASSERT(grouped_layout.is_contiguous());
|
||||
|
||||
// Type and shape checks
|
||||
const auto& [m, k] = get_shape<2>(a);
|
||||
const auto& [num_groups, n, k_] = get_shape<3>(b);
|
||||
const auto& [m_, n_] = get_shape<2>(d);
|
||||
const auto [m, k] = get_shape<2>(a);
|
||||
const auto [num_groups, n, k_] = get_shape<3>(b);
|
||||
const auto [m_, n_] = get_shape<2>(d);
|
||||
DG_HOST_ASSERT(m == m_ and n == n_ and k == k_);
|
||||
DG_HOST_ASSERT(n > 0 and k > 0 and num_groups > 0);
|
||||
DG_HOST_ASSERT(a.scalar_type() == torch::kBFloat16);
|
||||
@@ -445,10 +448,10 @@ static void m_grouped_bf16_gemm_nt_contiguous(const torch::Tensor& a, const torc
|
||||
|
||||
// Layout checks
|
||||
if (use_psum_layout) {
|
||||
const auto& [num_groups_] = get_shape<1>(grouped_layout);
|
||||
const auto [num_groups_] = get_shape<1>(grouped_layout);
|
||||
DG_HOST_ASSERT(num_groups == num_groups_);
|
||||
} else {
|
||||
const auto& [m__] = get_shape<1>(grouped_layout);
|
||||
const auto [m__] = get_shape<1>(grouped_layout);
|
||||
DG_HOST_ASSERT(m == m__);
|
||||
DG_HOST_ASSERT(not expected_m_for_psum_layout.has_value());
|
||||
}
|
||||
@@ -461,11 +464,11 @@ static void m_grouped_bf16_gemm_nt_contiguous(const torch::Tensor& a, const torc
|
||||
return;
|
||||
|
||||
// Dispatch implementation
|
||||
const auto& arch_major = device_runtime->get_arch_major();
|
||||
const auto arch_major = device_runtime->get_arch_major();
|
||||
if (arch_major == 9) {
|
||||
DG_HOST_ASSERT(not use_psum_layout);
|
||||
sm90_m_grouped_bf16_gemm_contiguous(a, b, d, grouped_layout,
|
||||
num_groups, m, n, k, major_a, major_b, compiled_dims);
|
||||
num_groups, m, n, k, major_a, major_b, compiled_dims,
|
||||
use_psum_layout, expected_m_for_psum_layout);
|
||||
} else if (arch_major == 10) {
|
||||
sm100_m_grouped_bf16_gemm_contiguous(a, b, d, grouped_layout,
|
||||
num_groups, m, n, k, major_a, major_b, compiled_dims,
|
||||
@@ -487,16 +490,16 @@ static void m_grouped_bf16_gemm_nt_masked(const torch::Tensor& a, const torch::T
|
||||
const torch::Tensor& d, const torch::Tensor& masked_m,
|
||||
const int& expected_m, const std::string& compiled_dims) {
|
||||
// Shape must be `[G, M, K] @ [G, N, K].mT`
|
||||
const auto& major_a = get_major_type_ab(a);
|
||||
const auto& major_b = get_major_type_ab(b);
|
||||
const auto major_a = get_major_type_ab(a);
|
||||
const auto major_b = get_major_type_ab(b);
|
||||
DG_HOST_ASSERT(major_a == cute::UMMA::Major::K and major_b == cute::UMMA::Major::K);
|
||||
DG_HOST_ASSERT(masked_m.is_contiguous());
|
||||
|
||||
// Type and shape checks
|
||||
const auto& [num_groups, m, k] = get_shape<3>(a);
|
||||
const auto& [num_groups_, n, k_] = get_shape<3>(b);
|
||||
const auto& [num_groups__, m_, n_] = get_shape<3>(d);
|
||||
const auto& num_groups___ = static_cast<int>(masked_m.numel());
|
||||
const auto [num_groups, m, k] = get_shape<3>(a);
|
||||
const auto [num_groups_, n, k_] = get_shape<3>(b);
|
||||
const auto [num_groups__, m_, n_] = get_shape<3>(d);
|
||||
const auto num_groups___ = static_cast<int>(masked_m.numel());
|
||||
DG_HOST_ASSERT(num_groups == num_groups_ and num_groups == num_groups__ and num_groups == num_groups___);
|
||||
DG_HOST_ASSERT(m == m_ and n == n_ and k == k_);
|
||||
DG_HOST_ASSERT(expected_m > 0 and m > 0 and n > 0 and k > 0 and num_groups > 0);
|
||||
@@ -509,7 +512,7 @@ static void m_grouped_bf16_gemm_nt_masked(const torch::Tensor& a, const torch::T
|
||||
check_major_type_cd(d);
|
||||
|
||||
// Dispatch implementation
|
||||
const auto& arch_major = device_runtime->get_arch_major();
|
||||
const auto arch_major = device_runtime->get_arch_major();
|
||||
if (arch_major == 9) {
|
||||
sm90_bf16_m_grouped_gemm_masked(a, b, d, masked_m,
|
||||
num_groups, m, n, k, expected_m, major_a, major_b, compiled_dims);
|
||||
@@ -529,9 +532,9 @@ static void k_grouped_bf16_gemm_tn_contiguous(const torch::Tensor& a,
|
||||
const std::optional<torch::Tensor>& c,
|
||||
const std::string& compiled_dims) {
|
||||
// Shape checks
|
||||
const auto& [num_groups, m, n] = get_shape<3>(d);
|
||||
const auto& [sum_k_ , m_] = get_shape<2>(a);
|
||||
const auto& [sum_k__, n_] = get_shape<2>(b);
|
||||
const auto [num_groups, m, n] = get_shape<3>(d);
|
||||
const auto [sum_k_ , m_] = get_shape<2>(a);
|
||||
const auto [sum_k__, n_] = get_shape<2>(b);
|
||||
const int sum_k = std::accumulate(ks.begin(), ks.end(), 0);
|
||||
DG_HOST_ASSERT(m == m_ and n == n_ and sum_k == sum_k_ and sum_k == sum_k__);
|
||||
|
||||
@@ -546,7 +549,7 @@ static void k_grouped_bf16_gemm_tn_contiguous(const torch::Tensor& a,
|
||||
return;
|
||||
|
||||
// Dispatch implementation
|
||||
const auto& arch_major = device_runtime->get_arch_major();
|
||||
const auto arch_major = device_runtime->get_arch_major();
|
||||
if (arch_major == 9) {
|
||||
sm90_bf16_k_grouped_gemm(a, b, c, d, m, n, ks, ks_tensor,
|
||||
cute::UMMA::Major::MN, cute::UMMA::Major::MN, compiled_dims);
|
||||
@@ -562,20 +565,20 @@ static void k_grouped_bf16_gemm_tn_contiguous(const torch::Tensor& a,
|
||||
static void cublaslt_gemm_nt(const torch::Tensor& a, const torch::Tensor& b,
|
||||
const torch::Tensor& d, const std::optional<torch::Tensor>& c) {
|
||||
// Shape must be `[M, K] @ [N, K].T`
|
||||
const auto& major_a = get_major_type_ab(a);
|
||||
const auto& major_b = get_major_type_ab(b);
|
||||
const auto major_a = get_major_type_ab(a);
|
||||
const auto major_b = get_major_type_ab(b);
|
||||
|
||||
// Type and shape checks
|
||||
const auto& [m , k ] = get_shape<2>(a);
|
||||
const auto& [n , k_] = get_shape<2>(b);
|
||||
const auto& [m_, n_] = get_shape<2>(d);
|
||||
const auto [m , k ] = get_shape<2>(a);
|
||||
const auto [n , k_] = get_shape<2>(b);
|
||||
const auto [m_, n_] = get_shape<2>(d);
|
||||
DG_HOST_ASSERT(m == m_ and n == n_ and k == k_);
|
||||
|
||||
// Early return for trivial cases
|
||||
if (early_return(m, n, k, d, c))
|
||||
return;
|
||||
|
||||
cublaslt_gemm(a, b, c, d, m, n, k, major_a, major_b);
|
||||
cublaslt_gemm(a, b, d, m, n, k, major_a, major_b, c.has_value());
|
||||
}
|
||||
|
||||
static void cublaslt_gemm_nn(const torch::Tensor& a, const torch::Tensor& b,
|
||||
|
||||
Reference in New Issue
Block a user