[kernel slimming] Move fast_hadamard_transform to jit_kernel (#18475)
This commit is contained in:
@@ -106,15 +106,6 @@ FetchContent_Declare(
|
||||
)
|
||||
FetchContent_Populate(repo-mscclpp)
|
||||
|
||||
# fast-hadamard-transform
|
||||
FetchContent_Declare(
|
||||
repo-fast-hadamard-transform
|
||||
GIT_REPOSITORY https://github.com/sgl-project/fast-hadamard-transform.git
|
||||
GIT_TAG 48f3c13764dc2ec662ade842a4696a90a137f1bc
|
||||
GIT_SHALLOW OFF
|
||||
)
|
||||
FetchContent_Populate(repo-fast-hadamard-transform)
|
||||
|
||||
# ccache option
|
||||
option(ENABLE_CCACHE "Whether to use ccache" ON)
|
||||
find_program(CCACHE_FOUND ccache)
|
||||
@@ -343,9 +334,6 @@ set(SOURCES
|
||||
"${repo-flashinfer_SOURCE_DIR}/csrc/renorm.cu"
|
||||
"${repo-flashinfer_SOURCE_DIR}/csrc/sampling.cu"
|
||||
|
||||
"${repo-fast-hadamard-transform_SOURCE_DIR}/csrc/fast_hadamard_transform_cuda.cu"
|
||||
"${repo-fast-hadamard-transform_SOURCE_DIR}/csrc/fast_hadamard_transform.cpp"
|
||||
|
||||
"${repo-flash-attention_SOURCE_DIR}/csrc/flash_attn/src/flash_fwd_sparse_hdim128_bf16_causal_sm80.cu"
|
||||
"${repo-flash-attention_SOURCE_DIR}/csrc/flash_attn/src/flash_fwd_sparse_hdim128_bf16_sm80.cu"
|
||||
"${repo-flash-attention_SOURCE_DIR}/csrc/flash_attn/src/flash_fwd_sparse_hdim128_fp16_causal_sm80.cu"
|
||||
|
||||
@@ -591,24 +591,6 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
|
||||
"es_sm100_mxfp8_blockscaled_grouped_quant(Tensor input, Tensor problem_sizes, Tensor expert_offsets, Tensor "
|
||||
"blockscale_offsets, Tensor quant_output, Tensor scale_factor) -> () ");
|
||||
m.impl("es_sm100_mxfp8_blockscaled_grouped_quant", &es_sm100_mxfp8_blockscaled_grouped_quant);
|
||||
|
||||
/*
|
||||
* From fast-hadamard-transform
|
||||
*/
|
||||
m.def("fast_hadamard_transform(Tensor x, float scale) -> Tensor");
|
||||
m.impl("fast_hadamard_transform", torch::kCUDA, &fast_hadamard_transform);
|
||||
|
||||
m.def("fast_hadamard_transform_12N(Tensor x, float scale) -> Tensor");
|
||||
m.impl("fast_hadamard_transform_12N", torch::kCUDA, &fast_hadamard_transform_12N);
|
||||
|
||||
m.def("fast_hadamard_transform_20N(Tensor x, float scale) -> Tensor");
|
||||
m.impl("fast_hadamard_transform_20N", torch::kCUDA, &fast_hadamard_transform_20N);
|
||||
|
||||
m.def("fast_hadamard_transform_28N(Tensor x, float scale) -> Tensor");
|
||||
m.impl("fast_hadamard_transform_28N", torch::kCUDA, &fast_hadamard_transform_28N);
|
||||
|
||||
m.def("fast_hadamard_transform_40N(Tensor x, float scale) -> Tensor");
|
||||
m.impl("fast_hadamard_transform_40N", torch::kCUDA, &fast_hadamard_transform_40N);
|
||||
}
|
||||
|
||||
REGISTER_EXTENSION(common_ops)
|
||||
|
||||
@@ -936,15 +936,6 @@ void es_sm100_mxfp8_blockscaled_grouped_quant(
|
||||
torch::Tensor& quant_output,
|
||||
torch::Tensor& scale_factor);
|
||||
|
||||
/*
|
||||
* From fast-hadamard-transform
|
||||
*/
|
||||
torch::Tensor fast_hadamard_transform(torch::Tensor& x, double scale);
|
||||
torch::Tensor fast_hadamard_transform_12N(torch::Tensor& x, double scale);
|
||||
torch::Tensor fast_hadamard_transform_20N(torch::Tensor& x, double scale);
|
||||
torch::Tensor fast_hadamard_transform_28N(torch::Tensor& x, double scale);
|
||||
torch::Tensor fast_hadamard_transform_40N(torch::Tensor& x, double scale);
|
||||
|
||||
/*
|
||||
* From flashmla
|
||||
*/
|
||||
|
||||
@@ -65,13 +65,6 @@ from sgl_kernel.gemm import (
|
||||
silu_and_mul_scaled_fp4_grouped_quant,
|
||||
)
|
||||
from sgl_kernel.grammar import apply_token_bitmask_inplace_cuda
|
||||
from sgl_kernel.hadamard import (
|
||||
hadamard_transform,
|
||||
hadamard_transform_12n,
|
||||
hadamard_transform_20n,
|
||||
hadamard_transform_28n,
|
||||
hadamard_transform_40n,
|
||||
)
|
||||
from sgl_kernel.kvcacheio import (
|
||||
transfer_kv_all_layer,
|
||||
transfer_kv_all_layer_mla,
|
||||
|
||||
@@ -1,21 +0,0 @@
|
||||
import torch
|
||||
|
||||
|
||||
def hadamard_transform(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
|
||||
return torch.ops.sgl_kernel.fast_hadamard_transform.default(x, scale)
|
||||
|
||||
|
||||
def hadamard_transform_12n(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
|
||||
return torch.ops.sgl_kernel.fast_hadamard_transform_12N.default(x, scale)
|
||||
|
||||
|
||||
def hadamard_transform_20n(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
|
||||
return torch.ops.sgl_kernel.fast_hadamard_transform_20N.default(x, scale)
|
||||
|
||||
|
||||
def hadamard_transform_28n(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
|
||||
return torch.ops.sgl_kernel.fast_hadamard_transform_28N.default(x, scale)
|
||||
|
||||
|
||||
def hadamard_transform_40n(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
|
||||
return torch.ops.sgl_kernel.fast_hadamard_transform_40N.default(x, scale)
|
||||
@@ -5,7 +5,14 @@ import torch
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange, repeat
|
||||
from scipy.linalg import hadamard
|
||||
from sgl_kernel import hadamard_transform
|
||||
|
||||
try:
|
||||
from sgl_kernel import hadamard_transform
|
||||
except Exception:
|
||||
pytest.skip(
|
||||
"sgl-kernel hadamard interface was removed (migrated to jit_kernel)",
|
||||
allow_module_level=True,
|
||||
)
|
||||
|
||||
|
||||
def hadamard_transform_ref(x, scale=1.0):
|
||||
|
||||
Reference in New Issue
Block a user