Add new moe wna16 marlin gemm (#14122)
This commit is contained in:
@@ -299,7 +299,8 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
|
|||||||
*/
|
*/
|
||||||
m.def(
|
m.def(
|
||||||
"moe_wna16_marlin_gemm(Tensor! a, Tensor? c_or_none,"
|
"moe_wna16_marlin_gemm(Tensor! a, Tensor? c_or_none,"
|
||||||
"Tensor! b_q_weight, Tensor! b_scales, Tensor? b_zeros_or_none,"
|
"Tensor! b_q_weight, Tensor? b_bias_or_none, Tensor! b_scales,"
|
||||||
|
"Tensor? global_scale_or_none, Tensor? b_zeros_or_none,"
|
||||||
"Tensor? g_idx_or_none, Tensor? perm_or_none, Tensor! workspace,"
|
"Tensor? g_idx_or_none, Tensor? perm_or_none, Tensor! workspace,"
|
||||||
"Tensor sorted_token_ids,"
|
"Tensor sorted_token_ids,"
|
||||||
"Tensor! expert_ids, Tensor! num_tokens_past_padded,"
|
"Tensor! expert_ids, Tensor! num_tokens_past_padded,"
|
||||||
|
|||||||
@@ -449,6 +449,51 @@ __device__ inline void dequant_fp8_scales<nv_bfloat162>(int q, nv_bfloat162* fra
|
|||||||
q <<= 8;
|
q <<= 8;
|
||||||
int Out2 = ((q & 0x80008000) >> 1) | ((q & MASK) >> RIGHT_SHIFT);
|
int Out2 = ((q & 0x80008000) >> 1) | ((q & MASK) >> RIGHT_SHIFT);
|
||||||
|
|
||||||
|
// Note: reverse indexing is intentional because weights are permuted
|
||||||
|
frag_b[1] = *reinterpret_cast<const nv_bfloat162*>(&Out1);
|
||||||
|
frag_b[0] = *reinterpret_cast<const nv_bfloat162*>(&Out2);
|
||||||
|
};
|
||||||
|
|
||||||
|
// New version with s_type_id parameter for marlin_moe_wna16_v2
|
||||||
|
template <typename scalar_t2, sglang::ScalarTypeId s_type_id>
|
||||||
|
__device__ inline void dequant_fp8_scales(int q, scalar_t2* frag_b);
|
||||||
|
|
||||||
|
template <>
|
||||||
|
__device__ inline void dequant_fp8_scales<half2, sglang::kFE4M3fn.id()>(int q, half2* frag_b) {
|
||||||
|
int Out1 = (q & 0xFF00FF00) >> 1;
|
||||||
|
;
|
||||||
|
q <<= 8;
|
||||||
|
int Out2 = (q & 0xFF00FF00) >> 1;
|
||||||
|
|
||||||
|
// Note: reverse indexing is intentional because weights are permuted
|
||||||
|
frag_b[1] = *reinterpret_cast<const half2*>(&Out1);
|
||||||
|
frag_b[0] = *reinterpret_cast<const half2*>(&Out2);
|
||||||
|
};
|
||||||
|
|
||||||
|
template <>
|
||||||
|
__device__ inline void dequant_fp8_scales<nv_bfloat162, sglang::kFE4M3fn.id()>(int q, nv_bfloat162* frag_b) {
|
||||||
|
constexpr int FP8_EXPONENT = 4, BF16_EXPONENT = 8;
|
||||||
|
constexpr int RIGHT_SHIFT = BF16_EXPONENT - FP8_EXPONENT;
|
||||||
|
constexpr int MASK = 0x7F007F00;
|
||||||
|
|
||||||
|
// Extract and shift FP8 values to BF16 format
|
||||||
|
int Out1 = ((q & 0x80008000) >> 1) | ((q & MASK) >> RIGHT_SHIFT);
|
||||||
|
q <<= 8;
|
||||||
|
int Out2 = ((q & 0x80008000) >> 1) | ((q & MASK) >> RIGHT_SHIFT);
|
||||||
|
|
||||||
|
// Note: reverse indexing is intentional because weights are permuted
|
||||||
|
frag_b[1] = *reinterpret_cast<const nv_bfloat162*>(&Out1);
|
||||||
|
frag_b[0] = *reinterpret_cast<const nv_bfloat162*>(&Out2);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <>
|
||||||
|
__device__ inline void dequant_fp8_scales<nv_bfloat162, sglang::kFE8M0fnu.id()>(int q, nv_bfloat162* frag_b) {
|
||||||
|
// In this conversion, 2 ** -127 in FP8E8M0 would become 0 in BF16,
|
||||||
|
// but we assume that such a extreme value would not occur in real models.
|
||||||
|
int Out1 = (q & 0xFF00FF00) >> 1;
|
||||||
|
q <<= 7;
|
||||||
|
int Out2 = q & 0x7F807F80;
|
||||||
|
|
||||||
// Note: reverse indexing is intentional because weights are permuted
|
// Note: reverse indexing is intentional because weights are permuted
|
||||||
frag_b[1] = *reinterpret_cast<const nv_bfloat162*>(&Out1);
|
frag_b[1] = *reinterpret_cast<const nv_bfloat162*>(&Out1);
|
||||||
frag_b[0] = *reinterpret_cast<const nv_bfloat162*>(&Out2);
|
frag_b[0] = *reinterpret_cast<const nv_bfloat162*>(&Out2);
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
# SPDX-License-Identifier: Apache-2.0
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||||
import glob
|
import glob
|
||||||
import itertools
|
import itertools
|
||||||
import os
|
import os
|
||||||
@@ -11,9 +12,6 @@ FILE_HEAD = """
|
|||||||
// clang-format off
|
// clang-format off
|
||||||
#pragma once
|
#pragma once
|
||||||
|
|
||||||
#include "kernel.h"
|
|
||||||
#include "marlin_template.h"
|
|
||||||
|
|
||||||
namespace MARLIN_NAMESPACE_NAME {
|
namespace MARLIN_NAMESPACE_NAME {
|
||||||
""".strip()
|
""".strip()
|
||||||
|
|
||||||
@@ -21,14 +19,13 @@ TEMPLATE = (
|
|||||||
"template __global__ void Marlin<"
|
"template __global__ void Marlin<"
|
||||||
"{{scalar_t}}, "
|
"{{scalar_t}}, "
|
||||||
"{{w_type_id}}, "
|
"{{w_type_id}}, "
|
||||||
|
"{{s_type_id}}, "
|
||||||
"{{threads}}, "
|
"{{threads}}, "
|
||||||
"{{thread_m_blocks}}, "
|
"{{thread_m_blocks}}, "
|
||||||
"{{thread_n_blocks}}, "
|
"{{thread_n_blocks}}, "
|
||||||
"{{thread_k_blocks}}, "
|
"{{thread_k_blocks}}, "
|
||||||
"{{'true' if m_block_size_8 else 'false'}}, "
|
"{{'true' if m_block_size_8 else 'false'}}, "
|
||||||
"{{stages}}, "
|
"{{stages}}, "
|
||||||
"{{'true' if has_act_order else 'false'}}, "
|
|
||||||
"{{'true' if has_zp else 'false'}}, "
|
|
||||||
"{{group_blocks}}, "
|
"{{group_blocks}}, "
|
||||||
"{{'true' if is_zp_float else 'false'}}>"
|
"{{'true' if is_zp_float else 'false'}}>"
|
||||||
"( MARLIN_KERNEL_PARAMS );"
|
"( MARLIN_KERNEL_PARAMS );"
|
||||||
@@ -47,10 +44,15 @@ KERNEL_FILE_NAME = "kernel_marlin.cuh"
|
|||||||
|
|
||||||
# int8 with zero point case (sglang::kU8) is also supported,
|
# int8 with zero point case (sglang::kU8) is also supported,
|
||||||
# we don't add it to reduce wheel size.
|
# we don't add it to reduce wheel size.
|
||||||
SCALAR_TYPES = ["sglang::kU4", "sglang::kU4B8", "sglang::kU8B128"]
|
# Only keep the most commonly used types to reduce compilation time
|
||||||
THREAD_CONFIGS = [(128, 128, 256), (64, 256, 256), (64, 128, 128)]
|
SCALAR_TYPES = [
|
||||||
|
"sglang::kU4",
|
||||||
|
"sglang::kU4B8",
|
||||||
|
"sglang::kU8B128",
|
||||||
|
]
|
||||||
|
THREAD_CONFIGS = [(128, 128, 256), (64, 256, 256)]
|
||||||
|
|
||||||
THREAD_M_BLOCKS = [0.5, 1, 2, 3, 4]
|
THREAD_M_BLOCKS = [0.5, 1, 2, 4]
|
||||||
# group_blocks:
|
# group_blocks:
|
||||||
# = 0 : act order case
|
# = 0 : act order case
|
||||||
# = -1 : channelwise quantization
|
# = -1 : channelwise quantization
|
||||||
@@ -67,39 +69,65 @@ def remove_old_kernels():
|
|||||||
def generate_new_kernels():
|
def generate_new_kernels():
|
||||||
kernel_files = set()
|
kernel_files = set()
|
||||||
for scalar_type, dtype in itertools.product(SCALAR_TYPES, DTYPES):
|
for scalar_type, dtype in itertools.product(SCALAR_TYPES, DTYPES):
|
||||||
has_zp = "B" not in scalar_type
|
|
||||||
all_template_str_list = []
|
all_template_str_list = []
|
||||||
|
|
||||||
for group_blocks, m_blocks, thread_configs in itertools.product(
|
for group_blocks, m_blocks, thread_configs in itertools.product(
|
||||||
GROUP_BLOCKS, THREAD_M_BLOCKS, THREAD_CONFIGS
|
GROUP_BLOCKS, THREAD_M_BLOCKS, THREAD_CONFIGS
|
||||||
):
|
):
|
||||||
|
# act order case only support gptq-int4 and gptq-int8
|
||||||
has_act_order = group_blocks == 0
|
if group_blocks == 0 and scalar_type not in [
|
||||||
if has_zp and has_act_order:
|
"sglang::kU4B8",
|
||||||
|
"sglang::kU8B128",
|
||||||
|
]:
|
||||||
continue
|
continue
|
||||||
if thread_configs[2] == 256:
|
if thread_configs[2] == 256:
|
||||||
|
# for small batch (m_blocks == 1), we only need (128, 128, 256)
|
||||||
|
# for large batch (m_blocks > 1), we only need (64, 256, 256)
|
||||||
if m_blocks <= 1 and thread_configs[0] != 128:
|
if m_blocks <= 1 and thread_configs[0] != 128:
|
||||||
continue
|
continue
|
||||||
if m_blocks > 1 and thread_configs[0] != 64:
|
if m_blocks > 1 and thread_configs[0] != 64:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
# we only support channelwise quantization and group_size == 128
|
||||||
|
# for fp8
|
||||||
|
if scalar_type == "sglang::kFE4M3fn" and group_blocks not in [-1, 8]:
|
||||||
|
continue
|
||||||
|
# nvfp4 only supports group_size == 16
|
||||||
|
# mxfp4 only supports group_size == 32
|
||||||
|
if scalar_type == "sglang::kFE2M1f" and group_blocks not in [1, 2]:
|
||||||
|
continue
|
||||||
|
# other quantization methods don't support group_size = 16
|
||||||
|
if scalar_type != "sglang::kFE2M1f" and group_blocks == 1:
|
||||||
|
continue
|
||||||
|
|
||||||
k_blocks = thread_configs[0] // 16
|
k_blocks = thread_configs[0] // 16
|
||||||
n_blocks = thread_configs[1] // 16
|
n_blocks = thread_configs[1] // 16
|
||||||
threads = thread_configs[2]
|
threads = thread_configs[2]
|
||||||
|
|
||||||
c_dtype = "half" if dtype == "fp16" else "nv_bfloat16"
|
c_dtype = "half" if dtype == "fp16" else "nv_bfloat16"
|
||||||
|
|
||||||
|
if scalar_type == "sglang::kFE2M1f" and group_blocks == 1:
|
||||||
|
s_type = "sglang::kFE4M3fn"
|
||||||
|
elif scalar_type == "sglang::kFE2M1f" and group_blocks == 2:
|
||||||
|
s_type = "sglang::kFE8M0fnu"
|
||||||
|
if dtype == "fp16":
|
||||||
|
# we cannot safely dequantize e8m0 to fp16, so skip this
|
||||||
|
continue
|
||||||
|
elif dtype == "fp16":
|
||||||
|
s_type = "sglang::kFloat16"
|
||||||
|
elif dtype == "bf16":
|
||||||
|
s_type = "sglang::kBFloat16"
|
||||||
|
|
||||||
template_str = jinja2.Template(TEMPLATE).render(
|
template_str = jinja2.Template(TEMPLATE).render(
|
||||||
scalar_t=c_dtype,
|
scalar_t=c_dtype,
|
||||||
w_type_id=scalar_type + ".id()",
|
w_type_id=scalar_type + ".id()",
|
||||||
|
s_type_id=s_type + ".id()",
|
||||||
threads=threads,
|
threads=threads,
|
||||||
thread_m_blocks=max(m_blocks, 1),
|
thread_m_blocks=max(m_blocks, 1),
|
||||||
thread_n_blocks=n_blocks,
|
thread_n_blocks=n_blocks,
|
||||||
thread_k_blocks=k_blocks,
|
thread_k_blocks=k_blocks,
|
||||||
m_block_size_8=m_blocks == 0.5,
|
m_block_size_8=m_blocks == 0.5,
|
||||||
stages="pipe_stages",
|
stages="pipe_stages",
|
||||||
has_act_order=has_act_order,
|
|
||||||
has_zp=has_zp,
|
|
||||||
group_blocks=group_blocks,
|
group_blocks=group_blocks,
|
||||||
is_zp_float=False,
|
is_zp_float=False,
|
||||||
)
|
)
|
||||||
@@ -108,6 +136,7 @@ def generate_new_kernels():
|
|||||||
|
|
||||||
file_content = FILE_HEAD + "\n\n"
|
file_content = FILE_HEAD + "\n\n"
|
||||||
file_content += "\n\n".join(all_template_str_list) + "\n\n}\n"
|
file_content += "\n\n".join(all_template_str_list) + "\n\n}\n"
|
||||||
|
# Remove "sglang::" prefix (8 chars) from scalar_type for filename
|
||||||
filename = f"kernel_{dtype}_{scalar_type[8:].lower()}.cuh"
|
filename = f"kernel_{dtype}_{scalar_type[8:].lower()}.cuh"
|
||||||
|
|
||||||
with open(os.path.join(os.path.dirname(__file__), filename), "w") as f:
|
with open(os.path.join(os.path.dirname(__file__), filename), "w") as f:
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
#pragma once
|
|
||||||
|
|
||||||
#ifndef MARLIN_NAMESPACE_NAME
|
#ifndef MARLIN_NAMESPACE_NAME
|
||||||
#define MARLIN_NAMESPACE_NAME marlin_moe_wna16
|
#define MARLIN_NAMESPACE_NAME marlin_moe_wna16_v2
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
#include "gemm/marlin/marlin.cuh"
|
#include "gemm/marlin/marlin.cuh"
|
||||||
@@ -10,16 +9,18 @@
|
|||||||
|
|
||||||
#define MARLIN_KERNEL_PARAMS \
|
#define MARLIN_KERNEL_PARAMS \
|
||||||
const int4 *__restrict__ A, const int4 *__restrict__ B, int4 *__restrict__ C, int4 *__restrict__ C_tmp, \
|
const int4 *__restrict__ A, const int4 *__restrict__ B, int4 *__restrict__ C, int4 *__restrict__ C_tmp, \
|
||||||
const int4 *__restrict__ scales_ptr, const int4 *__restrict__ zp_ptr, const int *__restrict__ g_idx, \
|
const int4 *__restrict__ b_bias_ptr, const int4 *__restrict__ scales_ptr, \
|
||||||
|
const uint16_t *__restrict__ scale2_ptr, const int4 *__restrict__ zp_ptr, const int *__restrict__ g_idx, \
|
||||||
const int32_t *__restrict__ sorted_token_ids_ptr, const int32_t *__restrict__ expert_ids_ptr, \
|
const int32_t *__restrict__ sorted_token_ids_ptr, const int32_t *__restrict__ expert_ids_ptr, \
|
||||||
const int32_t *__restrict__ num_tokens_past_padded_ptr, const float *__restrict__ topk_weights_ptr, int top_k, \
|
const int32_t *__restrict__ num_tokens_past_padded_ptr, const float *__restrict__ topk_weights_ptr, int top_k, \
|
||||||
bool mul_topk_weights, bool is_ep, int num_groups, int prob_m, int prob_n, int prob_k, int *locks, \
|
bool mul_topk_weights, bool is_ep, int num_groups, int prob_m, int prob_n, int prob_k, int *locks, \
|
||||||
bool use_atomic_add, bool use_fp32_reduce
|
bool has_bias, bool use_atomic_add, bool use_fp32_reduce, int max_shared_mem
|
||||||
|
|
||||||
namespace MARLIN_NAMESPACE_NAME {
|
namespace MARLIN_NAMESPACE_NAME {
|
||||||
template <
|
template <
|
||||||
typename scalar_t, // compute dtype, half or nv_float16
|
typename scalar_t, // compute dtype, half or nv_float16
|
||||||
const sglang::ScalarTypeId w_type_id, // weight ScalarType id
|
const sglang::ScalarTypeId w_type_id, // weight ScalarType id
|
||||||
|
const sglang::ScalarTypeId s_type_id, // weight scale ScalarType id
|
||||||
const int threads, // number of threads in a threadblock
|
const int threads, // number of threads in a threadblock
|
||||||
const int thread_m_blocks, // number of 16x16 blocks in the m
|
const int thread_m_blocks, // number of 16x16 blocks in the m
|
||||||
// dimension (batchsize) of the
|
// dimension (batchsize) of the
|
||||||
@@ -30,8 +31,6 @@ template <
|
|||||||
// only works when thread_m_blocks == 1
|
// only works when thread_m_blocks == 1
|
||||||
const int stages, // number of stages for the async global->shared
|
const int stages, // number of stages for the async global->shared
|
||||||
// fetch pipeline
|
// fetch pipeline
|
||||||
const bool has_act_order, // whether act_order is enabled
|
|
||||||
const bool has_zp, // whether zero-points are enabled
|
|
||||||
const int group_blocks, // number of consecutive 16x16 blocks
|
const int group_blocks, // number of consecutive 16x16 blocks
|
||||||
// with a separate quantization scale
|
// with a separate quantization scale
|
||||||
const bool is_zp_float // is zero point of float16 type?
|
const bool is_zp_float // is zero point of float16 type?
|
||||||
|
|||||||
@@ -2,89 +2,38 @@
|
|||||||
// clang-format off
|
// clang-format off
|
||||||
#pragma once
|
#pragma once
|
||||||
|
|
||||||
#include "kernel.h"
|
|
||||||
#include "marlin_template.h"
|
|
||||||
|
|
||||||
namespace MARLIN_NAMESPACE_NAME {
|
namespace MARLIN_NAMESPACE_NAME {
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 1, 8, 8, true, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, true, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 1, 8, 4, true, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 1, 8, 8, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 2, 16, 4, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 1, 8, 4, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 4, 16, 4, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 2, 16, 4, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 2, 8, 4, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 3, 16, 4, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 3, 8, 4, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 4, 16, 4, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 4, 8, 4, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 1, 8, 8, true, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 1, 8, 4, true, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 1, 8, 8, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 1, 8, 4, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 2, 16, 4, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 2, 8, 4, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), sglang::kBFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 3, 16, 4, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 3, 8, 4, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 4, 16, 4, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 4, 8, 4, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 1, 8, 8, true, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 1, 8, 4, true, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 1, 8, 8, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 1, 8, 4, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 2, 16, 4, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 2, 8, 4, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 3, 16, 4, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 3, 8, 4, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 4, 16, 4, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 4, 8, 4, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 1, 8, 8, true, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 1, 8, 4, true, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 1, 8, 8, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 1, 8, 4, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 2, 16, 4, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 2, 8, 4, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 3, 16, 4, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 3, 8, 4, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 256, 4, 16, 4, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4.id(), 128, 4, 8, 4, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,109 +2,46 @@
|
|||||||
// clang-format off
|
// clang-format off
|
||||||
#pragma once
|
#pragma once
|
||||||
|
|
||||||
#include "kernel.h"
|
|
||||||
#include "marlin_template.h"
|
|
||||||
|
|
||||||
namespace MARLIN_NAMESPACE_NAME {
|
namespace MARLIN_NAMESPACE_NAME {
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 1, 8, 8, true, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 1, 8, 4, true, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 1, 8, 8, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 1, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 2, 16, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, true, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 2, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 3, 16, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 2, 16, 4, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 3, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 4, 16, 4, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 4, 16, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 4, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 1, 8, 8, true, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 1, 8, 4, true, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 1, 8, 8, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 1, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 2, 16, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 2, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 3, 16, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 3, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 4, 16, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 4, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), sglang::kBFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 1, 8, 8, true, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 1, 8, 4, true, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 1, 8, 8, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 1, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 2, 16, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 2, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 3, 16, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 3, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 4, 16, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 4, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 1, 8, 8, true, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 1, 8, 4, true, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 1, 8, 8, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 1, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 2, 16, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 2, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 3, 16, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 3, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 4, 16, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 4, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 1, 8, 8, true, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 1, 8, 4, true, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 1, 8, 8, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 1, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 2, 16, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 2, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 3, 16, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 3, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 256, 4, 16, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU4B8.id(), 128, 4, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,109 +2,46 @@
|
|||||||
// clang-format off
|
// clang-format off
|
||||||
#pragma once
|
#pragma once
|
||||||
|
|
||||||
#include "kernel.h"
|
|
||||||
#include "marlin_template.h"
|
|
||||||
|
|
||||||
namespace MARLIN_NAMESPACE_NAME {
|
namespace MARLIN_NAMESPACE_NAME {
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 1, 8, 8, true, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 1, 8, 4, true, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 1, 8, 8, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 1, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 2, 16, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, true, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 2, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 3, 16, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 2, 16, 4, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 3, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 4, 16, 4, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 4, 16, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 4, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 1, 8, 8, true, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 1, 8, 4, true, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 1, 8, 8, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 1, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 2, 16, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 2, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 3, 16, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 3, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 4, 16, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 4, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), sglang::kBFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 1, 8, 8, true, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 1, 8, 4, true, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 1, 8, 8, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 1, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 2, 16, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 2, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 3, 16, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 3, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 4, 16, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 4, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 1, 8, 8, true, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 1, 8, 4, true, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 1, 8, 8, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 1, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 2, 16, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 2, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 3, 16, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 3, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 4, 16, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 4, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 1, 8, 8, true, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 1, 8, 4, true, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 1, 8, 8, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 1, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 2, 16, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 2, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 3, 16, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 3, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 256, 4, 16, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<nv_bfloat16, sglang::kU8B128.id(), 128, 4, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,89 +2,38 @@
|
|||||||
// clang-format off
|
// clang-format off
|
||||||
#pragma once
|
#pragma once
|
||||||
|
|
||||||
#include "kernel.h"
|
|
||||||
#include "marlin_template.h"
|
|
||||||
|
|
||||||
namespace MARLIN_NAMESPACE_NAME {
|
namespace MARLIN_NAMESPACE_NAME {
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 1, 8, 8, true, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 1, 8, 8, true, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 1, 8, 4, true, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 1, 8, 8, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 1, 8, 8, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 2, 16, 4, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 1, 8, 4, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 4, 16, 4, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 2, 16, 4, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 2, 8, 4, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 3, 16, 4, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 3, 8, 4, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 4, 16, 4, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 4, 8, 4, false, pipe_stages, false, true, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 1, 8, 8, true, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 1, 8, 4, true, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 1, 8, 8, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 1, 8, 4, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 2, 16, 4, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 2, 8, 4, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4.id(), sglang::kFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 3, 16, 4, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 3, 8, 4, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 4, 16, 4, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 4, 8, 4, false, pipe_stages, false, true, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 1, 8, 8, true, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 1, 8, 4, true, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 1, 8, 8, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 1, 8, 4, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 2, 16, 4, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 2, 8, 4, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 3, 16, 4, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 3, 8, 4, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 4, 16, 4, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 4, 8, 4, false, pipe_stages, false, true, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 1, 8, 8, true, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 1, 8, 4, true, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 1, 8, 8, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 1, 8, 4, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 2, 16, 4, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 2, 8, 4, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 3, 16, 4, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 3, 8, 4, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 256, 4, 16, 4, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4.id(), 128, 4, 8, 4, false, pipe_stages, false, true, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,109 +2,46 @@
|
|||||||
// clang-format off
|
// clang-format off
|
||||||
#pragma once
|
#pragma once
|
||||||
|
|
||||||
#include "kernel.h"
|
|
||||||
#include "marlin_template.h"
|
|
||||||
|
|
||||||
namespace MARLIN_NAMESPACE_NAME {
|
namespace MARLIN_NAMESPACE_NAME {
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 1, 8, 8, true, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 1, 8, 4, true, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 1, 8, 8, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 1, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 2, 16, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 1, 8, 8, true, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 2, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 1, 8, 8, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 3, 16, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 2, 16, 4, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 3, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 4, 16, 4, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 4, 16, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 4, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 1, 8, 8, true, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 1, 8, 4, true, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 1, 8, 8, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 1, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 2, 16, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 2, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 3, 16, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 3, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 4, 16, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 4, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU4B8.id(), sglang::kFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 1, 8, 8, true, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 1, 8, 4, true, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 1, 8, 8, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 1, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 2, 16, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 2, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 3, 16, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 3, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 4, 16, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 4, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 1, 8, 8, true, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 1, 8, 4, true, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 1, 8, 8, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 1, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 2, 16, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 2, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 3, 16, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 3, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 4, 16, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 4, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 1, 8, 8, true, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 1, 8, 4, true, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 1, 8, 8, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 1, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 2, 16, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 2, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 3, 16, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 3, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 256, 4, 16, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU4B8.id(), 128, 4, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,109 +2,46 @@
|
|||||||
// clang-format off
|
// clang-format off
|
||||||
#pragma once
|
#pragma once
|
||||||
|
|
||||||
#include "kernel.h"
|
|
||||||
#include "marlin_template.h"
|
|
||||||
|
|
||||||
namespace MARLIN_NAMESPACE_NAME {
|
namespace MARLIN_NAMESPACE_NAME {
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 1, 8, 8, true, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 1, 8, 4, true, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 1, 8, 8, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 1, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 0, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 2, 16, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 1, 8, 8, true, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 2, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 1, 8, 8, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 3, 16, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 2, 16, 4, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 3, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 4, 16, 4, false, pipe_stages, -1, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 4, 16, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 4, 8, 4, false, pipe_stages, true, false, 0, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 1, 8, 8, true, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 1, 8, 4, true, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 2, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 1, 8, 8, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 1, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 2, 16, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 2, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 4, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 3, 16, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 1, 8, 8, true, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 3, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 1, 8, 8, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 4, 16, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 2, 16, 4, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 4, 8, 4, false, pipe_stages, false, false, -1, false>( MARLIN_KERNEL_PARAMS );
|
template __global__ void Marlin<half, sglang::kU8B128.id(), sglang::kFloat16.id(), 256, 4, 16, 4, false, pipe_stages, 8, false>( MARLIN_KERNEL_PARAMS );
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 1, 8, 8, true, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 1, 8, 4, true, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 1, 8, 8, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 1, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 2, 16, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 2, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 3, 16, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 3, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 4, 16, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 4, 8, 4, false, pipe_stages, false, false, 2, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 1, 8, 8, true, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 1, 8, 4, true, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 1, 8, 8, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 1, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 2, 16, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 2, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 3, 16, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 3, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 4, 16, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 4, 8, 4, false, pipe_stages, false, false, 4, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 1, 8, 8, true, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 1, 8, 4, true, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 1, 8, 8, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 1, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 2, 16, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 2, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 3, 16, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 3, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 256, 4, 16, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
template __global__ void Marlin<half, sglang::kU8B128.id(), 128, 4, 8, 4, false, pipe_stages, false, false, 8, false>( MARLIN_KERNEL_PARAMS );
|
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -25,6 +25,7 @@
|
|||||||
|
|
||||||
#include "kernel.h"
|
#include "kernel.h"
|
||||||
#include "kernel_marlin.cuh"
|
#include "kernel_marlin.cuh"
|
||||||
|
#include "marlin_template.h"
|
||||||
|
|
||||||
#define STATIC_ASSERT_SCALAR_TYPE_VALID(scalar_t) \
|
#define STATIC_ASSERT_SCALAR_TYPE_VALID(scalar_t) \
|
||||||
static_assert( \
|
static_assert( \
|
||||||
@@ -50,13 +51,16 @@ __global__ void permute_cols_kernel(
|
|||||||
int size_m,
|
int size_m,
|
||||||
int size_k,
|
int size_k,
|
||||||
int top_k) {};
|
int top_k) {};
|
||||||
}
|
|
||||||
|
} // namespace marlin
|
||||||
|
|
||||||
torch::Tensor moe_wna16_marlin_gemm(
|
torch::Tensor moe_wna16_marlin_gemm(
|
||||||
torch::Tensor& a,
|
torch::Tensor& a,
|
||||||
std::optional<torch::Tensor> const& c_or_none,
|
std::optional<torch::Tensor> c_or_none,
|
||||||
torch::Tensor& b_q_weight,
|
torch::Tensor& b_q_weight,
|
||||||
|
std::optional<torch::Tensor> const& b_bias_or_none,
|
||||||
torch::Tensor& b_scales,
|
torch::Tensor& b_scales,
|
||||||
|
std::optional<torch::Tensor> const& global_scale_or_none,
|
||||||
std::optional<torch::Tensor> const& b_zeros_or_none,
|
std::optional<torch::Tensor> const& b_zeros_or_none,
|
||||||
std::optional<torch::Tensor> const& g_idx_or_none,
|
std::optional<torch::Tensor> const& g_idx_or_none,
|
||||||
std::optional<torch::Tensor> const& perm_or_none,
|
std::optional<torch::Tensor> const& perm_or_none,
|
||||||
@@ -131,7 +135,7 @@ __global__ void permute_cols_kernel(
|
|||||||
int base_k = 0;
|
int base_k = 0;
|
||||||
|
|
||||||
for (int i = 0; i < iters; i++) {
|
for (int i = 0; i < iters; i++) {
|
||||||
int cur_k = base_k + threadIdx.x;
|
auto cur_k = base_k + threadIdx.x;
|
||||||
int src_pos = perm_int_ptr[cur_k];
|
int src_pos = perm_int_ptr[cur_k];
|
||||||
|
|
||||||
out_half[cur_k] = a_row_half[src_pos];
|
out_half[cur_k] = a_row_half[src_pos];
|
||||||
@@ -141,7 +145,7 @@ __global__ void permute_cols_kernel(
|
|||||||
|
|
||||||
if (rest) {
|
if (rest) {
|
||||||
if (threadIdx.x < rest) {
|
if (threadIdx.x < rest) {
|
||||||
int cur_k = base_k + threadIdx.x;
|
auto cur_k = base_k + threadIdx.x;
|
||||||
int src_pos = perm_int_ptr[cur_k];
|
int src_pos = perm_int_ptr[cur_k];
|
||||||
|
|
||||||
out_half[cur_k] = a_row_half[src_pos];
|
out_half[cur_k] = a_row_half[src_pos];
|
||||||
@@ -215,7 +219,6 @@ int get_scales_cache_size(
|
|||||||
int load_groups = tb_groups * pipe_stages * 2; // Chunk size is 2x pipeline over dim K
|
int load_groups = tb_groups * pipe_stages * 2; // Chunk size is 2x pipeline over dim K
|
||||||
load_groups = max(load_groups, 32); // We load at least 32 scale groups
|
load_groups = max(load_groups, 32); // We load at least 32 scale groups
|
||||||
return load_groups * tb_n * 2;
|
return load_groups * tb_n * 2;
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
int tb_scales = tb_groups * tb_n * 2;
|
int tb_scales = tb_groups * tb_n * 2;
|
||||||
|
|
||||||
@@ -225,6 +228,7 @@ int get_scales_cache_size(
|
|||||||
|
|
||||||
int get_kernel_cache_size(
|
int get_kernel_cache_size(
|
||||||
thread_config_t const& th_config,
|
thread_config_t const& th_config,
|
||||||
|
bool m_block_size_8,
|
||||||
int thread_m_blocks,
|
int thread_m_blocks,
|
||||||
int prob_m,
|
int prob_m,
|
||||||
int prob_n,
|
int prob_n,
|
||||||
@@ -242,11 +246,16 @@ int get_kernel_cache_size(
|
|||||||
int tb_n = th_config.thread_n;
|
int tb_n = th_config.thread_n;
|
||||||
int tb_m = thread_m_blocks * 16;
|
int tb_m = thread_m_blocks * 16;
|
||||||
|
|
||||||
// shm size for block_sorted_ids/block_topk_weights
|
// shm size for block_sorted_ids/rd_block_sorted_ids/block_topk_weights
|
||||||
// both of them requires tb_m * 4 bytes (tb_m * int32 or tb_m * float32)
|
// both of them requires tb_m * 4 bytes (tb_m * int32 or tb_m * float32)
|
||||||
int sh_block_meta_size = tb_m * 4 * 2;
|
int sh_block_meta_size = tb_m * 4;
|
||||||
int sh_a_size = pipe_stages * (tb_m * tb_k) * 2;
|
int sh_a_size = pipe_stages * (tb_m * tb_k) * 2;
|
||||||
int sh_b_size = pipe_stages * (tb_k * tb_n / pack_factor) * 4;
|
int sh_b_size = pipe_stages * (tb_k * tb_n / pack_factor) * 4;
|
||||||
|
int sh_red_size = tb_m * (tb_n + 8) * 2;
|
||||||
|
int sh_bias_size = tb_n * 2;
|
||||||
|
int tmp_size = (sh_b_size > sh_red_size ? sh_red_size : sh_b_size) + sh_bias_size;
|
||||||
|
tmp_size = max(max(sh_b_size, sh_red_size), tmp_size);
|
||||||
|
|
||||||
int sh_s_size =
|
int sh_s_size =
|
||||||
get_scales_cache_size(th_config, prob_m, prob_n, prob_k, num_bits, group_size, has_act_order, is_k_full);
|
get_scales_cache_size(th_config, prob_m, prob_n, prob_k, num_bits, group_size, has_act_order, is_k_full);
|
||||||
int sh_g_idx_size = has_act_order && !is_k_full ? pipe_stages * tb_k / 4 : 0;
|
int sh_g_idx_size = has_act_order && !is_k_full ? pipe_stages * tb_k / 4 : 0;
|
||||||
@@ -260,13 +269,14 @@ int get_kernel_cache_size(
|
|||||||
sh_zp_size = sh_s_size / 2;
|
sh_zp_size = sh_s_size / 2;
|
||||||
}
|
}
|
||||||
|
|
||||||
int total_size = sh_a_size + sh_b_size + sh_s_size + sh_zp_size + sh_g_idx_size + sh_block_meta_size;
|
int total_size = tmp_size + sh_a_size + sh_s_size + sh_zp_size + sh_g_idx_size + sh_block_meta_size;
|
||||||
|
|
||||||
return total_size;
|
return total_size;
|
||||||
}
|
}
|
||||||
|
|
||||||
bool is_valid_config(
|
bool is_valid_config(
|
||||||
thread_config_t const& th_config,
|
thread_config_t const& th_config,
|
||||||
|
bool m_block_size_8,
|
||||||
int thread_m_blocks,
|
int thread_m_blocks,
|
||||||
int prob_m,
|
int prob_m,
|
||||||
int prob_n,
|
int prob_n,
|
||||||
@@ -301,6 +311,7 @@ bool is_valid_config(
|
|||||||
// Check that pipeline fits into cache
|
// Check that pipeline fits into cache
|
||||||
int cache_size = get_kernel_cache_size(
|
int cache_size = get_kernel_cache_size(
|
||||||
th_config,
|
th_config,
|
||||||
|
m_block_size_8,
|
||||||
thread_m_blocks,
|
thread_m_blocks,
|
||||||
prob_m,
|
prob_m,
|
||||||
prob_n,
|
prob_n,
|
||||||
@@ -311,105 +322,151 @@ bool is_valid_config(
|
|||||||
is_k_full,
|
is_k_full,
|
||||||
has_zp,
|
has_zp,
|
||||||
is_zp_float);
|
is_zp_float);
|
||||||
return cache_size <= max_shared_mem;
|
return cache_size + 512 <= max_shared_mem;
|
||||||
}
|
}
|
||||||
|
|
||||||
#define __GET_IF( \
|
#define _GET_IF( \
|
||||||
W_TYPE, \
|
W_TYPE, THREAD_M_BLOCKS, THREAD_N_BLOCKS, THREAD_K_BLOCKS, M_BLOCK_SIZE_8, GROUP_BLOCKS, NUM_THREADS, IS_ZP_FLOAT) \
|
||||||
THREAD_M_BLOCKS, \
|
else if ( \
|
||||||
THREAD_N_BLOCKS, \
|
q_type == W_TYPE && thread_m_blocks == THREAD_M_BLOCKS && thread_n_blocks == THREAD_N_BLOCKS && \
|
||||||
THREAD_K_BLOCKS, \
|
thread_k_blocks == THREAD_K_BLOCKS && m_block_size_8 == M_BLOCK_SIZE_8 && group_blocks == GROUP_BLOCKS && \
|
||||||
M_BLOCK_SIZE_8, \
|
num_threads == NUM_THREADS && is_zp_float == IS_ZP_FLOAT) { \
|
||||||
HAS_ACT_ORDER, \
|
constexpr auto S_TYPE = W_TYPE == sglang::kFE2M1f \
|
||||||
HAS_ZP, \
|
? (GROUP_BLOCKS == 1 ? sglang::kFE4M3fn : sglang::kFE8M0fnu) \
|
||||||
GROUP_BLOCKS, \
|
: (std::is_same<scalar_t, half>::value ? sglang::kFloat16 : sglang::kBFloat16); \
|
||||||
NUM_THREADS, \
|
kernel = Marlin< \
|
||||||
IS_ZP_FLOAT) \
|
scalar_t, \
|
||||||
else if ( \
|
W_TYPE.id(), \
|
||||||
q_type == W_TYPE && thread_m_blocks == THREAD_M_BLOCKS && thread_n_blocks == THREAD_N_BLOCKS && \
|
S_TYPE.id(), \
|
||||||
thread_k_blocks == THREAD_K_BLOCKS && m_block_size_8 == M_BLOCK_SIZE_8 && has_act_order == HAS_ACT_ORDER && \
|
NUM_THREADS, \
|
||||||
has_zp == HAS_ZP && group_blocks == GROUP_BLOCKS && num_threads == NUM_THREADS && is_zp_float == IS_ZP_FLOAT) { \
|
THREAD_M_BLOCKS, \
|
||||||
kernel = Marlin< \
|
THREAD_N_BLOCKS, \
|
||||||
scalar_t, \
|
THREAD_K_BLOCKS, \
|
||||||
W_TYPE.id(), \
|
M_BLOCK_SIZE_8, \
|
||||||
NUM_THREADS, \
|
pipe_stages, \
|
||||||
THREAD_M_BLOCKS, \
|
GROUP_BLOCKS, \
|
||||||
THREAD_N_BLOCKS, \
|
IS_ZP_FLOAT>; \
|
||||||
THREAD_K_BLOCKS, \
|
|
||||||
M_BLOCK_SIZE_8, \
|
|
||||||
pipe_stages, \
|
|
||||||
HAS_ACT_ORDER, \
|
|
||||||
HAS_ZP, \
|
|
||||||
GROUP_BLOCKS, \
|
|
||||||
IS_ZP_FLOAT>; \
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#define GPTQ_GET_IF_M1(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
// COMMON: cases for (group_blocks in [-1, 2, 4, 8] and is_zp_float == false)
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, true, false, 0, NUM_THREADS, false) \
|
// this is the most common cases
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, true, false, 0, NUM_THREADS, false) \
|
// BIGGROUP: cases for big group size (group_blocks in [-1, 8])
|
||||||
\
|
// FZP: cases for float-zero-point (is_zp_float = true)
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, false, false, -1, NUM_THREADS, false) \
|
// ACT: cases for act order case (group_blocks == 0)
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, false, false, 2, NUM_THREADS, false) \
|
// FP4: cases for nvfp4(e2m1) (group_blocks == 1)
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, false, false, 4, NUM_THREADS, false) \
|
#define COMMON_GET_IF_M1(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, false, false, 8, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, -1, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, false, false, -1, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, 2, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, false, false, 2, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, 4, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, false, false, 4, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, 8, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, false, false, 8, NUM_THREADS, false)
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, -1, NUM_THREADS, false) \
|
||||||
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, 2, NUM_THREADS, false) \
|
||||||
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, 4, NUM_THREADS, false) \
|
||||||
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, 8, NUM_THREADS, false)
|
||||||
|
|
||||||
#define GPTQ_GET_IF_M234(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
#define COMMON_GET_IF_M234(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
||||||
__GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, true, false, 0, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, -1, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, true, false, 0, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, 2, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, true, false, 0, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, 4, NUM_THREADS, false) \
|
||||||
\
|
_GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, 8, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, false, false, -1, NUM_THREADS, false) \
|
\
|
||||||
__GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, false, false, 2, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, -1, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, false, false, 4, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, 2, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, false, false, 8, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, 4, NUM_THREADS, false) \
|
||||||
\
|
_GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, 8, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, false, false, -1, NUM_THREADS, false) \
|
\
|
||||||
__GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, false, false, 2, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, -1, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, false, false, 4, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, 2, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, false, false, 8, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, 4, NUM_THREADS, false) \
|
||||||
\
|
_GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, 8, NUM_THREADS, false)
|
||||||
__GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, false, false, -1, NUM_THREADS, false) \
|
|
||||||
__GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, false, false, 2, NUM_THREADS, false) \
|
|
||||||
__GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, false, false, 4, NUM_THREADS, false) \
|
|
||||||
__GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, false, false, 8, NUM_THREADS, false)
|
|
||||||
|
|
||||||
#define AWQ_GET_IF_M1(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
#define COMMON_GET_IF(W_TYPE) \
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, false, true, -1, NUM_THREADS, false) \
|
COMMON_GET_IF_M1(W_TYPE, 8, 8, 256) \
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, false, true, 2, NUM_THREADS, false) \
|
COMMON_GET_IF_M1(W_TYPE, 8, 4, 128) \
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, false, true, 4, NUM_THREADS, false) \
|
COMMON_GET_IF_M234(W_TYPE, 16, 4, 256) \
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, false, true, 8, NUM_THREADS, false) \
|
COMMON_GET_IF_M234(W_TYPE, 8, 4, 128)
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, false, true, -1, NUM_THREADS, false) \
|
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, false, true, 2, NUM_THREADS, false) \
|
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, false, true, 4, NUM_THREADS, false) \
|
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, false, true, 8, NUM_THREADS, false)
|
|
||||||
|
|
||||||
#define AWQ_GET_IF_M234(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
#define BIGGROUP_GET_IF_M1(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
||||||
__GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, false, true, -1, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, -1, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, false, true, 2, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, 8, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, false, true, 4, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, -1, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, false, true, 8, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, 8, NUM_THREADS, false)
|
||||||
\
|
|
||||||
__GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, false, true, -1, NUM_THREADS, false) \
|
#define BIGGROUP_GET_IF_M234(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
||||||
__GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, false, true, 2, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, -1, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, false, true, 4, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, 8, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, false, true, 8, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, -1, NUM_THREADS, false) \
|
||||||
\
|
_GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, 8, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, false, true, -1, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, -1, NUM_THREADS, false) \
|
||||||
__GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, false, true, 2, NUM_THREADS, false) \
|
_GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, 8, NUM_THREADS, false)
|
||||||
__GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, false, true, 4, NUM_THREADS, false) \
|
|
||||||
__GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, false, true, 8, NUM_THREADS, false)
|
#define BIGGROUP_GET_IF(W_TYPE) \
|
||||||
|
BIGGROUP_GET_IF_M1(W_TYPE, 8, 8, 256) \
|
||||||
|
BIGGROUP_GET_IF_M1(W_TYPE, 8, 4, 128) \
|
||||||
|
BIGGROUP_GET_IF_M234(W_TYPE, 16, 4, 256) \
|
||||||
|
BIGGROUP_GET_IF_M234(W_TYPE, 8, 4, 128)
|
||||||
|
|
||||||
|
#define NVFP4_GET_IF_M1(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
||||||
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, 1, NUM_THREADS, false) \
|
||||||
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, 1, NUM_THREADS, false)
|
||||||
|
|
||||||
|
#define NVFP4_GET_IF_M234(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
||||||
|
_GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, 1, NUM_THREADS, false) \
|
||||||
|
_GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, 1, NUM_THREADS, false) \
|
||||||
|
_GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, 1, NUM_THREADS, false)
|
||||||
|
|
||||||
|
#define NVFP4_GET_IF(W_TYPE) \
|
||||||
|
NVFP4_GET_IF_M1(W_TYPE, 8, 8, 256) \
|
||||||
|
NVFP4_GET_IF_M1(W_TYPE, 8, 4, 128) \
|
||||||
|
NVFP4_GET_IF_M234(W_TYPE, 16, 4, 256) \
|
||||||
|
NVFP4_GET_IF_M234(W_TYPE, 8, 4, 128)
|
||||||
|
|
||||||
|
#define MXFP4_GET_IF_M1(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
||||||
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, 2, NUM_THREADS, false) \
|
||||||
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, 2, NUM_THREADS, false)
|
||||||
|
|
||||||
|
#define MXFP4_GET_IF_M234(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
||||||
|
_GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, 2, NUM_THREADS, false) \
|
||||||
|
_GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, 2, NUM_THREADS, false) \
|
||||||
|
_GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, 2, NUM_THREADS, false)
|
||||||
|
|
||||||
|
#define MXFP4_GET_IF(W_TYPE) \
|
||||||
|
MXFP4_GET_IF_M1(W_TYPE, 8, 8, 256) \
|
||||||
|
MXFP4_GET_IF_M1(W_TYPE, 8, 4, 128) \
|
||||||
|
MXFP4_GET_IF_M234(W_TYPE, 16, 4, 256) \
|
||||||
|
MXFP4_GET_IF_M234(W_TYPE, 8, 4, 128)
|
||||||
|
|
||||||
// We currently have 4-bit models only with group_blocks == 4
|
// We currently have 4-bit models only with group_blocks == 4
|
||||||
#define HQQ_GET_IF(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
#define FZP_GET_IF_M1(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, false, true, 4, NUM_THREADS, true) \
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, 4, NUM_THREADS, true) \
|
||||||
__GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, false, true, 4, NUM_THREADS, true) \
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, 4, NUM_THREADS, true)
|
||||||
__GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, false, true, 4, NUM_THREADS, true) \
|
|
||||||
__GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, false, true, 4, NUM_THREADS, true) \
|
#define FZP_GET_IF_M234(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
||||||
__GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, false, true, 4, NUM_THREADS, true)
|
_GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, 4, NUM_THREADS, true) \
|
||||||
|
_GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, 4, NUM_THREADS, true) \
|
||||||
|
_GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, 4, NUM_THREADS, true)
|
||||||
|
|
||||||
|
#define FZP_GET_IF(W_TYPE) \
|
||||||
|
FZP_GET_IF_M1(W_TYPE, 8, 8, 256) \
|
||||||
|
FZP_GET_IF_M1(W_TYPE, 8, 4, 128) \
|
||||||
|
FZP_GET_IF_M234(W_TYPE, 16, 4, 256) \
|
||||||
|
FZP_GET_IF_M234(W_TYPE, 8, 4, 128)
|
||||||
|
|
||||||
|
// We currently have 4-bit models only with group_blocks == 4
|
||||||
|
#define ACT_GET_IF_M1(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
||||||
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, true, 0, NUM_THREADS, false) \
|
||||||
|
_GET_IF(W_TYPE, 1, N_BLOCKS, K_BLOCKS, false, 0, NUM_THREADS, false)
|
||||||
|
|
||||||
|
#define ACT_GET_IF_M234(W_TYPE, N_BLOCKS, K_BLOCKS, NUM_THREADS) \
|
||||||
|
_GET_IF(W_TYPE, 2, N_BLOCKS, K_BLOCKS, false, 0, NUM_THREADS, false) \
|
||||||
|
_GET_IF(W_TYPE, 3, N_BLOCKS, K_BLOCKS, false, 0, NUM_THREADS, false) \
|
||||||
|
_GET_IF(W_TYPE, 4, N_BLOCKS, K_BLOCKS, false, 0, NUM_THREADS, false)
|
||||||
|
|
||||||
|
#define ACT_GET_IF(W_TYPE) \
|
||||||
|
ACT_GET_IF_M1(W_TYPE, 8, 8, 256) \
|
||||||
|
ACT_GET_IF_M1(W_TYPE, 8, 4, 128) \
|
||||||
|
ACT_GET_IF_M234(W_TYPE, 16, 4, 256) \
|
||||||
|
ACT_GET_IF_M234(W_TYPE, 8, 4, 128)
|
||||||
|
|
||||||
template <typename scalar_t>
|
template <typename scalar_t>
|
||||||
MarlinFuncPtr get_marlin_kernel(
|
MarlinFuncPtr get_marlin_kernel(
|
||||||
@@ -427,23 +484,22 @@ MarlinFuncPtr get_marlin_kernel(
|
|||||||
auto kernel = MarlinDefault;
|
auto kernel = MarlinDefault;
|
||||||
if (false) {
|
if (false) {
|
||||||
}
|
}
|
||||||
GPTQ_GET_IF_M1(sglang::kU4B8, 8, 8, 256)
|
|
||||||
GPTQ_GET_IF_M1(sglang::kU4B8, 8, 4, 128)
|
|
||||||
|
|
||||||
GPTQ_GET_IF_M234(sglang::kU4B8, 16, 4, 256)
|
COMMON_GET_IF(sglang::kU4)
|
||||||
GPTQ_GET_IF_M234(sglang::kU4B8, 8, 4, 128)
|
COMMON_GET_IF(sglang::kU4B8)
|
||||||
|
COMMON_GET_IF(sglang::kU8B128)
|
||||||
|
|
||||||
GPTQ_GET_IF_M1(sglang::kU8B128, 8, 8, 256)
|
NVFP4_GET_IF(sglang::kFE2M1f)
|
||||||
GPTQ_GET_IF_M1(sglang::kU8B128, 8, 4, 128)
|
|
||||||
|
|
||||||
GPTQ_GET_IF_M234(sglang::kU8B128, 16, 4, 256)
|
BIGGROUP_GET_IF(sglang::kFE4M3fn)
|
||||||
GPTQ_GET_IF_M234(sglang::kU8B128, 8, 4, 128)
|
|
||||||
|
|
||||||
AWQ_GET_IF_M1(sglang::kU4, 8, 8, 256)
|
ACT_GET_IF(sglang::kU4B8)
|
||||||
AWQ_GET_IF_M1(sglang::kU4, 8, 4, 128)
|
ACT_GET_IF(sglang::kU8B128)
|
||||||
|
if (std::is_same<scalar_t, nv_bfloat16>::value) {
|
||||||
AWQ_GET_IF_M234(sglang::kU4, 16, 4, 256)
|
if (false) {
|
||||||
AWQ_GET_IF_M234(sglang::kU4, 8, 4, 128)
|
}
|
||||||
|
MXFP4_GET_IF(sglang::kFE2M1f)
|
||||||
|
}
|
||||||
|
|
||||||
return kernel;
|
return kernel;
|
||||||
}
|
}
|
||||||
@@ -475,6 +531,7 @@ exec_config_t determine_exec_config(
|
|||||||
|
|
||||||
if (!is_valid_config(
|
if (!is_valid_config(
|
||||||
th_config,
|
th_config,
|
||||||
|
m_block_size_8,
|
||||||
thread_m_blocks,
|
thread_m_blocks,
|
||||||
prob_m,
|
prob_m,
|
||||||
prob_n,
|
prob_n,
|
||||||
@@ -491,6 +548,7 @@ exec_config_t determine_exec_config(
|
|||||||
|
|
||||||
int cache_size = get_kernel_cache_size(
|
int cache_size = get_kernel_cache_size(
|
||||||
th_config,
|
th_config,
|
||||||
|
m_block_size_8,
|
||||||
thread_m_blocks,
|
thread_m_blocks,
|
||||||
prob_m,
|
prob_m,
|
||||||
prob_n,
|
prob_n,
|
||||||
@@ -504,7 +562,7 @@ exec_config_t determine_exec_config(
|
|||||||
|
|
||||||
int group_blocks = 0;
|
int group_blocks = 0;
|
||||||
if (!has_act_order) {
|
if (!has_act_order) {
|
||||||
group_blocks = group_size == -1 ? -1 : group_size / 16;
|
group_blocks = group_size == -1 ? -1 : (group_size / 16);
|
||||||
}
|
}
|
||||||
|
|
||||||
auto kernel = get_marlin_kernel<scalar_t>(
|
auto kernel = get_marlin_kernel<scalar_t>(
|
||||||
@@ -546,7 +604,9 @@ void marlin_mm(
|
|||||||
const void* B,
|
const void* B,
|
||||||
void* C,
|
void* C,
|
||||||
void* C_tmp,
|
void* C_tmp,
|
||||||
|
void* b_bias,
|
||||||
void* s,
|
void* s,
|
||||||
|
void* s2,
|
||||||
void* zp,
|
void* zp,
|
||||||
void* g_idx,
|
void* g_idx,
|
||||||
void* perm,
|
void* perm,
|
||||||
@@ -564,6 +624,7 @@ void marlin_mm(
|
|||||||
int prob_k,
|
int prob_k,
|
||||||
void* workspace,
|
void* workspace,
|
||||||
sglang::ScalarType const& q_type,
|
sglang::ScalarType const& q_type,
|
||||||
|
bool has_bias,
|
||||||
bool has_act_order,
|
bool has_act_order,
|
||||||
bool is_k_full,
|
bool is_k_full,
|
||||||
bool has_zp,
|
bool has_zp,
|
||||||
@@ -587,8 +648,9 @@ void marlin_mm(
|
|||||||
q_type.str());
|
q_type.str());
|
||||||
} else {
|
} else {
|
||||||
TORCH_CHECK(
|
TORCH_CHECK(
|
||||||
q_type == sglang::kU4B8 || q_type == sglang::kU8B128,
|
q_type == sglang::kU4B8 || q_type == sglang::kU8B128 || q_type == sglang::kFE4M3fn || q_type == sglang::kFE2M1f,
|
||||||
"q_type must be uint4b8 or uint8b128 when has_zp = False. Got = ",
|
"q_type must be uint4b8, uint8b128, float8_e4m3fn or float4_e2m1f when "
|
||||||
|
"has_zp = False. Got = ",
|
||||||
q_type.str());
|
q_type.str());
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -620,7 +682,9 @@ void marlin_mm(
|
|||||||
const int4* B_ptr = (const int4*)B;
|
const int4* B_ptr = (const int4*)B;
|
||||||
int4* C_ptr = (int4*)C;
|
int4* C_ptr = (int4*)C;
|
||||||
int4* C_tmp_ptr = (int4*)C_tmp;
|
int4* C_tmp_ptr = (int4*)C_tmp;
|
||||||
|
const int4* bias_ptr = (const int4*)b_bias;
|
||||||
const int4* s_ptr = (const int4*)s;
|
const int4* s_ptr = (const int4*)s;
|
||||||
|
const uint16_t* s2_ptr = (const uint16_t*)s2;
|
||||||
const int4* zp_ptr = (const int4*)zp;
|
const int4* zp_ptr = (const int4*)zp;
|
||||||
const int* g_idx_ptr = (const int*)g_idx;
|
const int* g_idx_ptr = (const int*)g_idx;
|
||||||
const int* perm_ptr = (const int*)perm;
|
const int* perm_ptr = (const int*)perm;
|
||||||
@@ -705,6 +769,7 @@ void marlin_mm(
|
|||||||
TORCH_CHECK(
|
TORCH_CHECK(
|
||||||
is_valid_config(
|
is_valid_config(
|
||||||
thread_tfg,
|
thread_tfg,
|
||||||
|
m_block_size_8,
|
||||||
thread_m_blocks,
|
thread_m_blocks,
|
||||||
prob_m,
|
prob_m,
|
||||||
prob_n,
|
prob_n,
|
||||||
@@ -787,10 +852,10 @@ void marlin_mm(
|
|||||||
// avoid ">>>" being formatted to "> > >"
|
// avoid ">>>" being formatted to "> > >"
|
||||||
// clang-format off
|
// clang-format off
|
||||||
kernel<<<blocks, num_threads, max_shared_mem, stream>>>(
|
kernel<<<blocks, num_threads, max_shared_mem, stream>>>(
|
||||||
A_ptr, B_ptr, C_ptr, C_tmp_ptr, s_ptr, zp_ptr, g_idx_ptr,
|
A_ptr, B_ptr, C_ptr, C_tmp_ptr, bias_ptr, s_ptr, s2_ptr, zp_ptr, g_idx_ptr,
|
||||||
sorted_token_ids_ptr, expert_ids_ptr, num_tokens_past_padded_ptr,
|
sorted_token_ids_ptr, expert_ids_ptr, num_tokens_past_padded_ptr,
|
||||||
topk_weights_ptr, top_k, mul_topk_weights, is_ep, num_groups, prob_m,
|
topk_weights_ptr, top_k, mul_topk_weights, is_ep, num_groups, prob_m,
|
||||||
prob_n, prob_k, locks, use_atomic_add, use_fp32_reduce);
|
prob_n, prob_k, locks, has_bias, use_atomic_add, use_fp32_reduce, max_shared_mem);
|
||||||
// clang-format on
|
// clang-format on
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -800,7 +865,9 @@ torch::Tensor moe_wna16_marlin_gemm(
|
|||||||
torch::Tensor& a,
|
torch::Tensor& a,
|
||||||
std::optional<torch::Tensor> const& c_or_none,
|
std::optional<torch::Tensor> const& c_or_none,
|
||||||
torch::Tensor& b_q_weight,
|
torch::Tensor& b_q_weight,
|
||||||
|
std::optional<torch::Tensor> const& b_bias_or_none,
|
||||||
torch::Tensor& b_scales,
|
torch::Tensor& b_scales,
|
||||||
|
std::optional<torch::Tensor> const& global_scale_or_none,
|
||||||
std::optional<torch::Tensor> const& b_zeros_or_none,
|
std::optional<torch::Tensor> const& b_zeros_or_none,
|
||||||
std::optional<torch::Tensor> const& g_idx_or_none,
|
std::optional<torch::Tensor> const& g_idx_or_none,
|
||||||
std::optional<torch::Tensor> const& perm_or_none,
|
std::optional<torch::Tensor> const& perm_or_none,
|
||||||
@@ -915,7 +982,6 @@ torch::Tensor moe_wna16_marlin_gemm(
|
|||||||
num_groups = b_scales.size(1);
|
num_groups = b_scales.size(1);
|
||||||
|
|
||||||
torch::Tensor g_idx, perm, a_tmp;
|
torch::Tensor g_idx, perm, a_tmp;
|
||||||
;
|
|
||||||
if (g_idx_or_none.has_value() && perm_or_none.has_value()) {
|
if (g_idx_or_none.has_value() && perm_or_none.has_value()) {
|
||||||
g_idx = g_idx_or_none.value();
|
g_idx = g_idx_or_none.value();
|
||||||
perm = perm_or_none.value();
|
perm = perm_or_none.value();
|
||||||
@@ -962,6 +1028,29 @@ torch::Tensor moe_wna16_marlin_gemm(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
torch::Tensor global_scale;
|
||||||
|
if (global_scale_or_none.has_value()) {
|
||||||
|
global_scale = global_scale_or_none.value();
|
||||||
|
TORCH_CHECK(b_q_type == sglang::kFE2M1f && group_size == 16, "global_scale can only be used for nvfp4 format.");
|
||||||
|
} else {
|
||||||
|
global_scale = torch::empty({0}, options);
|
||||||
|
TORCH_CHECK(
|
||||||
|
!(b_q_type == sglang::kFE2M1f && group_size == 16),
|
||||||
|
"the global_scale parameter must be passed for nvfp4 format.");
|
||||||
|
}
|
||||||
|
|
||||||
|
bool has_bias = b_bias_or_none.has_value();
|
||||||
|
torch::Tensor b_bias;
|
||||||
|
if (has_bias) {
|
||||||
|
b_bias = b_bias_or_none.value();
|
||||||
|
TORCH_CHECK(b_bias.device().is_cuda(), "b_bias is not on GPU");
|
||||||
|
TORCH_CHECK(b_bias.is_contiguous(), "b_bias is not contiguous");
|
||||||
|
TORCH_CHECK(b_bias.size(1) == size_n, "b_bias.size(0) != size_n");
|
||||||
|
TORCH_CHECK(b_bias.stride(1) == 1, "b_bias.stride(1) != 1");
|
||||||
|
} else {
|
||||||
|
b_bias = torch::empty({0}, options);
|
||||||
|
}
|
||||||
|
|
||||||
torch::Tensor b_zeros;
|
torch::Tensor b_zeros;
|
||||||
if (b_zeros_or_none.has_value()) {
|
if (b_zeros_or_none.has_value()) {
|
||||||
b_zeros = b_zeros_or_none.value();
|
b_zeros = b_zeros_or_none.value();
|
||||||
@@ -971,13 +1060,18 @@ torch::Tensor moe_wna16_marlin_gemm(
|
|||||||
b_zeros = torch::empty({0}, options);
|
b_zeros = torch::empty({0}, options);
|
||||||
}
|
}
|
||||||
bool has_zp = b_zeros.size(-1) > 0;
|
bool has_zp = b_zeros.size(-1) > 0;
|
||||||
|
|
||||||
if (has_zp) {
|
if (has_zp) {
|
||||||
TORCH_CHECK(b_q_type == sglang::kU4, "b_q_type must be u4 when has_zp = True. Got = ", b_q_type.str());
|
TORCH_CHECK(
|
||||||
|
b_q_type == sglang::kU4 || b_q_type == sglang::kU8,
|
||||||
|
"b_q_type must be u4 or u8 when has_zp = True. Got = ",
|
||||||
|
b_q_type.str());
|
||||||
} else {
|
} else {
|
||||||
TORCH_CHECK(
|
TORCH_CHECK(
|
||||||
b_q_type == sglang::kU4B8 || b_q_type == sglang::kU8B128,
|
b_q_type == sglang::kU4B8 || b_q_type == sglang::kU8B128 || b_q_type == sglang::kFE4M3fn ||
|
||||||
"b_q_type must be uint4b8 or uint8b128 when has_zp = False. Got = ",
|
b_q_type == sglang::kFE2M1f,
|
||||||
|
"b_q_type must be uint4b8, uint8b128, float8_e4m3fn or "
|
||||||
|
"float4_e2m1f when "
|
||||||
|
"has_zp = False. Got = ",
|
||||||
b_q_type.str());
|
b_q_type.str());
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1028,12 +1122,26 @@ torch::Tensor moe_wna16_marlin_gemm(
|
|||||||
|
|
||||||
int dev = a.get_device();
|
int dev = a.get_device();
|
||||||
if (a.scalar_type() == at::ScalarType::Half) {
|
if (a.scalar_type() == at::ScalarType::Half) {
|
||||||
|
void* scales_ptr;
|
||||||
|
if (b_q_type == sglang::kFE2M1f) {
|
||||||
|
if (group_size == 16)
|
||||||
|
scales_ptr = b_scales.data_ptr<at::Float8_e4m3fn>();
|
||||||
|
else if (group_size == 32)
|
||||||
|
scales_ptr = b_scales.data_ptr<at::Float8_e8m0fnu>();
|
||||||
|
else
|
||||||
|
TORCH_CHECK(false, "float4_e2m1f only supports group_size == 16 (NVFP4) ", "and group_size == 32 (MXFP4)");
|
||||||
|
} else {
|
||||||
|
scales_ptr = b_scales.data_ptr<at::Half>();
|
||||||
|
}
|
||||||
|
|
||||||
MARLIN_NAMESPACE_NAME::marlin_mm<half>(
|
MARLIN_NAMESPACE_NAME::marlin_mm<half>(
|
||||||
a.data_ptr<at::Half>(),
|
a.data_ptr<at::Half>(),
|
||||||
b_q_weight.data_ptr(),
|
b_q_weight.data_ptr(),
|
||||||
c.data_ptr<at::Half>(),
|
c.data_ptr<at::Half>(),
|
||||||
c_tmp.data_ptr<float>(),
|
c_tmp.data_ptr<float>(),
|
||||||
b_scales.data_ptr<at::Half>(),
|
b_bias.data_ptr<at::Half>(),
|
||||||
|
scales_ptr,
|
||||||
|
global_scale.data_ptr<at::Half>(),
|
||||||
b_zeros.data_ptr(),
|
b_zeros.data_ptr(),
|
||||||
g_idx.data_ptr(),
|
g_idx.data_ptr(),
|
||||||
perm.data_ptr(),
|
perm.data_ptr(),
|
||||||
@@ -1051,6 +1159,7 @@ torch::Tensor moe_wna16_marlin_gemm(
|
|||||||
size_k,
|
size_k,
|
||||||
workspace.data_ptr(),
|
workspace.data_ptr(),
|
||||||
b_q_type,
|
b_q_type,
|
||||||
|
has_bias,
|
||||||
has_act_order,
|
has_act_order,
|
||||||
is_k_full,
|
is_k_full,
|
||||||
has_zp,
|
has_zp,
|
||||||
@@ -1065,12 +1174,26 @@ torch::Tensor moe_wna16_marlin_gemm(
|
|||||||
use_fp32_reduce,
|
use_fp32_reduce,
|
||||||
is_zp_float);
|
is_zp_float);
|
||||||
} else if (a.scalar_type() == at::ScalarType::BFloat16) {
|
} else if (a.scalar_type() == at::ScalarType::BFloat16) {
|
||||||
|
void* scales_ptr;
|
||||||
|
if (b_q_type == sglang::kFE2M1f) {
|
||||||
|
if (group_size == 16)
|
||||||
|
scales_ptr = b_scales.data_ptr<at::Float8_e4m3fn>();
|
||||||
|
else if (group_size == 32)
|
||||||
|
scales_ptr = b_scales.data_ptr<at::Float8_e8m0fnu>();
|
||||||
|
else
|
||||||
|
TORCH_CHECK(false, "float4_e2m1f only supports group_size == 16 (NVFP4) ", "and group_size == 32 (MXFP4)");
|
||||||
|
} else {
|
||||||
|
scales_ptr = b_scales.data_ptr<at::BFloat16>();
|
||||||
|
}
|
||||||
|
|
||||||
MARLIN_NAMESPACE_NAME::marlin_mm<nv_bfloat16>(
|
MARLIN_NAMESPACE_NAME::marlin_mm<nv_bfloat16>(
|
||||||
a.data_ptr<at::BFloat16>(),
|
a.data_ptr<at::BFloat16>(),
|
||||||
b_q_weight.data_ptr(),
|
b_q_weight.data_ptr(),
|
||||||
c.data_ptr<at::BFloat16>(),
|
c.data_ptr<at::BFloat16>(),
|
||||||
c_tmp.data_ptr<float>(),
|
c_tmp.data_ptr<float>(),
|
||||||
b_scales.data_ptr<at::BFloat16>(),
|
b_bias.data_ptr<at::BFloat16>(),
|
||||||
|
scales_ptr,
|
||||||
|
global_scale.data_ptr<at::BFloat16>(),
|
||||||
b_zeros.data_ptr(),
|
b_zeros.data_ptr(),
|
||||||
g_idx.data_ptr(),
|
g_idx.data_ptr(),
|
||||||
perm.data_ptr(),
|
perm.data_ptr(),
|
||||||
@@ -1088,6 +1211,7 @@ torch::Tensor moe_wna16_marlin_gemm(
|
|||||||
size_k,
|
size_k,
|
||||||
workspace.data_ptr(),
|
workspace.data_ptr(),
|
||||||
b_q_type,
|
b_q_type,
|
||||||
|
has_bias,
|
||||||
has_act_order,
|
has_act_order,
|
||||||
is_k_full,
|
is_k_full,
|
||||||
has_zp,
|
has_zp,
|
||||||
@@ -1109,3 +1233,5 @@ torch::Tensor moe_wna16_marlin_gemm(
|
|||||||
}
|
}
|
||||||
|
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
|
// Registration is done in common_extension.cc for v2 version
|
||||||
|
|||||||
@@ -301,6 +301,7 @@ static inline constexpr auto kU8B128 = ScalarType::uint(8, 128);
|
|||||||
static inline constexpr auto kFE2M1f = ScalarType::float_(2, 1, true, ScalarType::NAN_NONE);
|
static inline constexpr auto kFE2M1f = ScalarType::float_(2, 1, true, ScalarType::NAN_NONE);
|
||||||
static inline constexpr auto kFE3M2f = ScalarType::float_(3, 2, true, ScalarType::NAN_NONE);
|
static inline constexpr auto kFE3M2f = ScalarType::float_(3, 2, true, ScalarType::NAN_NONE);
|
||||||
static inline constexpr auto kFE4M3fn = ScalarType::float_(4, 3, true, ScalarType::NAN_EXTD_RANGE_MAX_MIN);
|
static inline constexpr auto kFE4M3fn = ScalarType::float_(4, 3, true, ScalarType::NAN_EXTD_RANGE_MAX_MIN);
|
||||||
|
static inline constexpr auto kFE8M0fnu = ScalarType(8, 0, false, 0, true, ScalarType::NAN_EXTD_RANGE_MAX_MIN);
|
||||||
static inline constexpr auto kFE5M2 = ScalarType::float_IEEE754(5, 2);
|
static inline constexpr auto kFE5M2 = ScalarType::float_IEEE754(5, 2);
|
||||||
static inline constexpr auto kFE8M7 = ScalarType::float_IEEE754(8, 7);
|
static inline constexpr auto kFE8M7 = ScalarType::float_IEEE754(8, 7);
|
||||||
static inline constexpr auto kFE5M10 = ScalarType::float_IEEE754(5, 10);
|
static inline constexpr auto kFE5M10 = ScalarType::float_IEEE754(5, 10);
|
||||||
|
|||||||
@@ -459,7 +459,9 @@ torch::Tensor moe_wna16_marlin_gemm(
|
|||||||
torch::Tensor& a,
|
torch::Tensor& a,
|
||||||
std::optional<torch::Tensor> const& c_or_none,
|
std::optional<torch::Tensor> const& c_or_none,
|
||||||
torch::Tensor& b_q_weight,
|
torch::Tensor& b_q_weight,
|
||||||
|
std::optional<torch::Tensor> const& b_bias_or_none,
|
||||||
torch::Tensor& b_scales,
|
torch::Tensor& b_scales,
|
||||||
|
std::optional<torch::Tensor> const& global_scale_or_none,
|
||||||
std::optional<torch::Tensor> const& b_zeros_or_none,
|
std::optional<torch::Tensor> const& b_zeros_or_none,
|
||||||
std::optional<torch::Tensor> const& g_idx_or_none,
|
std::optional<torch::Tensor> const& g_idx_or_none,
|
||||||
std::optional<torch::Tensor> const& perm_or_none,
|
std::optional<torch::Tensor> const& perm_or_none,
|
||||||
|
|||||||
@@ -7,7 +7,9 @@ def moe_wna16_marlin_gemm(
|
|||||||
a: torch.Tensor,
|
a: torch.Tensor,
|
||||||
c_or_none: Optional[torch.Tensor],
|
c_or_none: Optional[torch.Tensor],
|
||||||
b_q_weight: torch.Tensor,
|
b_q_weight: torch.Tensor,
|
||||||
|
b_bias_or_none: Optional[torch.Tensor],
|
||||||
b_scales: torch.Tensor,
|
b_scales: torch.Tensor,
|
||||||
|
global_scale_or_none: Optional[torch.Tensor],
|
||||||
b_zeros_or_none: Optional[torch.Tensor],
|
b_zeros_or_none: Optional[torch.Tensor],
|
||||||
g_idx_or_none: Optional[torch.Tensor],
|
g_idx_or_none: Optional[torch.Tensor],
|
||||||
perm_or_none: Optional[torch.Tensor],
|
perm_or_none: Optional[torch.Tensor],
|
||||||
@@ -33,7 +35,9 @@ def moe_wna16_marlin_gemm(
|
|||||||
a,
|
a,
|
||||||
c_or_none,
|
c_or_none,
|
||||||
b_q_weight,
|
b_q_weight,
|
||||||
|
b_bias_or_none,
|
||||||
b_scales,
|
b_scales,
|
||||||
|
global_scale_or_none,
|
||||||
b_zeros_or_none,
|
b_zeros_or_none,
|
||||||
g_idx_or_none,
|
g_idx_or_none,
|
||||||
perm_or_none,
|
perm_or_none,
|
||||||
|
|||||||
Reference in New Issue
Block a user