update pre-commit config (#18860)
This commit is contained in:
@@ -14,7 +14,9 @@ from sglang.srt.layers.quantization.fp8_kernel import (
|
||||
from sglang.srt.layers.quantization.fp8_kernel import (
|
||||
per_token_group_quant_8bit as triton_per_token_group_quant_8bit,
|
||||
)
|
||||
from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_8bit
|
||||
from sglang.srt.layers.quantization.fp8_kernel import (
|
||||
sglang_per_token_group_quant_8bit,
|
||||
)
|
||||
from sglang.srt.utils import is_hip
|
||||
from sglang.srt.utils.bench_utils import bench_kineto
|
||||
|
||||
|
||||
@@ -176,8 +176,10 @@ void merge_attn_states_launcher(
|
||||
LAUNCH_MERGE_ATTN_STATES(scalar_t, NUM_THREADS);
|
||||
}
|
||||
|
||||
#define CALL_MERGE_ATTN_STATES_LAUNCHER(scalar_t) \
|
||||
{ merge_attn_states_launcher<scalar_t>(v_a, s_a, v_b, s_b, v_merged, s_merged); }
|
||||
#define CALL_MERGE_ATTN_STATES_LAUNCHER(scalar_t) \
|
||||
{ \
|
||||
merge_attn_states_launcher<scalar_t>(v_a, s_a, v_b, s_b, v_merged, s_merged); \
|
||||
}
|
||||
|
||||
void merge_state_v2(
|
||||
at::Tensor v_a, at::Tensor s_a, at::Tensor v_b, at::Tensor s_b, at::Tensor v_merged, at::Tensor s_merged) {
|
||||
|
||||
@@ -65,9 +65,9 @@ TORCH_LIBRARY_EXPAND(sgl_kernel, m) {
|
||||
m.impl("all_reduce_unreg", torch::kCUDA, &all_reduce_unreg);
|
||||
|
||||
// Deterministic all-reduce for ROCm
|
||||
extern void deterministic_all_reduce_reg(int64_t _fa, torch::Tensor & inp, torch::Tensor & out);
|
||||
extern void deterministic_all_reduce_reg(int64_t _fa, torch::Tensor& inp, torch::Tensor& out);
|
||||
extern void deterministic_all_reduce_unreg(
|
||||
int64_t _fa, torch::Tensor & inp, torch::Tensor & reg_buffer, torch::Tensor & out);
|
||||
int64_t _fa, torch::Tensor& inp, torch::Tensor& reg_buffer, torch::Tensor& out);
|
||||
|
||||
m.def("deterministic_all_reduce_reg(int fa, Tensor inp, Tensor! out) -> ()");
|
||||
m.impl("deterministic_all_reduce_reg", torch::kCUDA, &deterministic_all_reduce_reg);
|
||||
|
||||
@@ -452,7 +452,7 @@ void sgl_per_token_group_quant_8bit_v2(
|
||||
#define LAUNCH_KERNEL(GROUP_SIZE, T, DST_DTYPE) \
|
||||
do { \
|
||||
constexpr int THREADS_PER_SUBWARP = GROUP_SIZE / 16; \
|
||||
TORCH_CHECK(THREADS_PER_SUBWARP* INPUT_PRIMARY_VEC_NUM_BYTES == group_size * sizeof(T)); \
|
||||
TORCH_CHECK(THREADS_PER_SUBWARP * INPUT_PRIMARY_VEC_NUM_BYTES == group_size * sizeof(T)); \
|
||||
\
|
||||
using dst_dtype_info = DtypeInfo<DST_DTYPE>; \
|
||||
CHECK_EQ(dst_dtype_info::MIN, min_8bit); \
|
||||
|
||||
@@ -454,7 +454,7 @@ def scaled_fp4_experts_quant(
|
||||
input_tensor.ndim == 2
|
||||
), f"input.ndim needs to be == 2, but got {input_tensor.ndim}."
|
||||
if expert_map is not None:
|
||||
(m, k) = input_tensor.shape
|
||||
m, k = input_tensor.shape
|
||||
output_tensor_shape = (m * topk, k)
|
||||
input_tensor = shuffle_rows(input_tensor, expert_map, output_tensor_shape)
|
||||
m_numtopk, k = input_tensor.shape
|
||||
|
||||
@@ -13,7 +13,9 @@ from sgl_kernel.test_utils import (
|
||||
from sglang.srt.layers.quantization.fp8_kernel import (
|
||||
per_token_group_quant_8bit as triton_per_token_group_quant_8bit,
|
||||
)
|
||||
from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_8bit
|
||||
from sglang.srt.layers.quantization.fp8_kernel import (
|
||||
sglang_per_token_group_quant_8bit,
|
||||
)
|
||||
from sglang.srt.utils import get_bool_env_var, is_hip
|
||||
|
||||
_is_hip = is_hip()
|
||||
|
||||
Reference in New Issue
Block a user