TP: quantized KV cache support (#23792)
* TP: quantized KV cache support * fix partial view * remove overly strict assert
This commit is contained in:
@@ -3403,10 +3403,6 @@ llama_context * llama_init_from_model(
|
||||
LLAMA_LOG_ERROR("%s: SPLIT_MODE_TENSOR requires flash_attn to be enabled\n", __func__);
|
||||
return nullptr;
|
||||
}
|
||||
if (ggml_is_quantized(params.type_k) || ggml_is_quantized(params.type_v)) {
|
||||
LLAMA_LOG_ERROR("%s: simultaneous use of SPLIT_MODE_TENSOR and KV cache quantization not implemented\n", __func__);
|
||||
return nullptr;
|
||||
}
|
||||
}
|
||||
|
||||
if (params.flash_attn_type != LLAMA_FLASH_ATTN_TYPE_DISABLED && ggml_is_quantized(params.type_k)) {
|
||||
|
||||
+23
-20
@@ -488,7 +488,7 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str
|
||||
return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_MIRRORED);
|
||||
};
|
||||
|
||||
auto get_split_segments = [&](int axis, uint32_t il) -> std::vector<int64_t> {
|
||||
auto get_split_segments = [&](int axis, uint32_t il) -> std::vector<std::pair<int64_t, uint32_t>> {
|
||||
if (ud->model->arch == LLM_ARCH_QWEN3NEXT || ud->model->arch == LLM_ARCH_QWEN35 || ud->model->arch == LLM_ARCH_QWEN35MOE) {
|
||||
const int64_t head_k_dim = hparams.ssm_d_state;
|
||||
const int64_t head_v_dim = hparams.ssm_d_state;
|
||||
@@ -503,26 +503,26 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str
|
||||
if (ud->model->arch == LLM_ARCH_QWEN3NEXT) {
|
||||
if (std::regex_match(tensor_name, pattern_qkv_weight) || std::regex_match(tensor_name, pattern_ssm_conv1d)) {
|
||||
GGML_ASSERT(tensor->ne[axis] == 2*key_dim + value_dim);
|
||||
return {key_dim, key_dim, value_dim};
|
||||
return {{key_dim, 2}, {value_dim, 1}};
|
||||
}
|
||||
} else {
|
||||
const int64_t head_ratio = n_v_heads / n_k_heads;
|
||||
if (std::regex_match(tensor_name, pattern_qkv_weight) || std::regex_match(tensor_name, pattern_ssm_conv1d)) {
|
||||
GGML_ASSERT(tensor->ne[axis] == 2*key_dim + value_dim);
|
||||
return std::vector<int64_t>(2 + head_ratio, key_dim);
|
||||
return {{key_dim, 2 + head_ratio}};
|
||||
}
|
||||
if (std::regex_match(tensor_name, pattern_attn_gate_weight) || std::regex_match(tensor_name, pattern_ssm_out_weight)) {
|
||||
return std::vector<int64_t>(head_ratio, key_dim);
|
||||
return {{key_dim, head_ratio}};
|
||||
}
|
||||
if (std::regex_match(tensor_name, pattern_ssm_dt) || std::regex_match(tensor_name, pattern_ssm_a) ||
|
||||
std::regex_match(tensor_name, pattern_ssm_alpha) || std::regex_match(tensor_name, pattern_ssm_beta)) {
|
||||
return std::vector<int64_t>(head_ratio, n_k_heads);
|
||||
return {{n_k_heads, head_ratio}};
|
||||
}
|
||||
if (std::regex_match(tensor_name, pattern_r_cache)) {
|
||||
return std::vector<int64_t>(2 + head_ratio, key_dim * (hparams.ssm_d_conv - 1));
|
||||
return {{key_dim * (hparams.ssm_d_conv - 1), 2 + head_ratio}};
|
||||
}
|
||||
if (std::regex_match(tensor_name, pattern_s_cache)) {
|
||||
return std::vector<int64_t>(head_ratio, n_k_heads * head_v_dim * head_v_dim);
|
||||
return {{n_k_heads * head_v_dim * head_v_dim, head_ratio}};
|
||||
}
|
||||
}
|
||||
|
||||
@@ -530,9 +530,9 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str
|
||||
if (std::regex_match(tensor_name, pattern_ffn_gate_up_weight)) {
|
||||
const int64_t n_ff_exp = hparams.n_ff_exp;
|
||||
GGML_ASSERT(tensor->ne[axis] == 2*n_ff_exp);
|
||||
return {n_ff_exp, n_ff_exp};
|
||||
return {{n_ff_exp, 2}};
|
||||
}
|
||||
return {tensor->ne[axis]};
|
||||
return {{tensor->ne[axis], 1}};
|
||||
}
|
||||
|
||||
if (std::regex_match(tensor_name, pattern_qkv_weight) || std::regex_match(tensor_name, pattern_qkv_bias)) {
|
||||
@@ -540,17 +540,17 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str
|
||||
const int64_t n_embd_gqa = hparams.n_embd_v_gqa(il);
|
||||
GGML_ASSERT(hparams.n_embd_k_gqa() == n_embd_gqa);
|
||||
GGML_ASSERT(tensor->ne[axis] == n_embd + 2*n_embd_gqa);
|
||||
return {n_embd, n_embd_gqa, n_embd_gqa};
|
||||
return {{n_embd, 1}, {n_embd_gqa, 2}};
|
||||
}
|
||||
if (std::regex_match(tensor_name, pattern_ffn_gate_up_weight)) {
|
||||
const int64_t n_ff_exp = hparams.n_ff_exp;
|
||||
GGML_ASSERT(tensor->ne[axis] == 2*n_ff_exp);
|
||||
return {n_ff_exp, n_ff_exp};
|
||||
return {{n_ff_exp, 2}};
|
||||
}
|
||||
return {tensor->ne[axis]};
|
||||
return {{tensor->ne[axis], 1}};
|
||||
};
|
||||
|
||||
auto get_split_granularity = [&](int64_t blck_size, uint32_t il, const std::vector<int64_t> & segments) -> std::vector<int64_t> {
|
||||
auto get_split_granularity = [&](int64_t blck_size, uint32_t il, const std::vector<std::pair<int64_t, uint32_t>> & segments) -> std::vector<int64_t> {
|
||||
if (hparams.is_recurrent(il)) {
|
||||
// linear attention
|
||||
const int64_t head_dim = hparams.ssm_d_state;
|
||||
@@ -603,16 +603,16 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str
|
||||
return {granularity_kv};
|
||||
}
|
||||
if (std::regex_match(tensor_name, pattern_qkv_weight) || std::regex_match(tensor_name, pattern_qkv_bias)) {
|
||||
GGML_ASSERT(segments.size() == 3);
|
||||
return {granularity_q, granularity_kv, granularity_kv};
|
||||
GGML_ASSERT(segments.size() == 2);
|
||||
return {granularity_q, granularity_kv};
|
||||
}
|
||||
}
|
||||
|
||||
// FFN
|
||||
if (std::regex_match(tensor_name, pattern_ffn_up_gate_weight) || std::regex_match(tensor_name, pattern_ffn_up_gate_bias) ||
|
||||
std::regex_match(tensor_name, pattern_ffn_gate_up_weight) || std::regex_match(tensor_name, pattern_ffn_down_weight)) {
|
||||
GGML_ASSERT(segments.size() <= 2);
|
||||
return std::vector<int64_t>(segments.size(), blck_size);
|
||||
GGML_ASSERT(segments.size() == 1);
|
||||
return {blck_size};
|
||||
}
|
||||
|
||||
// everything else
|
||||
@@ -636,11 +636,12 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str
|
||||
tensor_split_scan[j] += tensor_split_scan[j - 1];
|
||||
}
|
||||
}
|
||||
const std::vector<int64_t> segments = get_split_segments(split_state.axis, tc.il);
|
||||
const std::vector<std::pair<int64_t, uint32_t>> segments = get_split_segments(split_state.axis, tc.il);
|
||||
const std::vector<int64_t> granularity = get_split_granularity(blck_size, tc.il, segments);
|
||||
for (size_t is = 0; is < segments.size(); is++) {
|
||||
const int64_t ne_s = segments[is];
|
||||
const int64_t g_s = granularity[is];
|
||||
const int64_t ne_s = segments[is].first;
|
||||
const uint32_t nr_s = segments[is].second;
|
||||
const int64_t g_s = granularity[is];
|
||||
GGML_ASSERT(ne_full % g_s == 0);
|
||||
int64_t low = 0;
|
||||
size_t j = 0;
|
||||
@@ -654,10 +655,12 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str
|
||||
low = high;
|
||||
}
|
||||
split_state.ne[is*ud->n_devices + (j + tc.rotation) % ud->n_devices] = ne_s - low;
|
||||
split_state.nr[is] = nr_s;
|
||||
}
|
||||
split_state.n_segments = segments.size();
|
||||
} else {
|
||||
memset(split_state.ne, 0, sizeof(split_state.ne));
|
||||
split_state.nr[0] = 1;
|
||||
split_state.n_segments = 1;
|
||||
}
|
||||
return split_state;
|
||||
|
||||
Reference in New Issue
Block a user