[CPU][INT4] Add INT4 kernels for CPU (#8226)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
parent
88f7759402
commit
c35aa0238c
+83
-30
@@ -908,14 +908,10 @@ static inline void check_moe_scales(
|
||||
bool use_fp8_w8a16,
|
||||
const std::optional<at::Tensor>& w1_scale,
|
||||
const std::optional<at::Tensor>& w2_scale,
|
||||
const std::optional<std::vector<int64_t>> block_size,
|
||||
const std::optional<at::Tensor>& a1_scale,
|
||||
const std::optional<at::Tensor>& a2_scale) {
|
||||
const std::optional<std::vector<int64_t>> block_size) {
|
||||
if (use_int8_w8a8) {
|
||||
TORCH_CHECK(w1_scale.has_value(), "missing w1_scale for int8 w8a8.");
|
||||
TORCH_CHECK(w2_scale.has_value(), "missing w2_scale for int8 w8a8.");
|
||||
TORCH_CHECK(!a1_scale.has_value(), "static quantization for activation not supported.");
|
||||
TORCH_CHECK(!a2_scale.has_value(), "static quantization for activation not supported.");
|
||||
}
|
||||
if (use_fp8_w8a16) {
|
||||
TORCH_CHECK(w1_scale.has_value(), "missing w1_scale for fp8 w8a16.");
|
||||
@@ -942,6 +938,7 @@ static inline void check_moe_scales(
|
||||
// topk_weights: [M, topk]
|
||||
// topk_ids: [M, topk] (int32_t)
|
||||
//
|
||||
|
||||
at::Tensor fused_experts_cpu(
|
||||
at::Tensor& hidden_states,
|
||||
at::Tensor& w1,
|
||||
@@ -949,13 +946,12 @@ at::Tensor fused_experts_cpu(
|
||||
at::Tensor& topk_weights,
|
||||
at::Tensor& topk_ids,
|
||||
bool inplace,
|
||||
bool use_int8_w8a8,
|
||||
bool use_fp8_w8a16,
|
||||
int64_t moe_comp_method,
|
||||
const std::optional<at::Tensor>& w1_scale,
|
||||
const std::optional<at::Tensor>& w2_scale,
|
||||
const std::optional<at::Tensor>& w1_zero,
|
||||
const std::optional<at::Tensor>& w2_zero,
|
||||
const std::optional<std::vector<int64_t>> block_size,
|
||||
const std::optional<at::Tensor>& a1_scale,
|
||||
const std::optional<at::Tensor>& a2_scale,
|
||||
bool is_vnni) {
|
||||
RECORD_FUNCTION(
|
||||
"sgl-kernel::fused_experts_cpu", std::vector<c10::IValue>({hidden_states, w1, w2, topk_weights, topk_ids}));
|
||||
@@ -972,8 +968,13 @@ at::Tensor fused_experts_cpu(
|
||||
CHECK_INPUT(w2);
|
||||
CHECK_EQ(topk_weights.sizes(), topk_ids.sizes());
|
||||
CHECK_DIM(2, hidden_states);
|
||||
CHECK_DIM(3, w1);
|
||||
CHECK_DIM(3, w2);
|
||||
if (moe_comp_method == CPUQuantMethod::INT4_W4A8 && is_vnni) {
|
||||
CHECK_DIM(4, w1);
|
||||
CHECK_DIM(4, w2);
|
||||
} else {
|
||||
CHECK_DIM(3, w1);
|
||||
CHECK_DIM(3, w2);
|
||||
}
|
||||
CHECK_DIM(2, topk_weights);
|
||||
CHECK_DIM(2, topk_ids);
|
||||
|
||||
@@ -987,22 +988,29 @@ at::Tensor fused_experts_cpu(
|
||||
|
||||
int64_t M = hidden_states.size(0);
|
||||
int64_t K = hidden_states.size(1);
|
||||
int64_t N = w1.size(1) / 2;
|
||||
int64_t N = moe_comp_method == CPUQuantMethod::INT4_W4A8 ? w1_scale.value().size(1) * w1_scale.value().size(3) / 2
|
||||
: w1.size(1) / 2;
|
||||
int64_t E = w1.size(0);
|
||||
int64_t topk = topk_weights_.size(1);
|
||||
|
||||
// we use int32_t compensation for int8 w8a8
|
||||
int64_t packed_K = get_row_size(K, use_int8_w8a8);
|
||||
int64_t packed_N = get_row_size(N, use_int8_w8a8);
|
||||
int64_t packed_K = get_row_size(K, moe_comp_method == CPUQuantMethod::INT8_W8A8);
|
||||
int64_t packed_N = get_row_size(N, moe_comp_method == CPUQuantMethod::INT8_W8A8);
|
||||
|
||||
// check weight shapes
|
||||
CHECK_EQ(w2.size(0), E);
|
||||
CHECK_EQ(w2.size(1), K);
|
||||
CHECK_EQ(packed_w1.size(2), packed_K);
|
||||
CHECK_EQ(packed_w2.size(2), packed_N);
|
||||
|
||||
if (!(moe_comp_method == CPUQuantMethod::INT4_W4A8)) {
|
||||
CHECK_EQ(w2.size(1), K);
|
||||
CHECK_EQ(packed_w1.size(2), packed_K / (moe_comp_method == CPUQuantMethod::INT4_W4A8 ? 2 : 1));
|
||||
CHECK_EQ(packed_w2.size(2), packed_N / (moe_comp_method == CPUQuantMethod::INT4_W4A8 ? 2 : 1));
|
||||
}
|
||||
// check scales
|
||||
check_moe_scales(use_int8_w8a8, use_fp8_w8a16, w1_scale, w2_scale, block_size, a1_scale, a2_scale);
|
||||
check_moe_scales(
|
||||
moe_comp_method == CPUQuantMethod::INT8_W8A8,
|
||||
moe_comp_method == CPUQuantMethod::FP8_W8A16,
|
||||
w1_scale,
|
||||
w2_scale,
|
||||
block_size);
|
||||
|
||||
at::Tensor out_hidden_states = inplace ? hidden_states : at::empty_like(hidden_states);
|
||||
|
||||
@@ -1058,24 +1066,29 @@ at::Tensor fused_experts_cpu(
|
||||
// 7. intermediate_cache0 : [M * topk, 2N]
|
||||
// 8. B_tmp : [T, MAX_CACHE_BLOCK_SIZE, BLOCK_N, std::max(K, N)]
|
||||
//
|
||||
int64_t buffer_size_nbytes = M * topk * N * 2 + M * topk * K * 2 +
|
||||
num_threads * BLOCK_M * K * (use_int8_w8a8 ? 1 : 2) +
|
||||
num_threads * 2 * BLOCK_M * BLOCK_N * sizeof(float);
|
||||
int64_t buffer_size_nbytes =
|
||||
M * topk * N * 2 + M * topk * K * 2 +
|
||||
num_threads * BLOCK_M * K *
|
||||
(moe_comp_method == CPUQuantMethod::INT8_W8A8 | moe_comp_method == CPUQuantMethod::INT4_W4A8 ? 1 : 2) +
|
||||
num_threads * 2 * BLOCK_M * BLOCK_N * sizeof(float);
|
||||
|
||||
if (use_int8_w8a8) {
|
||||
if (moe_comp_method == CPUQuantMethod::INT8_W8A8) {
|
||||
buffer_size_nbytes += std::max(M * K, M * topk * N) + M * topk * sizeof(float);
|
||||
}
|
||||
if (use_fp8_w8a16) {
|
||||
if (moe_comp_method == CPUQuantMethod::FP8_W8A16) {
|
||||
buffer_size_nbytes += M * topk * 2 * N * 2 + num_threads * MAX_CACHE_BLOCK_SIZE * BLOCK_N * std::max(K, N) * 2;
|
||||
}
|
||||
|
||||
if (moe_comp_method == CPUQuantMethod::INT4_W4A8) {
|
||||
buffer_size_nbytes += M * topk * 2 * N * 2 + std::max(M * K, M * topk * N) + M * topk * sizeof(float) +
|
||||
num_threads * 2 * get_4bit_block_k_size(K / w1_scale.value().size(2)) * BLOCK_N;
|
||||
}
|
||||
auto buffer2 = at::empty({buffer_size_nbytes}, hidden_states.options().dtype(at::kChar));
|
||||
|
||||
AT_DISPATCH_REDUCED_FLOATING_TYPES(st, "fused_experts_kernel_impl", [&] {
|
||||
scalar_t* __restrict__ intermediate_cache1 = (scalar_t*)((void*)(buffer2.data_ptr<int8_t>()));
|
||||
scalar_t* __restrict__ intermediate_cache2 = intermediate_cache1 + M * topk * N;
|
||||
|
||||
if (use_int8_w8a8) {
|
||||
if (moe_comp_method == CPUQuantMethod::INT8_W8A8) {
|
||||
uint8_t* __restrict__ A_tmp = (uint8_t*)((void*)(intermediate_cache2 + M * topk * K));
|
||||
float* __restrict__ C_tmp = (float*)((void*)(A_tmp + num_threads * BLOCK_M * K));
|
||||
uint8_t* __restrict__ Aq_tmp = (uint8_t*)((void*)(C_tmp + num_threads * 2 * BLOCK_M * BLOCK_N));
|
||||
@@ -1109,7 +1122,7 @@ at::Tensor fused_experts_cpu(
|
||||
E,
|
||||
topk,
|
||||
num_tokens_post_pad);
|
||||
} else if (use_fp8_w8a16) {
|
||||
} else if (moe_comp_method == CPUQuantMethod::FP8_W8A16) {
|
||||
// here we just ignore C_tmp as it is not used
|
||||
scalar_t* __restrict__ A_tmp = (scalar_t*)((void*)(intermediate_cache2 + M * topk * K));
|
||||
float* __restrict__ C_tmp = (float*)((void*)(A_tmp + num_threads * BLOCK_M * K));
|
||||
@@ -1142,6 +1155,48 @@ at::Tensor fused_experts_cpu(
|
||||
E,
|
||||
topk,
|
||||
num_tokens_post_pad);
|
||||
} else if (moe_comp_method == CPUQuantMethod::INT4_W4A8) {
|
||||
uint8_t* __restrict__ A_tmp = (uint8_t*)((void*)(intermediate_cache2 + M * topk * K));
|
||||
float* __restrict__ C_tmp = (float*)((void*)(A_tmp + num_threads * BLOCK_M * K));
|
||||
scalar_t* __restrict__ intermediate_cache0 = (scalar_t*)((void*)(C_tmp + num_threads * 2 * BLOCK_M * BLOCK_N));
|
||||
uint8_t* __restrict__ Aq_tmp = (uint8_t*)((void*)(intermediate_cache0 + M * topk * 2 * N));
|
||||
float* __restrict__ As_tmp = (float*)((void*)(Aq_tmp + std::max(M * K, M * topk * N)));
|
||||
int8_t* __restrict__ dqB_tmp = (int8_t*)((void*)(As_tmp + M * topk));
|
||||
|
||||
// weight + compensation shape = [Nc, Kc, block_n * block_k / 2 + block_n*sizeof(int32_t)]
|
||||
// scales/qzeros shape = [E, Nc, G, block_n]
|
||||
int64_t num_groups = w1_scale.value().size(2);
|
||||
const int group_size = K / num_groups;
|
||||
// TODO: check scales and zeros
|
||||
fused_experts_int4_w4a8_kernel_impl<scalar_t>(
|
||||
out_hidden_states.data_ptr<scalar_t>(),
|
||||
intermediate_cache0,
|
||||
intermediate_cache1,
|
||||
intermediate_cache2,
|
||||
A_tmp,
|
||||
Aq_tmp,
|
||||
As_tmp,
|
||||
nullptr,
|
||||
C_tmp,
|
||||
dqB_tmp,
|
||||
hidden_states.data_ptr<scalar_t>(),
|
||||
packed_w1.data_ptr<uint8_t>(),
|
||||
packed_w2.data_ptr<uint8_t>(),
|
||||
w1_zero.value().data_ptr<int8_t>(),
|
||||
w2_zero.value().data_ptr<int8_t>(),
|
||||
w1_scale.value().data_ptr<float>(),
|
||||
w2_scale.value().data_ptr<float>(),
|
||||
group_size,
|
||||
topk_weights.data_ptr<float>(),
|
||||
sorted_ids,
|
||||
expert_ids,
|
||||
offsets,
|
||||
M,
|
||||
N,
|
||||
K,
|
||||
E,
|
||||
topk,
|
||||
num_tokens_post_pad);
|
||||
} else {
|
||||
scalar_t* __restrict__ A_tmp = intermediate_cache2 + M * topk * K;
|
||||
float* __restrict__ C_tmp = (float*)((void*)(A_tmp + num_threads * BLOCK_M * K));
|
||||
@@ -1188,8 +1243,6 @@ at::Tensor shared_expert_cpu(
|
||||
const std::optional<at::Tensor>& w1_scale,
|
||||
const std::optional<at::Tensor>& w2_scale,
|
||||
const std::optional<std::vector<int64_t>> block_size,
|
||||
const std::optional<at::Tensor>& a1_scale,
|
||||
const std::optional<at::Tensor>& a2_scale,
|
||||
bool is_vnni) {
|
||||
RECORD_FUNCTION("sgl-kernel::shared_expert_cpu", std::vector<c10::IValue>({hidden_states, w1, w2}));
|
||||
|
||||
@@ -1224,7 +1277,7 @@ at::Tensor shared_expert_cpu(
|
||||
CHECK_EQ(packed_w2.size(1), packed_N);
|
||||
|
||||
// check scales
|
||||
check_moe_scales(use_int8_w8a8, use_fp8_w8a16, w1_scale, w2_scale, block_size, a1_scale, a2_scale);
|
||||
check_moe_scales(use_int8_w8a8, use_fp8_w8a16, w1_scale, w2_scale, block_size);
|
||||
|
||||
at::Tensor out_hidden_states = inplace ? hidden_states : at::empty_like(hidden_states);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user