Multiple updates and refactorings (#231)

This commit is contained in:
Ray Wang
2025-11-21 17:49:47 +08:00
committed by GitHub
parent bb4424aad4
commit 38f8ef73a4
80 changed files with 3767 additions and 2103 deletions
+46 -33
View File
@@ -1,5 +1,8 @@
import os
import subprocess
import torch
from torch.version import cuda as cuda_version
from packaging import version
# Set some default environment provided at setup
try:
@@ -12,53 +15,63 @@ except ImportError:
pass
# Configs
import deep_gemm_cpp
from deep_gemm_cpp import (
from . import _C
from ._C import (
set_num_sms,
get_num_sms,
set_tc_util,
get_tc_util,
)
# Kernels
from deep_gemm_cpp import (
# FP8 GEMMs
fp8_gemm_nt, fp8_gemm_nn,
fp8_gemm_tn, fp8_gemm_tt,
fp8_gemm_nt_skip_head_mid,
m_grouped_fp8_gemm_nt_contiguous,
m_grouped_fp8_gemm_nn_contiguous,
m_grouped_fp8_gemm_nt_masked,
k_grouped_fp8_gemm_nt_contiguous,
k_grouped_fp8_gemm_tn_contiguous,
# BF16 GEMMs
bf16_gemm_nt, bf16_gemm_nn,
bf16_gemm_tn, bf16_gemm_tt,
m_grouped_bf16_gemm_nt_contiguous,
m_grouped_bf16_gemm_nt_masked,
# cuBLASLt GEMMs
# cuBLASLt Kernels
from ._C import (
cublaslt_gemm_nt, cublaslt_gemm_nn,
cublaslt_gemm_tn, cublaslt_gemm_tt,
# Einsum kernels
einsum,
# Attention kernels
fp8_mqa_logits,
get_paged_mqa_logits_metadata,
fp8_paged_mqa_logits,
# Layout kernels
transform_sf_into_required_layout
)
# Some alias for legacy supports
# TODO: remove these later
fp8_m_grouped_gemm_nt_masked = m_grouped_fp8_gemm_nt_masked
bf16_m_grouped_gemm_nt_masked = m_grouped_bf16_gemm_nt_masked
if version.parse(cuda_version) >= version.parse('12.1'):
# DeepGEMM Kernels
from ._C import (
# FP8 GEMMs
fp8_gemm_nt, fp8_gemm_nn,
fp8_gemm_tn, fp8_gemm_tt,
fp8_gemm_nt_skip_head_mid,
m_grouped_fp8_gemm_nt_contiguous,
m_grouped_fp8_gemm_nn_contiguous,
m_grouped_fp8_gemm_nt_masked,
k_grouped_fp8_gemm_nt_contiguous,
k_grouped_fp8_gemm_tn_contiguous,
# BF16 GEMMs
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 kernels
einsum,
fp8_einsum,
# Attention kernels
fp8_mqa_logits,
get_paged_mqa_logits_metadata,
fp8_paged_mqa_logits,
# Layout kernels
transform_sf_into_required_layout,
get_mk_alignment_for_contiguous_layout
)
# Some alias for legacy supports
# TODO: remove these later
fp8_m_grouped_gemm_nt_masked = m_grouped_fp8_gemm_nt_masked
bf16_m_grouped_gemm_nt_masked = m_grouped_bf16_gemm_nt_masked
# Some utils
from . import testing
from . import utils
from .utils import *
# Legacy Triton kernels for A100
from . import legacy
# Initialize CPP modules
def _find_cuda_home() -> str:
@@ -79,9 +92,9 @@ def _find_cuda_home() -> str:
return cuda_home
deep_gemm_cpp.init(
_C.init(
os.path.dirname(os.path.abspath(__file__)), # Library root directory path
_find_cuda_home() # CUDA home
)
__version__ = '2.1.1'
__version__ = '2.2.0'