diff --git a/common/common.h b/common/common.h index 1eaf6644..28999cd4 100644 --- a/common/common.h +++ b/common/common.h @@ -421,7 +421,7 @@ struct gpt_params { bool rope_cache = false; // if to use RoPE cache (for supported models) bool graph_reuse = true; // if to reuse compute graphs bool dsa = false; // enable GLM DSA sparse attention (off by default; opt-in via --dsa) - bool fused_idx_topk = false; // enable the fused indexer topk op (off by default; opt-in via -fidx pr --fused-indexer-topk) + bool fused_idx_topk = true; // enable the fused indexer topk op (off by default; opt-in via -fidx pr --fused-indexer-topk) int dsa_top_k = -1; // DSA top-k override (<0 => use the model's configured indexer_top_k) int min_experts = -1; float thresh_experts = 0; diff --git a/convert_hf_to_gguf.py b/convert_hf_to_gguf.py index f3540dea..b0e1d03f 100644 --- a/convert_hf_to_gguf.py +++ b/convert_hf_to_gguf.py @@ -4591,6 +4591,29 @@ class DeepseekV2Model(Model): raise ValueError(f"Unprocessed experts: {experts}") +@Model.register("DeepseekV4ForCausalLM") +@Model.register("DeepseekV4FlashForCausalLM") +@Model.register("DeepseekV4ProForCausalLM") +class DeepseekV4Model(DeepseekV2Model): + model_arch = gguf.MODEL_ARCH.DEEPSEEK4 + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.block_count = self.hparams["num_hidden_layers"] + self.hparams.get("num_nextn_predict_layers", 0) + self.tensor_map = gguf.get_tensor_name_map(self.model_arch, self.block_count) + + def set_gguf_parameters(self): + super().set_gguf_parameters() + + if (indexer_heads := self.hparams.get("num_indexer_heads")) is not None: + self.gguf_writer.add_attention_indexer_head_count(indexer_heads) + if (indexer_dim := self.hparams.get("indexer_head_dim")) is not None: + self.gguf_writer.add_attention_indexer_key_length(indexer_dim) + if (indexer_top_k := self.hparams.get("indexer_topk")) is not None: + self.gguf_writer.add_attention_indexer_top_k(indexer_top_k) + if (nextn_layers := self.hparams.get("num_nextn_predict_layers")) is not None: + self.gguf_writer.add_nextn_predict_layers(nextn_layers) + @Model.register("OpenPanguV2ForCausalLM") class OpenPanguV2Model(DeepseekV2Model): # openPangu-2.0-Flash: MLA + DSA/SWA hybrid + MoE + mHC(Hyper-Connections) + MoME convs. @@ -4729,7 +4752,7 @@ class OpenPanguV2Model(DeepseekV2Model): (self.map_tensor_name(name_vb), v_b), ] - # Everything else (attn/norms/mHC/MoME conv/param-sink/indexer/nextn) maps by name. + # Everything else maps by name. return [(self.map_tensor_name(name), data_torch)] diff --git a/examples/server/server-context.cpp b/examples/server/server-context.cpp index 5ed5a395..e86e97eb 100644 --- a/examples/server/server-context.cpp +++ b/examples/server/server-context.cpp @@ -281,6 +281,10 @@ bool server_context::load_model(const gpt_params& params_) { void server_context::init() { const int32_t n_ctx_slot = n_ctx / params_base.n_parallel; + if (!system_prompt.empty() && std::strcmp(llama_model_arch_string(model), "deepseek4") == 0) { + throw std::runtime_error("DeepSeek4 server system prompts are unsupported because seq_cp does not copy private cache state"); + } + LOG_INFO("initializing slots", { {"n_slots", params_base.n_parallel} }); if (params_base.has_mtp) { @@ -383,7 +387,7 @@ void server_context::init() { metrics.init(); - if (params_base.cache_ram_mib != 0) { + if (params_base.cache_ram_mib != 0 && llama_model_supports_partial_kv_reuse(model)) { if (params_base.cache_ram_mib < 0) { LLAMA_LOG_INFO("prompt cache is enabled, size limit: %s\n", "no limit"); } @@ -395,7 +399,11 @@ void server_context::init() { prompt_cache = std::make_unique(ctx, params_base.cache_ram_mib, 0); } else { - LLAMA_LOG_INFO("%s", "prompt cache is disabled - use `--cache-ram N` to enable it\n"); + if (params_base.cache_ram_mib != 0) { + LLAMA_LOG_WARN("prompt cache is disabled because this model has private state outside the generic KV cache\n"); + } else { + LLAMA_LOG_INFO("%s", "prompt cache is disabled - use `--cache-ram N` to enable it\n"); + } } // populate chat template params @@ -2073,6 +2081,11 @@ void server_context::system_prompt_update() { } bool server_context::system_prompt_set(const std::string& sys_prompt) { + if (!sys_prompt.empty() && model != nullptr && std::strcmp(llama_model_arch_string(model), "deepseek4") == 0) { + LOG_ERROR("DeepSeek4 server system prompts are unsupported because seq_cp does not copy private cache state", {}); + return false; + } + system_prompt = sys_prompt; LOG_VERBOSE("system prompt process", { @@ -2799,7 +2812,10 @@ void server_context::process_single_task(server_task&& task) { if (task.data.contains("system_prompt")) { std::string sys_prompt = json_value(task.data, "system_prompt", std::string()); - system_prompt_set(sys_prompt); + if (!system_prompt_set(sys_prompt)) { + send_error(task, "DeepSeek4 server system prompts are unsupported", ERROR_TYPE_INVALID_REQUEST); + break; + } for (server_slot& slot : slots) { slot.n_past = 0; diff --git a/examples/sweep-bench/sweep-bench.cpp b/examples/sweep-bench/sweep-bench.cpp index 741f2230..fe44fb6d 100644 --- a/examples/sweep-bench/sweep-bench.cpp +++ b/examples/sweep-bench/sweep-bench.cpp @@ -191,6 +191,8 @@ int main(int argc, char ** argv) { // first measure token generation performance at this context size const auto t_tg_start = ggml_time_us(); + //printf("======================================== tg_start for n_kv = %u\n", n_kv); + //fprintf(stderr, "======================================== tg_start for n_kv = %u\n", n_kv); for (int irep = 0; irep < nrep; ++irep) { @@ -207,6 +209,8 @@ int main(int argc, char ** argv) { } } + //printf("======================================== tg_end for n_kv = %u\n", n_kv); + //fprintf(stderr, "======================================== tg_end for n_kv = %u\n", n_kv); const auto t_tg_end = ggml_time_us(); diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index 38be6e64..87ca109d 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -707,6 +707,9 @@ extern "C" { GGML_OP_INDEXER_TOPK, GGML_OP_MASK_TOPK, GGML_OP_SINKHORN, + GGML_OP_HC_PRE, + GGML_OP_HC_POST, + GGML_OP_MASK_TO_IDX, GGML_OP_COUNT, }; @@ -730,6 +733,7 @@ extern "C" { GGML_UNARY_OP_GELU, GGML_UNARY_OP_EXP, GGML_UNARY_OP_SOFTPLUS, + GGML_UNARY_OP_SQRT_SOFTPLUS, GGML_UNARY_OP_COUNT, }; @@ -1226,6 +1230,14 @@ extern "C" { struct ggml_context * ctx, struct ggml_tensor * a); + GGML_API struct ggml_tensor * ggml_sqrt_softplus( + struct ggml_context * ctx, + struct ggml_tensor * a); + + GGML_API struct ggml_tensor * ggml_sqrt_softplus_inplace( + struct ggml_context * ctx, + struct ggml_tensor * a); + // return scalar GGML_API struct ggml_tensor * ggml_sum( struct ggml_context * ctx, @@ -2476,6 +2488,10 @@ extern "C" { float max_bias, float softcap); + // Backend hint stored in ggml_flash_attn_ext op_params slot 4. + // Negative values request the generic implementation instead of IQK FA. + #define GGML_FLASH_ATTN_EXT_IQK_DISABLED (-1) + GGML_API void ggml_flash_attn_ext_set_prec( struct ggml_tensor * a, enum ggml_prec prec); @@ -2605,6 +2621,28 @@ extern "C" { float eps, bool output_transposed); + GGML_API struct ggml_tensor * ggml_hc_pre( + struct ggml_context * ctx, + struct ggml_tensor * x, + struct ggml_tensor * scale, + struct ggml_tensor * bias, + int S, + int n_iters, + float eps); + + GGML_API struct ggml_tensor * ggml_hc_post( + struct ggml_context * ctx, + struct ggml_tensor * x, + struct ggml_tensor * post, + struct ggml_tensor * res, + struct ggml_tensor * comb); + + GGML_API struct ggml_tensor * ggml_mask_to_index( + struct ggml_context * ctx, + struct ggml_tensor * mask, + int max_row_size); + + // custom operators typedef void (*ggml_unary_op_f32_t) (const int, float *, const float *); diff --git a/ggml/src/ggml-cuda.cu b/ggml/src/ggml-cuda.cu index 7e66572e..f5d60400 100644 --- a/ggml/src/ggml-cuda.cu +++ b/ggml/src/ggml-cuda.cu @@ -3836,6 +3836,9 @@ static bool ggml_cuda_compute_forward(ggml_backend_cuda_context & ctx, struct gg case GGML_UNARY_OP_SOFTPLUS: ggml_cuda_op_softplus(ctx, dst); break; + case GGML_UNARY_OP_SQRT_SOFTPLUS: + ggml_cuda_op_sqrt_softplus(ctx, dst); + break; default: return -1; } @@ -4131,6 +4134,12 @@ static bool ggml_cuda_compute_forward(ggml_backend_cuda_context & ctx, struct gg case GGML_OP_SINKHORN: ggml_cuda_op_sinkhorn(ctx, dst); break; + case GGML_OP_HC_PRE: + ggml_cuda_op_hc_pre(ctx, dst); + break; + case GGML_OP_HC_POST: + ggml_cuda_op_hc_post(ctx, dst); + break; case GGML_OP_FLASH_ATTN_EXT: ggml_cuda_flash_attn_ext(ctx, dst); break; @@ -4140,6 +4149,9 @@ static bool ggml_cuda_compute_forward(ggml_backend_cuda_context & ctx, struct gg case GGML_OP_MASK_TOPK: ggml_cuda_op_indexer_mask(ctx, dst); break; + case GGML_OP_MASK_TO_IDX: + ggml_cuda_op_mask_to_index(ctx, dst); + break; default: return false; } @@ -4718,6 +4730,7 @@ GGML_CALL static bool ggml_backend_cuda_supports_op(ggml_backend_t backend, cons case GGML_UNARY_OP_TANH: case GGML_UNARY_OP_EXP: case GGML_UNARY_OP_SOFTPLUS: + case GGML_UNARY_OP_SQRT_SOFTPLUS: case GGML_UNARY_OP_NEG: return ggml_is_contiguous(op->src[0]); default: @@ -4829,6 +4842,8 @@ GGML_CALL static bool ggml_backend_cuda_supports_op(ggml_backend_t backend, cons case GGML_TYPE_Q5_1: case GGML_TYPE_Q8_0: return true; + case GGML_TYPE_I32: + return op->src[0]->type == op->type; default: return false; } @@ -5045,6 +5060,10 @@ GGML_CALL static bool ggml_backend_cuda_supports_op(ggml_backend_t backend, cons case GGML_OP_DELTA_NET: case GGML_OP_INDEXER_TOPK: case GGML_OP_MASK_TOPK: + case GGML_OP_MASK_TO_IDX: + return true; + case GGML_OP_HC_PRE: + case GGML_OP_HC_POST: return true; case GGML_OP_SINKHORN: { const int sink_s = op->op_params[0]; diff --git a/ggml/src/ggml-cuda/dsa_attn.cu b/ggml/src/ggml-cuda/dsa_attn.cu index ba463960..8dc9c04d 100644 --- a/ggml/src/ggml-cuda/dsa_attn.cu +++ b/ggml/src/ggml-cuda/dsa_attn.cu @@ -14,7 +14,8 @@ static __global__ void k_prepare_mask(int nidx, const int * __restrict__ idx, co int row = blockIdx.x; int col = blockIdx.y*blockDim.x + threadIdx.x; idx += row*stride_idx; - m_out[row*nidx + col] = m_in[row*stride_m + idx[col]]; + int ii = idx[col]; + m_out[row*nidx + col] = ii >= 0 ? m_in[row*stride_m + ii] : __float2half(-INFINITY); } static __global__ void k_prepare_one_batch_kv(int nk, int ncol, const int * idx, const char * k_in, @@ -22,6 +23,10 @@ static __global__ void k_prepare_one_batch_kv(int nk, int ncol, const int * idx, int row = blockIdx.y; int col = blockIdx.x; int i = idx[row*stride_idx + col]; + if (i < 0) { + i = 0; + //return; + } auto k_row = (const half *)(k_in + stride_k * i); k_out += (row*ncol + col)*nk; for (int j = threadIdx.x; j < nk; j += blockDim.x) { @@ -49,12 +54,12 @@ static __global__ void k_copy_dst(int nelem, const half * kqv16, float * dst) { } template -static __global__ void soft_max_f16_simple(half * x, const half * mask, const int ncols_par, const int nrows_y, const float scale) { +static __global__ void soft_max_f16_simple(half * x, const half * mask, const float * sinks, const int ncols_par, const int nrows_y, const float scale) { const int ncols = ncols_template == 0 ? ncols_par : ncols_template; const int tid = threadIdx.x; const int rowx = blockIdx.x; - const int rowy = rowx / nrows_y; // broadcast the mask in the row dimension + const int rowy = rowx / nrows_y; const int block_size = block_size_template == 0 ? blockDim.x : block_size_template; @@ -66,7 +71,7 @@ static __global__ void soft_max_f16_simple(half * x, const half * mask, const in // shared memory buffer to cache values between iterations: float * vals = buf_iw + WARP_SIZE; - float max_val = -INFINITY; + float max_val = sinks ? sinks[rowx % nrows_y] : -INFINITY; #pragma unroll for (int col0 = 0; col0 < ncols; col0 += block_size) { @@ -135,6 +140,10 @@ static __global__ void soft_max_f16_simple(half * x, const half * mask, const in tmp = warp_reduce_sum(tmp); } + if (sinks) { + tmp += expf(sinks[rowx % nrows_y] - max_val); + } + const float inv_sum = 1.0f / tmp; #pragma unroll @@ -152,7 +161,9 @@ static __global__ void soft_max_f16_simple(half * x, const half * mask, const in #define CUDA_SOFT_MAX_BLOCK_SIZE 1024 -static void soft_max_f16_cuda_simple(half * x, const half * mask, const int ncols_x, const int nrows_x, +// nrows_y is Q->ne[2] +// nrows_x is Q->ne[2] * nrows +static void soft_max_f16_cuda_simple(half * x, const half * mask, const float * sinks, const int ncols_x, const int nrows_x, const int nrows_y, const float scale, cudaStream_t stream) { int nth = WARP_SIZE; while (nth < ncols_x && nth < CUDA_SOFT_MAX_BLOCK_SIZE) nth *= 2; @@ -165,31 +176,31 @@ static void soft_max_f16_cuda_simple(half * x, const half * mask, const int ncol switch (ncols_x) { case 32: - soft_max_f16_simple<32, 32><<>>(x, mask, ncols_x, nrows_y, scale); + soft_max_f16_simple<32, 32><<>>(x, mask, sinks, ncols_x, nrows_y, scale); break; case 64: - soft_max_f16_simple<64, 64><<>>(x, mask, ncols_x, nrows_y, scale); + soft_max_f16_simple<64, 64><<>>(x, mask, sinks, ncols_x, nrows_y, scale); break; case 128: - soft_max_f16_simple<128, 128><<>>(x, mask, ncols_x, nrows_y, scale); + soft_max_f16_simple<128, 128><<>>(x, mask, sinks, ncols_x, nrows_y, scale); break; case 256: - soft_max_f16_simple<256, 256><<>>(x, mask, ncols_x, nrows_y, scale); + soft_max_f16_simple<256, 256><<>>(x, mask, sinks, ncols_x, nrows_y, scale); break; case 512: - soft_max_f16_simple<512, 512><<>>(x, mask, ncols_x, nrows_y, scale); + soft_max_f16_simple<512, 512><<>>(x, mask, sinks, ncols_x, nrows_y, scale); break; case 1024: - soft_max_f16_simple<1024, 1024><<>>(x, mask, ncols_x, nrows_y, scale); + soft_max_f16_simple<1024, 1024><<>>(x, mask, sinks, ncols_x, nrows_y, scale); break; case 2048: - soft_max_f16_simple<2048, 1024><<>>(x, mask, ncols_x, nrows_y, scale); + soft_max_f16_simple<2048, 1024><<>>(x, mask, sinks, ncols_x, nrows_y, scale); break; case 4096: - soft_max_f16_simple<4096, 1024><<>>(x, mask, ncols_x, nrows_y, scale); + soft_max_f16_simple<4096, 1024><<>>(x, mask, sinks, ncols_x, nrows_y, scale); break; default: - soft_max_f16_simple<0, 0><<>>(x, mask, ncols_x, nrows_y, scale); + soft_max_f16_simple<0, 0><<>>(x, mask, sinks, ncols_x, nrows_y, scale); break; } } @@ -206,16 +217,17 @@ bool ggml_cuda_dsa_attn_ext(ggml_backend_cuda_context & ctx, ggml_tensor * dst) const ggml_tensor * sink = dst->src[4]; const ggml_tensor * indexer = dst->src[5]; - if (sink) return false; // We do not support sinks at this point if (!Q || !K || !V || !mask || !indexer) return false; - if (indexer->ne[0] % 256 != 0) return false; // lazyness to add checks and handle tailes in case of not multiple of 256 + if (indexer->ne[0] % 256 != 0) return false; // lazyness to add checks and handle tails in case of not multiple of 256 // But are there DSA variants where top_k is not a multiple of 256? if (K->ne[1] < 4*indexer->ne[0]) return false; // for efficiency if (K->ne[2] > 1 || K->ne[3] > 1 || mask->ne[2] > 1 || mask->ne[3] > 1 || Q->ne[3] > 1) return false; if (K->type != GGML_TYPE_F16 || V->type != GGML_TYPE_F16 || mask->type != GGML_TYPE_F16 || Q->type != GGML_TYPE_F32) return false; if (K->ne[0] != Q->ne[0]) return false; + //printf("%s(%s)\n", __func__, dst->name); + float scale; memcpy(&scale, dst->op_params, sizeof(float)); @@ -280,7 +292,9 @@ bool ggml_cuda_dsa_attn_ext(ggml_backend_cuda_context & ctx, ggml_tensor * dst) q16.get(), Q->ne[0], Q->ne[0]*Q->ne[2], &beta, kq16.get(), indexer->ne[0], indexer->ne[0]*Q->ne[2], nrows)); - soft_max_f16_cuda_simple(kq16.get(), mask16.get() + first*indexer->ne[0], indexer->ne[0], Q->ne[2]*nrows, + soft_max_f16_cuda_simple(kq16.get(), mask16.get() + first*indexer->ne[0], + sink ? (const float *)sink->data : nullptr, + indexer->ne[0], Q->ne[2]*nrows, Q->ne[2], scale, ctx.stream()); CUDA_CHECK(cudaGetLastError()); diff --git a/ggml/src/ggml-cuda/getrows.cu b/ggml/src/ggml-cuda/getrows.cu index f734271a..fd67dce1 100644 --- a/ggml/src/ggml-cuda/getrows.cu +++ b/ggml/src/ggml-cuda/getrows.cu @@ -149,7 +149,7 @@ void ggml_cuda_op_get_rows(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { GGML_ASSERT(src1->type == GGML_TYPE_I32); - GGML_ASSERT(dst->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32 || (src0->type == GGML_TYPE_I32 && dst->type == GGML_TYPE_I32)); GGML_ASSERT(src0->nb[0] == ggml_type_size(src0->type)); GGML_ASSERT(src1->nb[0] == ggml_type_size(src1->type)); @@ -162,6 +162,7 @@ void ggml_cuda_op_get_rows(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { get_rows_cuda_float(src0, src1, dst, (const half *)src0_d, src1_i32, dst_d, stream); break; case GGML_TYPE_F32: + case GGML_TYPE_I32: get_rows_cuda_float(src0, src1, dst, src0_d, src1_i32, dst_d, stream); break; case GGML_TYPE_Q4_0: diff --git a/ggml/src/ggml-cuda/indexer_topk.cu b/ggml/src/ggml-cuda/indexer_topk.cu index a110a7ac..adc3cb82 100644 --- a/ggml/src/ggml-cuda/indexer_topk.cu +++ b/ggml/src/ggml-cuda/indexer_topk.cu @@ -253,7 +253,7 @@ static __global__ void k_indexer_mask(int ne0, int ne1, int ne2, int ntopk, int void ggml_cuda_op_indexer_mask(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { auto mask = dst->src[0]; auto topk = dst->src[1]; - GGML_ASSERT(mask->ne[0] >= topk->ne[1]); + GGML_ASSERT(mask->ne[0] >= topk->ne[0]); GGML_ASSERT(mask->ne[1] >= topk->ne[1] && mask->ne[2] == topk->ne[2] && mask->ne[3] == topk->ne[3]); GGML_ASSERT(ggml_are_same_shape(mask, dst)); GGML_ASSERT(mask->type == GGML_TYPE_F16 || mask->type == GGML_TYPE_F32); @@ -277,3 +277,66 @@ void ggml_cuda_op_indexer_mask(ggml_backend_cuda_context & ctx, ggml_tensor * ds } } + +template +static __global__ void k_mask_to_index(int ne00, [[maybe_unused]] int ne0, + size_t nb01, size_t nb02, size_t nb03, + size_t nb1, size_t nb2, size_t nb3, + const mask_t * __restrict__ mask, int * __restrict__ idx) { + int i1 = blockIdx.x; + int i2 = blockIdx.y; + int i3 = blockIdx.z; + + mask_t zero; + if constexpr (std::is_same_v) { + zero = __float2half(0.0f); + } else { + zero = 0.0f; + } + __shared__ int counts[WARP_SIZE]; + const mask_t * mask_r = (const mask_t *)((const char *)mask + i1*nb01 + i2*nb02 + i3*nb03); + int * idx_r = (int *)((char *)idx + i1*nb1 + i2*nb2 + i3*nb3); + + for (int j = threadIdx.x; j < ne0; j += WARP_SIZE) { + idx_r[j] = -1; + } + + int nOn = 0; + for (int j = threadIdx.x; j < ne00; j += WARP_SIZE) { + nOn += (mask_r[j] == zero ? 1 : 0); + } + counts[threadIdx.x] = nOn; + __syncthreads(); + int cum[WARP_SIZE]; + int start = 0; + for (int i = 0; i < WARP_SIZE; ++i) { + cum[i] = start; + start += counts[i]; + } + start = cum[threadIdx.x]; + for (int j = threadIdx.x; j < ne00; j += WARP_SIZE) { + if (mask_r[j] == zero) idx_r[start++] = j; + } +} + +void ggml_cuda_op_mask_to_index(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { + auto src = dst->src[0]; + GGML_ASSERT(dst->type == GGML_TYPE_I32); + GGML_ASSERT(src->type == GGML_TYPE_F32 || src->type == GGML_TYPE_F16); + GGML_ASSERT(src->ne[1] == dst->ne[1] && src->ne[2] == dst->ne[2] && src->ne[3] == dst->ne[3]); + GGML_ASSERT(src->ne[0] >= dst->ne[0]); + + dim3 grid(dst->ne[1], dst->ne[2], dst->ne[3]); + if (src->type == GGML_TYPE_F16) { + k_mask_to_index<<>>(src->ne[0], dst->ne[0], + src->nb[1], src->nb[2], src->nb[3], + dst->nb[1], dst->nb[2], dst->nb[3], + (const half *)src->data, (int *)dst->data); + } else { + k_mask_to_index<<>>(src->ne[0], dst->ne[0], + src->nb[1], src->nb[2], src->nb[3], + dst->nb[1], dst->nb[2], dst->nb[3], + (const float *)src->data, (int *)dst->data); + } + +} diff --git a/ggml/src/ggml-cuda/indexer_topk.cuh b/ggml/src/ggml-cuda/indexer_topk.cuh index e131d26d..32325c89 100644 --- a/ggml/src/ggml-cuda/indexer_topk.cuh +++ b/ggml/src/ggml-cuda/indexer_topk.cuh @@ -8,3 +8,5 @@ void ggml_cuda_op_indexer_topk(ggml_backend_cuda_context & ctx, ggml_tensor * dst); void ggml_cuda_op_indexer_mask(ggml_backend_cuda_context & ctx, ggml_tensor * dst); + +void ggml_cuda_op_mask_to_index(ggml_backend_cuda_context & ctx, ggml_tensor * dst); diff --git a/ggml/src/ggml-cuda/sinkhorn.cu b/ggml/src/ggml-cuda/sinkhorn.cu index f02a8a44..a8d144bb 100644 --- a/ggml/src/ggml-cuda/sinkhorn.cu +++ b/ggml/src/ggml-cuda/sinkhorn.cu @@ -62,6 +62,113 @@ static __global__ void k_sinkhorn(const float * __restrict__ x, float * __restri } } +template +static __global__ void k_hc_pre(const float * __restrict__ x, const float * __restrict__ scale, + const float * __restrict__ bias, float * __restrict__ dst, + const int64_t T, const int iters, const float eps, const int64_t nb1) { + const int64_t t = (int64_t) blockIdx.x*blockDim.x + threadIdx.x; + if (t >= T) { + return; + } + + float m[S*S]; + + const float * x_pre = (const float *)((const char *) x + t*nb1); + const float * x_post = x_pre + S; + const float * x_comb = x_pre + 2*S; + + float * y_pre = dst + S*t; + float * y_post = dst + S*(T + t); + float * y_comb = dst + 2*S*T + S*S*t; + + #pragma unroll + for (int i = 0; i < S; ++i) { + float val = x_pre[i] * scale[0] + bias[i]; + y_pre[i] = 1.f / (1.f + expf(-val)) + eps; + } + + #pragma unroll + for (int i = 0; i < S; ++i) { + float val = x_post[i] * scale[1] + bias[S + i]; + y_post[i] = 2.f / (1.f + expf(-val)); + } + + #pragma unroll + for (int i = 0; i < S*S; ++i) { + m[i] = x_comb[i] * scale[2] + bias[2*S + i]; + } + + #pragma unroll + for (int r = 0; r < S; ++r) { + float mx = m[r*S]; + for (int c = 1; c < S; ++c) mx = fmaxf(mx, m[r*S + c]); + float sum = 0.0f; + for (int c = 0; c < S; ++c) { m[r*S + c] = expf(m[r*S + c] - mx); sum += m[r*S + c]; } + for (int c = 0; c < S; ++c) m[r*S + c] = m[r*S + c]/sum + eps; + } + #pragma unroll + for (int c = 0; c < S; ++c) { + float sum = eps; + for (int r = 0; r < S; ++r) sum += m[r*S + c]; + for (int r = 0; r < S; ++r) m[r*S + c] /= sum; + } + for (int i = 0; i < iters - 1; ++i) { + #pragma unroll + for (int r = 0; r < S; ++r) { + float sum = eps; + for (int c = 0; c < S; ++c) sum += m[r*S + c]; + for (int c = 0; c < S; ++c) m[r*S + c] /= sum; + } + #pragma unroll + for (int c = 0; c < S; ++c) { + float sum = eps; + for (int r = 0; r < S; ++r) sum += m[r*S + c]; + for (int r = 0; r < S; ++r) m[r*S + c] /= sum; + } + } + + #pragma unroll + for (int i = 0; i < S*S; ++i) y_comb[i] = m[i]; +} + +template +static __global__ void k_hc_post(int ne0, const float * __restrict__ x, const float * __restrict__ post, + const float * __restrict__ res, const float * __restrict__ comb, float * __restrict__ dst, + const int nelem, size_t nbx1, size_t nbp1, size_t nbc2, size_t nbr1, size_t nbr2, + size_t nb1, size_t nb2) { + + int ii = blockIdx.x*blockDim.x + threadIdx.x; + if (ii >= nelem) { + return; + } + int i0 = ii % ne0; + int i1 = ii / ne0; + + float r[S]; + + const float * x_r = (const float *)((const char *)x + i1*nbx1); + const float * post_r = (const float *)((const char *)post + i1*nbp1); + const float * comb_r = (const float *)((const char *)comb + i1*nbc2); + + #pragma unroll + for (int j = 0; j < S; ++j) { + const float * res_r = (const float *)((const char *)res + i1*nbr2 + j*nbr1); + r[j] = res_r[i0]; + } + + #pragma unroll + for (int i = 0; i < S; ++i) { + float sum = x_r[i0] * post_r[i]; + #pragma unroll + for (int j = 0; j < S; ++j) { + sum += comb_r[j*S + i] * r[j]; + } + float * dst_r = (float *)((char *)dst + i1*nb2 + i*nb1); + dst_r[i0] = sum; + } + +} + void ggml_cuda_op_sinkhorn(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { const ggml_tensor * src0 = dst->src[0]; @@ -101,3 +208,93 @@ void ggml_cuda_op_sinkhorn(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { default: GGML_ABORT("sinkhorn: unsupported S"); } } + +void ggml_cuda_op_hc_pre(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { + const ggml_tensor * src0 = dst->src[0]; + const ggml_tensor * src1 = dst->src[1]; + const ggml_tensor * src2 = dst->src[2]; + + const int S = dst->op_params[0]; + const int iters = dst->op_params[1]; + float eps; + memcpy(&eps, &dst->op_params[2], sizeof(float)); + const int64_t T = src0->ne[1]; + + GGML_ASSERT(src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F32 && src2->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + GGML_ASSERT(S >= 1 && S <= 8); + GGML_ASSERT(src0->ne[0] == (int64_t) S * S + 2 * S); + GGML_ASSERT(ggml_is_contiguous(dst)); + GGML_ASSERT(ggml_is_contiguous(src0)); + GGML_ASSERT(src1->ne[0] >= 3 && ggml_nrows(src1) == 1); + GGML_ASSERT(src2->ne[0] >= S*(S+2) && ggml_nrows(src2) == 1); + + if (T == 0) { + return; + } + + const int block = 256; + const int64_t grid = (T + block - 1)/block; + cudaStream_t stream = ctx.stream(); + + const float * x = (const float *) src0->data; + const float * s = (const float *) src1->data; + const float * b = (const float *) src2->data; + float * y = (float *) dst->data; + + switch (S) { + case 1: k_hc_pre<1><<>>(x, s, b, y, T, iters, eps, src0->nb[1]); break; + case 2: k_hc_pre<2><<>>(x, s, b, y, T, iters, eps, src0->nb[1]); break; + case 3: k_hc_pre<3><<>>(x, s, b, y, T, iters, eps, src0->nb[1]); break; + case 4: k_hc_pre<4><<>>(x, s, b, y, T, iters, eps, src0->nb[1]); break; + case 5: k_hc_pre<5><<>>(x, s, b, y, T, iters, eps, src0->nb[1]); break; + case 6: k_hc_pre<6><<>>(x, s, b, y, T, iters, eps, src0->nb[1]); break; + case 7: k_hc_pre<7><<>>(x, s, b, y, T, iters, eps, src0->nb[1]); break; + case 8: k_hc_pre<8><<>>(x, s, b, y, T, iters, eps, src0->nb[1]); break; + default: GGML_ABORT("hc_pre: unsupported S"); + } +} + +void ggml_cuda_op_hc_post(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { + const ggml_tensor * x = dst->src[0]; + const ggml_tensor * post = dst->src[1]; + const ggml_tensor * res = dst->src[2]; + const ggml_tensor * comb = dst->src[3]; + + GGML_ASSERT(x->type == GGML_TYPE_F32 && post->type == GGML_TYPE_F32 && res->type == GGML_TYPE_F32 && comb->type == GGML_TYPE_F32); + GGML_ASSERT(x->ne[0] == res->ne[0]); + GGML_ASSERT(x->ne[1] == res->ne[2] && x->ne[1] == post->ne[1] && x->ne[1] == comb->ne[2]); + GGML_ASSERT(x->ne[2] == 1 && x->ne[3] == 1 && post->ne[2] == 1 && post->ne[3] == 1 && res->ne[3] == 1 && comb->ne[3] == 1); + GGML_ASSERT(post->ne[0] == res->ne[1] && post->ne[0] == comb->ne[0] && post->ne[0] == comb->ne[1]); + + int S = post->ne[0]; + + constexpr int kBlockSize = 256; + int nelem = x->ne[0]*x->ne[1]; + int nblock = (nelem + kBlockSize - 1)/kBlockSize; + auto x_d = (const float *)x->data; + auto post_d = (const float *)post->data; + auto res_d = (const float *)res->data; + auto comb_d = (const float *)comb->data; + auto dst_d = (float *)dst->data; + + switch (S) { + case 1: k_hc_post<1><<>>(x->ne[0], x_d, post_d, res_d, comb_d, dst_d, + nelem, x->nb[1], post->nb[1], comb->nb[2], res->nb[1], res->nb[2], dst->nb[1], dst->nb[2]); break; + case 2: k_hc_post<2><<>>(x->ne[0], x_d, post_d, res_d, comb_d, dst_d, + nelem, x->nb[1], post->nb[1], comb->nb[2], res->nb[1], res->nb[2], dst->nb[1], dst->nb[2]); break; + case 3: k_hc_post<3><<>>(x->ne[0], x_d, post_d, res_d, comb_d, dst_d, + nelem, x->nb[1], post->nb[1], comb->nb[2], res->nb[1], res->nb[2], dst->nb[1], dst->nb[2]); break; + case 4: k_hc_post<4><<>>(x->ne[0], x_d, post_d, res_d, comb_d, dst_d, + nelem, x->nb[1], post->nb[1], comb->nb[2], res->nb[1], res->nb[2], dst->nb[1], dst->nb[2]); break; + case 5: k_hc_post<5><<>>(x->ne[0], x_d, post_d, res_d, comb_d, dst_d, + nelem, x->nb[1], post->nb[1], comb->nb[2], res->nb[1], res->nb[2], dst->nb[1], dst->nb[2]); break; + case 6: k_hc_post<6><<>>(x->ne[0], x_d, post_d, res_d, comb_d, dst_d, + nelem, x->nb[1], post->nb[1], comb->nb[2], res->nb[1], res->nb[2], dst->nb[1], dst->nb[2]); break; + case 7: k_hc_post<7><<>>(x->ne[0], x_d, post_d, res_d, comb_d, dst_d, + nelem, x->nb[1], post->nb[1], comb->nb[2], res->nb[1], res->nb[2], dst->nb[1], dst->nb[2]); break; + case 8: k_hc_post<8><<>>(x->ne[0], x_d, post_d, res_d, comb_d, dst_d, + nelem, x->nb[1], post->nb[1], comb->nb[2], res->nb[1], res->nb[2], dst->nb[1], dst->nb[2]); break; + default: GGML_ABORT("hc_post: unsupported S"); + } +} diff --git a/ggml/src/ggml-cuda/sinkhorn.cuh b/ggml/src/ggml-cuda/sinkhorn.cuh index c9f86186..583b7ea2 100644 --- a/ggml/src/ggml-cuda/sinkhorn.cuh +++ b/ggml/src/ggml-cuda/sinkhorn.cuh @@ -1,3 +1,7 @@ #include "common.cuh" void ggml_cuda_op_sinkhorn(ggml_backend_cuda_context & ctx, ggml_tensor * dst); + +void ggml_cuda_op_hc_pre(ggml_backend_cuda_context & ctx, ggml_tensor * dst); + +void ggml_cuda_op_hc_post(ggml_backend_cuda_context & ctx, ggml_tensor * dst); diff --git a/ggml/src/ggml-cuda/unary.cu b/ggml/src/ggml-cuda/unary.cu index 2098eee7..5104d2ed 100644 --- a/ggml/src/ggml-cuda/unary.cu +++ b/ggml/src/ggml-cuda/unary.cu @@ -740,6 +740,10 @@ static __device__ __forceinline__ float op_softplus(float x) { return (x > 20.0f) ? x : logf(1.0f + expf(x)); } +static __device__ __forceinline__ float op_sqrt_softplus(float x) { + return (x > 20.0f) ? sqrtf(x) : sqrtf(logf(1.0f + expf(x))); +} + static __device__ __forceinline__ float op_sin(float x) { return sinf(x); } @@ -840,6 +844,10 @@ void ggml_cuda_op_softplus(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { ggml_cuda_op_unary(ctx, dst); } +void ggml_cuda_op_sqrt_softplus(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { + ggml_cuda_op_unary(ctx, dst); +} + // === gated ops template diff --git a/ggml/src/ggml-cuda/unary.cuh b/ggml/src/ggml-cuda/unary.cuh index 6e09d780..e6c3d504 100644 --- a/ggml/src/ggml-cuda/unary.cuh +++ b/ggml/src/ggml-cuda/unary.cuh @@ -55,6 +55,8 @@ void ggml_cuda_op_exp(ggml_backend_cuda_context & ctx, ggml_tensor * dst); void ggml_cuda_op_softplus(ggml_backend_cuda_context & ctx, ggml_tensor * dst); +void ggml_cuda_op_sqrt_softplus(ggml_backend_cuda_context & ctx, ggml_tensor * dst); + void ggml_cuda_op_hardswish(ggml_backend_cuda_context & ctx, ggml_tensor * dst); void ggml_cuda_op_leaky_relu(ggml_backend_cuda_context & ctx, ggml_tensor * dst); diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index eff90454..c3bcf181 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -3308,8 +3308,10 @@ inline static void ggml_vec_relu_f32 (const int n, float * y, const float * x) { inline static void ggml_vec_leaky_relu_f32 (const int n, float * y, const float * x, const float ns) { for (int i = 0; i < n; ++i) y[i] = ((x[i] > 0.f) ? x[i] : 0.f) + ns * ((x[i] < 0.0f) ? x[i] : 0.f); } inline static void ggml_vec_sigmoid_f32 (const int n, float * y, const float * x) { for (int i = 0; i < n; ++i) y[i] = 1.f / (1.f + expf(-x[i])); } inline static float ggml_compute_softplus_f32(const float x) { return x > 20.0f ? x : logf(1.0f + expf(x)); } +inline static float ggml_compute_sqrt_softplus_f32(const float x) { return sqrtf(ggml_compute_softplus_f32(x)); } inline static void ggml_vec_exp_f32 (const int n, float * y, const float * x) { for (int i = 0; i < n; ++i) y[i] = expf(x[i]); } inline static void ggml_vec_softplus_f32 (const int n, float * y, const float * x) { for (int i = 0; i < n; ++i) y[i] = ggml_compute_softplus_f32(x[i]); } +inline static void ggml_vec_sqrt_softplus_f32 (const int n, float * y, const float * x) { for (int i = 0; i < n; ++i) y[i] = ggml_compute_sqrt_softplus_f32(x[i]); } // TODO: optimize performance inline static void ggml_vec_hardswish_f32 (const int n, float * y, const float * x) { for (int i = 0; i < n; ++i) y[i] = x[i] * fminf(1.0f, fmaxf(0.0f, (x[i] + 3.0f) / 6.0f)); } inline static void ggml_vec_hardsigmoid_f32 (const int n, float * y, const float * x) { for (int i = 0; i < n; ++i) y[i] = fminf(1.0f, fmaxf(0.0f, (x[i] + 3.0f) / 6.0f)); } @@ -4335,9 +4337,12 @@ static const char * GGML_OP_NAME[GGML_OP_COUNT] = { "INDEXER_TOPK", "MASK_TOPK", "SINKHORN", + "HC_PRE", + "HC_POST", + "MASK_TO_IDX", }; -static_assert(GGML_OP_COUNT == 106, "GGML_OP_COUNT != 106"); +static_assert(GGML_OP_COUNT == 109, "GGML_OP_COUNT != 109"); static const char * GGML_OP_SYMBOL[GGML_OP_COUNT] = { "none", @@ -4459,10 +4464,13 @@ static const char * GGML_OP_SYMBOL[GGML_OP_COUNT] = { "indexer_topk(k, q, w, mask)", "mask_topk(mask, topk)", "sinkhorn(x)", + "hc_pre(x,s,b)", + "hc_post(x,p,r,c)", + "mask_to_idx(masl)", }; -static_assert(GGML_OP_COUNT == 106, "GGML_OP_COUNT != 106"); +static_assert(GGML_OP_COUNT == 109, "GGML_OP_COUNT != 109"); static_assert(GGML_OP_POOL_COUNT == 2, "GGML_OP_POOL_COUNT != 2"); @@ -4486,9 +4494,10 @@ static const char * GGML_UNARY_OP_NAME[GGML_UNARY_OP_COUNT] = { "GELU_ERF", "EXP", "SOFTPLUS", + "SQRT_SOFTPLUS", }; -static_assert(GGML_UNARY_OP_COUNT == 18, "GGML_UNARY_OP_COUNT != 18"); +static_assert(GGML_UNARY_OP_COUNT == 19, "GGML_UNARY_OP_COUNT != 19"); static_assert(sizeof(struct ggml_object)%GGML_MEM_ALIGN == 0, "ggml_object size must be a multiple of GGML_MEM_ALIGN"); @@ -7169,7 +7178,19 @@ struct ggml_tensor * ggml_softplus( struct ggml_tensor * ggml_softplus_inplace( struct ggml_context * ctx, struct ggml_tensor * a) { - return ggml_unary_inplace(ctx, a, GGML_UNARY_OP_SOFTPLUS); + return ggml_unary_inplace(ctx, a, GGML_UNARY_OP_SQRT_SOFTPLUS); +} + +struct ggml_tensor * ggml_sqrt_softplus( + struct ggml_context * ctx, + struct ggml_tensor * a) { + return ggml_unary(ctx, a, GGML_UNARY_OP_SQRT_SOFTPLUS); +} + +struct ggml_tensor * ggml_sqrt_softplus_inplace( + struct ggml_context * ctx, + struct ggml_tensor * a) { + return ggml_unary_inplace(ctx, a, GGML_UNARY_OP_SQRT_SOFTPLUS); } // ggml_gelu @@ -10189,6 +10210,70 @@ struct ggml_tensor * ggml_sinkhorn( return result; } +struct ggml_tensor * ggml_hc_pre( + struct ggml_context * ctx, + struct ggml_tensor * x, + struct ggml_tensor * scale, + struct ggml_tensor * bias, + int S, + int n_iters, + float eps) { + GGML_ASSERT(x->type == GGML_TYPE_F32 && scale->type == GGML_TYPE_F32 && bias->type == GGML_TYPE_F32); + int ntot = S*S + 2*S; + GGML_ASSERT(scale->ne[0] == 3 && ggml_nrows(scale) == 1); + GGML_ASSERT(bias->ne[0] == ntot && ggml_nrows(bias) == 1); + GGML_ASSERT(ggml_is_contiguous(x)); + GGML_ASSERT(x->ne[0] == ntot && x->ne[2] == 1 && x->ne[3] == 1); + + struct ggml_tensor * result = ggml_new_tensor_1d(ctx,GGML_TYPE_F32, ntot*x->ne[1]); + result->op = GGML_OP_HC_PRE; + result->src[0] = x; + result->src[1] = scale; + result->src[2] = bias; + result->op_params[0] = S; + result->op_params[1] = n_iters; + memcpy(&result->op_params[2], &eps, sizeof(eps)); + + return result; +} + +struct ggml_tensor * ggml_hc_post( + struct ggml_context * ctx, + struct ggml_tensor * x, + struct ggml_tensor * post, + struct ggml_tensor * res, + struct ggml_tensor * comb) { + GGML_ASSERT(x->type == GGML_TYPE_F32 && post->type == GGML_TYPE_F32 && res->type == GGML_TYPE_F32 && comb->type == GGML_TYPE_F32); + GGML_ASSERT(x->ne[0] == res->ne[0]); + GGML_ASSERT(x->ne[1] == res->ne[2] && x->ne[1] == post->ne[1] && x->ne[1] == comb->ne[2]); + GGML_ASSERT(x->ne[2] == 1 && x->ne[3] == 1 && post->ne[2] == 1 && post->ne[3] == 1 && res->ne[3] == 1 && comb->ne[3] == 1); + GGML_ASSERT(post->ne[0] == res->ne[1] && post->ne[0] == comb->ne[0] && post->ne[0] == comb->ne[1]); + + struct ggml_tensor * result = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, x->ne[0], post->ne[0], x->ne[1]); + result->src[0] = x; + result->src[1] = post; + result->src[2] = res; + result->src[3] = comb; + result->op = GGML_OP_HC_POST; + + return result; +} + +struct ggml_tensor * ggml_mask_to_index( + struct ggml_context * ctx, + struct ggml_tensor * mask, + int max_row_size) { + GGML_ASSERT(mask->type == GGML_TYPE_F16 || mask->type == GGML_TYPE_F32); + + int64_t ne0 = MIN(mask->ne[0], max_row_size); + struct ggml_tensor * result = ggml_new_tensor_4d(ctx, GGML_TYPE_I32, ne0, mask->ne[1], mask->ne[2], mask->ne[3]); + result->src[0] = mask; + result->op = GGML_OP_MASK_TO_IDX; + + return result; +} + + // ggml_fill static struct ggml_tensor * ggml_fill_impl( @@ -10196,7 +10281,7 @@ static struct ggml_tensor * ggml_fill_impl( struct ggml_tensor * a, float c, bool inplace) { - GGML_ASSERT(a->type == GGML_TYPE_F32); + GGML_ASSERT(a->type == GGML_TYPE_F32 || a->type == GGML_TYPE_F16); GGML_ASSERT(ggml_is_contiguous(a)); bool is_node = false; @@ -12778,6 +12863,57 @@ static void ggml_compute_forward_add_f16_f32( } } +static void ggml_compute_forward_add_f32_f16( + const struct ggml_compute_params * params, + struct ggml_tensor * dst) { + + const struct ggml_tensor * src0 = dst->src[0]; + const struct ggml_tensor * src1 = dst->src[1]; + + GGML_ASSERT(ggml_are_same_shape(src0, src1) && ggml_are_same_shape(src0, dst)); + + const int ith = params->ith; + const int nth = params->nth; + + const int nr = ggml_nrows(src0); + + GGML_TENSOR_BINARY_OP_LOCALS + + GGML_ASSERT(src0->type == GGML_TYPE_F32); + GGML_ASSERT(src1->type == GGML_TYPE_F16); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + GGML_ASSERT( nb0 == sizeof(float)); + GGML_ASSERT(nb00 == sizeof(float)); + + // rows per thread + const int dr = (nr + nth - 1)/nth; + + // row range for this thread + const int ir0 = dr*ith; + const int ir1 = MIN(ir0 + dr, nr); + + if (nb10 == sizeof(ggml_fp16_t)) { + for (int ir = ir0; ir < ir1; ++ir) { + // src0, src1 and dst are same shape => same indices + const int i3 = ir/(ne2*ne1); + const int i2 = (ir - i3*ne2*ne1)/ne1; + const int i1 = (ir - i3*ne2*ne1 - i2*ne1); + + float * dst_ptr = (float *) ((char *) dst->data + i3*nb3 + i2*nb2 + i1*nb1); + float * src0_ptr = (float *) ((char *) src0->data + i3*nb03 + i2*nb02 + i1*nb01); + ggml_fp16_t * src1_ptr = (ggml_fp16_t *) ((char *) src1->data + i3*nb13 + i2*nb12 + i1*nb11); + + for (int i = 0; i < ne0; i++) { + dst_ptr[i] = src0_ptr[i] + GGML_FP16_TO_FP32(src1_ptr[i]); + } + } + } + else { + // src1 is not contiguous + GGML_ABORT("fatal error"); + } +} + static void ggml_compute_forward_add_bf16_f32( const struct ggml_compute_params * params, struct ggml_tensor * dst) { @@ -13047,7 +13183,7 @@ static void ggml_compute_forward_add( ggml_compute_forward_add_f32(params, dst); } else { - GGML_ABORT("fatal error"); + ggml_compute_forward_add_f32_f16(params, dst); } } break; case GGML_TYPE_F16: @@ -15080,6 +15216,93 @@ static void ggml_compute_forward_repeat_back( // ggml_compute_forward_concat +static bool ggml_compute_forward_concat_any_opt( + const struct ggml_compute_params * params, + struct ggml_tensor * dst) { + + + const struct ggml_tensor * src0 = dst->src[0]; + const struct ggml_tensor * src1 = dst->src[1]; + + if (ggml_is_quantized(src0->type)) return false; + + GGML_ASSERT(src0->type == src1->type && src0->type == dst->type); + GGML_ASSERT(!ggml_is_quantized(src0->type)); + + const int ith = params->ith; + const int nth = params->nth; + + const int32_t dim = ggml_get_op_params_i32(dst, 0); + + GGML_ASSERT(dim >= 0 && dim < 4); + + if (dim == 0) { + size_t elem_size = ggml_element_size(dst); + if (src0->nb[0] == elem_size && src1->nb[0] == elem_size && dst->nb[0] == elem_size) { + GGML_ASSERT(dst->ne[0] == src0->ne[0] + src1->ne[0]); + int nrows = 1; + for (int d = 1; d < GGML_MAX_DIMS; ++d) { + GGML_ASSERT(src0->ne[d] == src1->ne[d]); + GGML_ASSERT(dst->ne[d] == src0->ne[d]); + nrows *= dst->ne[d]; + } + size_t row_size_0 = ggml_row_size(dst->type, src0->ne[0]); + size_t row_size_1 = ggml_row_size(dst->type, src1->ne[0]); + int npt = (nrows + nth - 1)/nth; + int first = ith*npt; + int last = MIN(first + npt, nrows); + for (int ir = first; ir < last; ++ir) { + int ii = ir; + int i3 = ii/(dst->ne[1]*dst->ne[2]); ii -= i3*dst->ne[1]*dst->ne[2]; + int i2 = ii/dst->ne[1]; + int i1 = ii - i2*dst->ne[1]; + const char * c0 = (const char *)src0->data + i1*src0->nb[1] + i2*src0->nb[2] + i3*src0->nb[3]; + const char * c1 = (const char *)src1->data + i1*src1->nb[1] + i2*src1->nb[2] + i3*src1->nb[3]; + char * y = (char *)dst->data + i1*dst->nb[1] + i2*dst->nb[2] + i3*dst->nb[3]; + memcpy(y, c0, row_size_0); + memcpy(y + row_size_0, c1, row_size_1); + } + return true; + } + } + if (dim != 0) { + GGML_ASSERT(dst->ne[dim] == src0->ne[dim] + src1->ne[dim]); + int nrows = 1; + for (int d = 0; d < GGML_MAX_DIMS; ++d) { + if (d != dim) { + GGML_ASSERT(src0->ne[d] == src1->ne[d]); + GGML_ASSERT(dst->ne[d] == src0->ne[d]); + } + if (d > 0) nrows *= dst->ne[d]; + } + size_t row_size = ggml_row_size(dst->type, dst->ne[0]); + if (src0->nb[1] == row_size && src1->nb[1] == row_size) { + int npt = (nrows + nth - 1)/nth; + int first = ith*npt; + int last = MIN(first + npt, nrows); + int idx[4]; + for (int ir = first; ir < last; ++ir) { + int ii = ir; + idx[3] = ii/(dst->ne[1]*dst->ne[2]); ii -= idx[3]*dst->ne[1]*dst->ne[2]; + idx[2] = ii/dst->ne[1]; ii -= idx[2]*dst->ne[1]; + idx[1] = ii; + char * y = (char *)dst->data + idx[1]*dst->nb[1] + idx[2]*dst->nb[2] + idx[3]*dst->nb[3]; + const char * x; + if (idx[dim] < src0->ne[dim]) { + x = (const char *)src0->data + idx[1]*src0->nb[1] + idx[2]*src0->nb[2] + idx[3]*src0->nb[3]; + } else { + idx[dim] -= src0->ne[dim]; + x = (const char *)src1->data + idx[1]*src1->nb[1] + idx[2]*src1->nb[2] + idx[3]*src1->nb[3]; + } + memcpy(y, x, row_size); + } + return true; + } + } + + return false; +} + static void ggml_compute_forward_concat_f32( const struct ggml_compute_params * params, struct ggml_tensor * dst) { @@ -15143,6 +15366,10 @@ static void ggml_compute_forward_concat_f32( return; } + if (ggml_compute_forward_concat_any_opt(params, dst)) { + return; + } + int64_t o[4] = {0, 0, 0, 0}; o[dim] = src0->ne[dim]; @@ -15172,61 +15399,80 @@ static void ggml_compute_forward_concat_any( const struct ggml_compute_params * params, struct ggml_tensor * dst) { + if (ggml_compute_forward_concat_any_opt(params, dst)) { + return; + } + const struct ggml_tensor * src0 = dst->src[0]; const struct ggml_tensor * src1 = dst->src[1]; GGML_ASSERT(src0->type == src1->type && src0->type == dst->type); + GGML_ASSERT(!ggml_is_quantized(src0->type)); + + const size_t len = ggml_type_size(src0->type); + + const int ith = params->ith; + const int nth = params->nth; const int32_t dim = ggml_get_op_params_i32(dst, 0); - int ith = params->ith; - int nth = params->nth; + GGML_ASSERT(dim >= 0 && dim < 4); - int64_t nrows = ggml_nrows(dst); - int64_t nrows_per_thread = (nrows + nth - 1)/nth; - int64_t first_row = ith*nrows_per_thread; - int64_t last_row = MIN(first_row + nrows_per_thread, nrows); - if (first_row >= last_row) return; + //if (ith == 0 && (strcmp(dst->name, "csa_kq_mask-2") == 0 || strcmp(dst->name, "hca_kq_mask-3") == 0)) { + // printf("%s: %s, %s: %ld x %ld x %ld x %ld + %ld x %ld x %ld x %ld\n", __func__, dst->name, ggml_type_name(dst->type), + // src0->ne[0], src0->ne[1], src0->ne[2], src0->ne[3], + // src1->ne[0], src1->ne[1], src1->ne[2], src1->ne[3]); + // const uint16_t * M = (const uint16_t *)src0->data; + // int n = src0->ne[0]; + // int first = 0, last = n; + // for ( ; first < n; ++first) if (M[first] == 0) break; + // for ( ; last > first; --last) if (M[last-1] == 0) break; + // int nOn = 0; + // for (int j = first; j < last; ++j) if (M[j] == 0) ++nOn; + // printf(" src0: %d ON in %d...%d\n", nOn, first, last); + // M = (const uint16_t *)src1->data; + // n = src1->ne[0]; + // first = 0; last = n; + // for ( ; first < n; ++first) if (M[first] == 0) break; + // for ( ; last > first; --last) if (M[last-1] == 0) break; + // nOn = 0; + // for (int j = first; j < last; ++j) if (M[j] == 0) ++nOn; + // printf(" src1: %d ON in %d...%d\n", nOn, first, last); + //} - int64_t src0_row_size = ggml_row_size(src0->type, src0->ne[0]); - int64_t src1_row_size = ggml_row_size(src1->type, src1->ne[0]); + GGML_TENSOR_BINARY_OP_LOCALS - if (dim == 0) { - for (int64_t row = first_row; row < last_row; ++row) { - int64_t i3 = row/(dst->ne[1]*dst->ne[2]); - int64_t i2 = (row - i3*dst->ne[1]*dst->ne[2])/dst->ne[1]; - int64_t i1 = row - i3*dst->ne[1]*dst->ne[2] - i2*dst->ne[1]; - char * y = (char *)dst->data + i1*dst->nb[1] + i2*dst->nb[2] + i3*dst->nb[3]; - const char * x0 = (const char *)src0->data + i1*src0->nb[1] + i2*src0->nb[2] + i3*src0->nb[3]; - const char * x1 = (const char *)src1->data + i1*src1->nb[1] + i2*src1->nb[2] + i3*src1->nb[3]; - memcpy(y, x0, src0_row_size); - memcpy(y + src0_row_size, x1, src1_row_size); - } - } - else { - GGML_ASSERT(src0_row_size == src1_row_size); - for (int64_t row = first_row; row < last_row; ++row) { - int64_t i3 = row/(dst->ne[1]*dst->ne[2]); - int64_t i2 = (row - i3*dst->ne[1]*dst->ne[2])/dst->ne[1]; - int64_t i1 = row - i3*dst->ne[1]*dst->ne[2] - i2*dst->ne[1]; - char * y = (char *)dst->data + i1*dst->nb[1] + i2*dst->nb[2] + i3*dst->nb[3]; - const char * x; - if (dim == 1) { - x = i1 < src0->ne[1] ? (const char *)src0->data + i1*src0->nb[1] + i2*src0->nb[2] + i3*src0->nb[3] - : (const char *)src1->data + (i1 - src0->ne[1])*src1->nb[1] + i2*src1->nb[2] + i3*src1->nb[3]; - } - else if (dim == 2) { - x = i2 < src0->ne[2] ? (const char *)src0->data + i1*src0->nb[1] + i2*src0->nb[2] + i3*src0->nb[3] - : (const char *)src1->data + i1*src1->nb[1] + (i2 - src0->ne[2])*src1->nb[2] + i3*src1->nb[3]; - } - else { - x = i3 < src0->ne[3] ? (const char *)src0->data + i1*src0->nb[1] + i2*src0->nb[2] + i3*src0->nb[3] - : (const char *)src1->data + i1*src1->nb[1] + i2*src1->nb[2] + (i3 - src0->ne[3])*src1->nb[3]; - } - memcpy(y, x, src0_row_size); + for (int d = 0; d < 4; ++d) { + if (d == dim) { + GGML_ASSERT(dst->ne[d] == src0->ne[d] + src1->ne[d]); + } else { + GGML_ASSERT(src0->ne[d] == src1->ne[d]); + GGML_ASSERT(dst->ne[d] == src0->ne[d]); } } + int64_t o[4] = { 0, 0, 0, 0 }; + o[dim] = src0->ne[dim]; + + const char * x; + + // Support strided DSV4 views. + for (int64_t i3 = 0; i3 < ne3; ++i3) { + for (int64_t i2 = ith; i2 < ne2; i2 += nth) { + for (int64_t i1 = 0; i1 < ne1; ++i1) { + for (int64_t i0 = 0; i0 < ne0; ++i0) { + if (i0 < ne00 && i1 < ne01 && i2 < ne02 && i3 < ne03) { + x = (const char *) src0->data + (i0 )*nb00 + (i1 )*nb01 + (i2 )*nb02 + (i3 )*nb03; + } else { + x = (const char *) src1->data + (i0 - o[0])*nb10 + (i1 - o[1])*nb11 + (i2 - o[2])*nb12 + (i3 - o[3])*nb13; + } + + char * y = (char *) dst->data + i0*nb0 + i1*nb1 + i2*nb2 + i3*nb3; + memcpy(y, x, len); + } + } + } + } } static void ggml_compute_forward_concat( @@ -15699,6 +15945,53 @@ static void ggml_compute_forward_softplus( } } +// ggml_compute_forward_sqrt_softplus + +static void ggml_compute_forward_sqrt_softplus_f32( + const struct ggml_compute_params * params, + struct ggml_tensor * dst) { + + const struct ggml_tensor * src0 = dst->src[0]; + + assert(ggml_is_contiguous_1(src0)); + assert(ggml_is_contiguous_1(dst)); + assert(ggml_are_same_shape(src0, dst)); + + const int ith = params->ith; + const int nth = params->nth; + + const int nc = src0->ne[0]; + const int nr = ggml_nrows(src0); + + const int dr = (nr + nth - 1)/nth; + const int ir0 = dr*ith; + const int ir1 = MIN(ir0 + dr, nr); + + for (int i1 = ir0; i1 < ir1; i1++) { + ggml_vec_sqrt_softplus_f32(nc, + (float *) ((char *) dst->data + i1*( dst->nb[1])), + (float *) ((char *) src0->data + i1*(src0->nb[1]))); + } +} + +static void ggml_compute_forward_sqrt_softplus( + const struct ggml_compute_params * params, + struct ggml_tensor * dst) { + + const struct ggml_tensor * src0 = dst->src[0]; + + switch (src0->type) { + case GGML_TYPE_F32: + { + ggml_compute_forward_sqrt_softplus_f32(params, dst); + } break; + default: + { + GGML_ABORT("fatal error"); + } + } +} + // ggml_compute_forward_gelu static void ggml_compute_forward_gelu_f32( @@ -15839,8 +16132,37 @@ static void ggml_compute_forward_fill_f32(const struct ggml_compute_params * par } } +static void ggml_compute_forward_fill_f16(const struct ggml_compute_params * params, struct ggml_tensor * dst) { + const ggml_fp16_t c = GGML_FP32_TO_FP16(ggml_get_op_params_f32(dst, 0)); + + GGML_TENSOR_LOCALS(int64_t, ne, dst, ne); + GGML_TENSOR_LOCALS(size_t, nb, dst, nb); + + const int ith = params->ith; + const int nth = params->nth; + const int64_t nr = ne1*ne2*ne3; + + for (int64_t ir = ith; ir < nr; ir += nth) { + const int64_t i03 = ir/(ne2*ne1); + const int64_t i02 = (ir - i03*ne2*ne1)/ne1; + const int64_t i01 = ir - i03*ne2*ne1 - i02*ne1; + + ggml_fp16_t * dst_ptr = (ggml_fp16_t *) ((char *) dst->data + i03*nb3 + i02*nb2 + i01*nb1); + ggml_vec_set_f16(ne0, dst_ptr, c); + } +} + static void ggml_compute_forward_fill(const struct ggml_compute_params * params, struct ggml_tensor * dst) { - ggml_compute_forward_fill_f32(params, dst); + switch (dst->type) { + case GGML_TYPE_F32: + ggml_compute_forward_fill_f32(params, dst); + break; + case GGML_TYPE_F16: + ggml_compute_forward_fill_f16(params, dst); + break; + default: + GGML_ABORT("unsupported type for ggml_fill: %s", ggml_type_name(dst->type)); + } } // ggml_compute_forward_tri @@ -22160,7 +22482,9 @@ static void ggml_compute_forward_flash_attn_ext_f16( #if GGML_USE_IQK_MULMAT // For now we do not implement sinks in the iqk FA implementation - if (iqk_flash_attn_noalibi(q->type, mask ? mask->type : GGML_TYPE_F16, max_bias, + // DSV4 marks its FA nodes with the shared generic-FA backend hint. + const bool use_iqk_fa = dst->op_params[4] != GGML_FLASH_ATTN_EXT_IQK_DISABLED; + if (use_iqk_fa && iqk_flash_attn_noalibi(q->type, mask ? mask->type : GGML_TYPE_F16, max_bias, q->ne[3], q->ne[2], q->nb[3], q->nb[2], k->ne[3], k->ne[2], k->nb[3], k->nb[2], v->ne[3], v->ne[2], v->nb[3], v->nb[2], @@ -22350,7 +22674,7 @@ static void ggml_compute_forward_flash_attn_ext_f16( } // V /= S - const float S_inv = 1.0f/S; + const float S_inv = S == 0.0f ? 0.0f : 1.0f/S; ggml_vec_scale_f32(Dv, VKQ32, S_inv); // dst indices @@ -23302,6 +23626,291 @@ static void ggml_compute_forward_sinkhorn( } } +// ggml_compute_forward_hc_pre + +static void ggml_compute_forward_hc_pre_f32( + const struct ggml_compute_params * params, + struct ggml_tensor * dst) { + const struct ggml_tensor * src0 = dst->src[0]; + const struct ggml_tensor * src1 = dst->src[1]; + const struct ggml_tensor * src2 = dst->src[2]; + + const int S = dst->op_params[0]; + const int iters = dst->op_params[1]; + float eps; + memcpy(&eps, &dst->op_params[2], sizeof(float)); + const int64_t T = src0->ne[1]; + + GGML_ASSERT(S >= 1 && S <= 8); + GGML_ASSERT(src0->ne[0] == (int64_t) S * S + 2 * S); + GGML_ASSERT(iters > 1); + + int ith = params->ith; + int nth = params->nth; + + int64_t npt = (T + nth - 1)/nth; + + // one token is S*S floats (16 at S=4): parallelize over tokens only + const int64_t t0 = ith * npt; + const int64_t t1 = MIN(t0 + npt, T); + + float m[64]; + + const float * scale = (const float *)src1->data; + const float * bias = (const float *)src2->data; + + for (int64_t t = t0; t < t1; ++t) { + const float * x_pre = (const float *)((const char *)src0->data + t*src0->nb[1]); + const float * x_post = x_pre + S; + const float * x_comb = x_pre + 2*S; + float * y_pre = (float *)dst->data + S*t; + float * y_post = (float *)dst->data + S*(T + t); + float * y_comb = (float *)dst->data + 2*S*T + S*S*t; + + for (int i = 0; i < S; ++i) { + float val = x_pre[i] * scale[0] + bias[i]; + val = 1.f / (1.f + expf(-val)); + y_pre[i] = val + eps; + + val = x_post[i] * scale[1] + bias[S + i]; + val = 1.f / (1.f + expf(-val)); + y_post[i] = 2.0f * val; + } + + for (int i = 0; i < S*S; ++i) { + m[i] = x_comb[i] * scale[2] + bias[2*S + i]; + } + + for (int r = 0; r < S; ++r) { + float mx = m[r*S]; + for (int c = 1; c < S; ++c) mx = MAX(mx, m[r*S + c]); + float sum = 0.0f; + for (int c = 0; c < S; ++c) { m[r*S + c] = expf(m[r*S + c] - mx); sum += m[r*S + c]; } + for (int c = 0; c < S; ++c) m[r*S + c] = m[r*S + c]/sum + eps; + } + + for (int c = 0; c < S; ++c) { + float sum = eps; + for (int r = 0; r < S; ++r) sum += m[r*S + c]; + for (int r = 0; r < S; ++r) m[r*S + c] /= sum; + } + + for (int i = 0; i < iters - 1; ++i) { + for (int r = 0; r < S; ++r) { + float sum = eps; + for (int c = 0; c < S; ++c) sum += m[r*S + c]; + for (int c = 0; c < S; ++c) m[r*S + c] /= sum; + } + for (int c = 0; c < S; ++c) { + float sum = eps; + for (int r = 0; r < S; ++r) sum += m[r*S + c]; + for (int r = 0; r < S; ++r) m[r*S + c] /= sum; + } + } + for (int i = 0; i < S*S; ++i) { + y_comb[i] = m[i]; + } + //if (transposed) { + // // dst is [row, col, T] (ne0 = row): transpose on write + // for (int c = 0; c < S; ++c) { + // for (int r = 0; r < S; ++r) y[c*S + r] = m[r*S + c]; + // } + //} else { + // for (int k = 0; k < S*S; ++k) y[k] = m[k]; + //} + } +} + +// ggml_compute_forward_hc_post + +//build_mhc_post: x = 4096 x 4096 x 1 x 1, post = 4 x 4096 x 1 x 1, residual = 4096 x 4 x 4096 x 1, comb = 4 x 4 x 4096 x 1 +//build_mhc_post: x = 4096 x 1 x 1 x 1, post = 4 x 1 x 1 x 1, residual = 4096 x 4 x 1 x 1, comb = 4 x 4 x 1 x 1 + +static void ggml_compute_forward_hc_post_f32( + const struct ggml_compute_params * params, + struct ggml_tensor * dst) { + const struct ggml_tensor * x = dst->src[0]; + const struct ggml_tensor * post = dst->src[1]; + const struct ggml_tensor * res = dst->src[2]; + const struct ggml_tensor * comb = dst->src[3]; + + int ith = params->ith; + int nth = params->nth; + + GGML_ASSERT(x->type == GGML_TYPE_F32 && post->type == GGML_TYPE_F32 && res->type == GGML_TYPE_F32 && comb->type == GGML_TYPE_F32); + GGML_ASSERT(x->ne[0] == res->ne[0]); + GGML_ASSERT(x->ne[1] == res->ne[2] && x->ne[1] == post->ne[1] && x->ne[1] == comb->ne[2]); + GGML_ASSERT(x->ne[2] == 1 && x->ne[3] == 1 && post->ne[2] == 1 && post->ne[3] == 1 && res->ne[3] == 1 && comb->ne[3] == 1); + GGML_ASSERT(post->ne[0] == res->ne[1] && post->ne[0] == comb->ne[0] && post->ne[0] == comb->ne[1]); + + int ne0 = x->ne[0]; + int S = post->ne[0]; + GGML_ASSERT(S <= 8); + + float r[8]; + + int T = x->ne[1]; + //const int i0_chunk = 64; + //if (T < nth && ne0 > i0_chunk && ne0 % i0_chunk == 0) { + // int ne0_64 = ne0/i0_chunk; + // int nchunk = T*ne0_64; + // int npt = (nchunk + nth - 1)/nth; + // int first = ith * npt; + // int last = MIN(first + npt, nchunk); + + // for (int ic = first; ic < last; ++ic) { + // int t = ic / ne0_64; + // int i0_first = ic - t * ne0_64; + + // const float * x_r = (const float *)((const char *)x->data + t*x->nb[1]) + i0_first; + // const float * post_r = (const float *)((const char *)post->data + t*post->nb[1]); + // const float * comb_r = (const float *)((const char *)comb->data + t*comb->nb[2]); + + // for (int i0 = 0; i0 < i0_chunk; ++i0) { + // for (int j = 0; j < S; ++j) { + // const float * res_r = (const float *)((const char *)res->data + t*res->nb[2] + j*res->nb[1]); + // r[j] = res_r[i0_first + i0]; + // } + // for (int i = 0; i < S; ++i) { + // float sum = x_r[i0] * post_r[i]; + // for (int j = 0; j < S; ++j) { + // sum += comb_r[j*S + i] * r[j]; + // } + // float * dst_r = (float *)((char *)dst->data + t*dst->nb[2] + i*dst->nb[1]) + i0_first; + // dst_r[i0] = sum; + // } + // } + // } + // return; + //} + + if (T == 1) { + const float * x_r = (const float *)((const char *)x->data); + const float * post_r = (const float *)((const char *)post->data); + const float * comb_r = (const float *)((const char *)comb->data); + + int nchunk = (ne0 + nth - 1)/nth; + + for (int ic = ith; ic < nchunk; ic += nth) { + int first = 64*ic; + for (int i0 = first; i0 < first + 64 && i0 < ne0; ++i0) { + for (int j = 0; j < S; ++j) { + const float * res_r = (const float *)((const char *)res->data + j*res->nb[1]); + r[j] = res_r[i0]; + } + for (int i = 0; i < S; ++i) { + float sum = x_r[i0] * post_r[i]; + for (int j = 0; j < S; ++j) { + sum += comb_r[j*S + i] * r[j]; + } + float * dst_r = (float *)((char *)dst->data + i*dst->nb[1]); + dst_r[i0] = sum; + } + } + } + return; + } + + int64_t npt = (T + nth - 1)/nth; + + // one token is S*S floats (16 at S=4): parallelize over tokens only + const int t0 = ith * npt; + const int t1 = MIN(t0 + npt, T); + + for (int64_t t = t0; t < t1; ++t) { + const float * x_r = (const float *)((const char *)x->data + t*x->nb[1]); + const float * post_r = (const float *)((const char *)post->data + t*post->nb[1]); + const float * comb_r = (const float *)((const char *)comb->data + t*comb->nb[2]); + + for (int i0 = 0; i0 < ne0; ++i0) { + for (int j = 0; j < S; ++j) { + const float * res_r = (const float *)((const char *)res->data + t*res->nb[2] + j*res->nb[1]); + r[j] = res_r[i0]; + } + for (int i = 0; i < S; ++i) { + float sum = x_r[i0] * post_r[i]; + for (int j = 0; j < S; ++j) { + //sum += comb_r[i*S + j] * r[j]; + sum += comb_r[j*S + i] * r[j]; + } + float * dst_r = (float *)((char *)dst->data + t*dst->nb[2] + i*dst->nb[1]); + dst_r[i0] = sum; + } + } + } +} + +static void ggml_compute_forward_hc_post( + const struct ggml_compute_params * params, + struct ggml_tensor * dst) { + switch (dst->src[0]->type) { + case GGML_TYPE_F32: + ggml_compute_forward_hc_post_f32(params, dst); + break; + default: + GGML_ABORT("fatal error"); + } +} + +static void ggml_compute_forward_hc_pre( + const struct ggml_compute_params * params, + struct ggml_tensor * dst) { + switch (dst->src[0]->type) { + case GGML_TYPE_F32: + ggml_compute_forward_hc_pre_f32(params, dst); + break; + default: + GGML_ABORT("fatal error"); + } +} + +static void ggml_compute_forward_mask_to_idx(const struct ggml_compute_params * params, + struct ggml_tensor * dst) { + struct ggml_tensor * src = dst->src[0]; + GGML_ASSERT(src->type == GGML_TYPE_F16 || src->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_I32); + GGML_ASSERT(src->ne[1] == dst->ne[1] && src->ne[2] == dst->ne[2] && src->ne[3] == dst->ne[3]); + + int ith = params->ith; + int nth = params->nth; + + int nrows = ggml_nrows(dst); + int npt = (nrows + nth - 1)/nth; + int first = npt*ith; + int last = MIN(first + npt, nrows); + + int ne0 = dst->ne[0]; + int ne00 = src->ne[0]; + + for (int ir = first; ir < last; ++ir) { + int ii = ir; + int i3 = ii/(src->ne[1]*src->ne[2]); ii -= i3*src->ne[1]*src->ne[2]; + int i2 = ii/src->ne[1]; ii -= i2*src->ne[1]; + int i1 = ii; + + int32_t * y = (int32_t *)((char *)dst->data + i1*dst->nb[1] + i2*dst->nb[2] + i3*dst->nb[3]); + //memset(y, 0, ne0*sizeof(int)); + for (int j = 0; j < ne0; ++j) y[j] = -1; + + if (src->type == GGML_TYPE_F16) { + const uint16_t * x = (const uint16_t *)((const char *)src->data + i1*src->nb[1] + i2*src->nb[2] + i3*src->nb[3]); + int idx = 0; + for (int j = 0; j < ne00; ++j) { + assert(idx < ne0); + if (x[j] == 0) y[idx++] = j; + } + } else { + const float * x = (const float *)((const char *)src->data + i1*src->nb[1] + i2*src->nb[2] + i3*src->nb[3]); + int idx = 0; + for (int j = 0; j < ne00; ++j) { + assert(idx < ne0); + if (x[j] == 0.0f) y[idx++] = j; + } + } + } +} + + // ggml_compute_forward_win_part static void ggml_compute_forward_win_part_f32( @@ -23507,6 +24116,10 @@ static void ggml_compute_forward_unary( { ggml_compute_forward_softplus(params, dst); } break; + case GGML_UNARY_OP_SQRT_SOFTPLUS: + { + ggml_compute_forward_sqrt_softplus(params, dst); + } break; default: { GGML_ABORT("fatal error"); @@ -25048,6 +25661,18 @@ static int ggml_compute_forward(struct ggml_compute_params * params, struct ggml { ggml_compute_forward_sinkhorn(params, tensor); } break; + case GGML_OP_HC_PRE: + { + ggml_compute_forward_hc_pre(params, tensor); + } break; + case GGML_OP_HC_POST: + { + ggml_compute_forward_hc_post(params, tensor); + } break; + case GGML_OP_MASK_TO_IDX: + { + ggml_compute_forward_mask_to_idx(params, tensor); + } break; case GGML_OP_INDEXER_TOPK: { if (!iqk_indexer_topk(tensor, params->wdata, (barrier_t)ggml_barrier, (void *)params->shared, params->ith, params->nth)) { @@ -26122,6 +26747,9 @@ static void ggml_compute_backward(struct ggml_context * ctx, struct ggml_tensor case GGML_OP_INDEXER_TOPK: case GGML_OP_MASK_TOPK: case GGML_OP_SINKHORN: + case GGML_OP_HC_PRE: + case GGML_OP_HC_POST: + case GGML_OP_MASK_TO_IDX: { GGML_ABORT("fatal error"); // TODO: not implemented } @@ -26257,6 +26885,10 @@ static void ggml_compute_backward(struct ggml_context * ctx, struct ggml_tensor { GGML_ABORT("fatal error"); // TODO: not implemented } + case GGML_UNARY_OP_SQRT_SOFTPLUS: + { + GGML_ABORT("fatal error"); // TODO: not implemented + } case GGML_UNARY_OP_SILU: { // necessary for llama @@ -26809,6 +27441,7 @@ static int ggml_get_n_tasks(struct ggml_tensor * node, int n_threads) { case GGML_UNARY_OP_SILU: case GGML_UNARY_OP_EXP: case GGML_UNARY_OP_SOFTPLUS: + case GGML_UNARY_OP_SQRT_SOFTPLUS: case GGML_UNARY_OP_SWIGLU: case GGML_UNARY_OP_SWIGLU_OAI: { @@ -26861,6 +27494,9 @@ static int ggml_get_n_tasks(struct ggml_tensor * node, int n_threads) { case GGML_OP_INDEXER_TOPK: case GGML_OP_MASK_TOPK: case GGML_OP_SINKHORN: + case GGML_OP_HC_PRE: + case GGML_OP_HC_POST: + case GGML_OP_MASK_TO_IDX: { n_tasks = n_threads; } break; diff --git a/ggml/src/iqk/fa/iqk_fa_128_128.cpp b/ggml/src/iqk/fa/iqk_fa_128_128.cpp index e5db1720..6400a43e 100644 --- a/ggml/src/iqk/fa/iqk_fa_128_128.cpp +++ b/ggml/src/iqk/fa/iqk_fa_128_128.cpp @@ -19,26 +19,26 @@ IQK_FA_CASE(iqk_fa_128_128) { if (type_v != GGML_TYPE_BF16) return false; // we do not support mixing bf16 k-cache with other types if (nk%64 == 0) { iqk_flash_helper_T<128, 128, 64>(nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } iqk_flash_helper_T<128, 128, 32>(nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } #endif if (nk%128 == 0) { return iqk_flash_helper_T<128, 128, 128>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } if (nk%64 == 0) { return iqk_flash_helper_T<128, 128, 64>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } return iqk_flash_helper_T<128, 128, 32>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } diff --git a/ggml/src/iqk/fa/iqk_fa_192_128.cpp b/ggml/src/iqk/fa/iqk_fa_192_128.cpp index 9cd62afd..48105792 100644 --- a/ggml/src/iqk/fa/iqk_fa_192_128.cpp +++ b/ggml/src/iqk/fa/iqk_fa_192_128.cpp @@ -19,26 +19,26 @@ IQK_FA_CASE(iqk_fa_192_128) { if (type_v != GGML_TYPE_BF16) return false; // we do not support mixing bf16 k-cache with other types if (nk%64 == 0) { iqk_flash_helper_T<192, 128, 64>(nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } iqk_flash_helper_T<192, 128, 32>(nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } #endif if (nk%128 == 0) { return iqk_flash_helper_T<192, 128, 128>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } if (nk%64 == 0) { return iqk_flash_helper_T<192, 128, 64>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } return iqk_flash_helper_T<192, 128, 32>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } diff --git a/ggml/src/iqk/fa/iqk_fa_192_192.cpp b/ggml/src/iqk/fa/iqk_fa_192_192.cpp index 21fe033c..024290cf 100644 --- a/ggml/src/iqk/fa/iqk_fa_192_192.cpp +++ b/ggml/src/iqk/fa/iqk_fa_192_192.cpp @@ -19,26 +19,26 @@ IQK_FA_CASE(iqk_fa_192_192) { if (type_v != GGML_TYPE_BF16) return false; // we do not support mixing bf16 k-cache with other types if (nk%64 == 0) { iqk_flash_helper_T<192, 192, 64>(nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } iqk_flash_helper_T<192, 192, 32>(nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } #endif if (nk%128 == 0) { return iqk_flash_helper_T<192, 192, 128>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } if (nk%64 == 0) { return iqk_flash_helper_T<192, 192, 64>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } return iqk_flash_helper_T<192, 192, 32>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } diff --git a/ggml/src/iqk/fa/iqk_fa_256_256.cpp b/ggml/src/iqk/fa/iqk_fa_256_256.cpp index a8565de2..46535ec0 100644 --- a/ggml/src/iqk/fa/iqk_fa_256_256.cpp +++ b/ggml/src/iqk/fa/iqk_fa_256_256.cpp @@ -19,26 +19,26 @@ IQK_FA_CASE(iqk_fa_256_256) { if (type_v != GGML_TYPE_BF16) return false; // we do not support mixing bf16 k-cache with other types if (nk%64 == 0) { iqk_flash_helper_T<256, 256, 64>(nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } iqk_flash_helper_T<256, 256, 32>(nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } #endif if (nk%128 == 0) { return iqk_flash_helper_T<256, 256, 128>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } if (nk%64 == 0) { return iqk_flash_helper_T<256, 256, 64>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } return iqk_flash_helper_T<256, 256, 32>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } diff --git a/ggml/src/iqk/fa/iqk_fa_320_256.cpp b/ggml/src/iqk/fa/iqk_fa_320_256.cpp index 03a9456d..f12ff099 100644 --- a/ggml/src/iqk/fa/iqk_fa_320_256.cpp +++ b/ggml/src/iqk/fa/iqk_fa_320_256.cpp @@ -10,43 +10,44 @@ template inline void iqk_deepseek_helper(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q, int stride_m, int stride_qkv, const float * q, const char * mask, float scale, float softcap, float * qkv, - const float * sinkf, float * M, float * S) { - auto update = [&nq1, &mask, &q, &qkv, &M, &S, stride_q, stride_m, stride_qkv] (int n) { + const float * sinkf, int sink_stride, float * M, float * S) { + auto update = [&nq1, &mask, &q, &qkv, &M, &S, stride_q, stride_m, stride_qkv, &sinkf, sink_stride] (int n) { nq1 -= n; if (nq1 == 0) return true; q += n*stride_q; mask += n*stride_m; qkv += n*stride_qkv; if (M && S) { M += n; S += n; } + if (sinkf) sinkf += n*sink_stride; return false; }; if (nq1 >= 16) { int n_step = nq1/16; - FlashAttn<320, 256, 16, step_k> fa(scale, softcap, sinkf); + FlashAttn<320, 256, 16, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 16*n_step, nk1, stride_q, stride_m, stride_qkv, q, mask, qkv, M, S); if (update(16*n_step)) return; } if (nq1 >= 8) { int n_step = nq1/8; - FlashAttn<320, 256, 8, step_k> fa(scale, softcap, sinkf); + FlashAttn<320, 256, 8, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 8*n_step, nk1, stride_q, stride_m, stride_qkv, q, mask, qkv, M, S); if (update(8*n_step)) return; } if (nq1 >= 4) { int n_step = nq1/4; - FlashAttn<320, 256, 4, step_k> fa(scale, softcap, sinkf); + FlashAttn<320, 256, 4, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 4*n_step, nk1, stride_q, stride_m, stride_qkv, q, mask, qkv, M, S); if (update(4*n_step)) return; } if (nq1 == 3) { - FlashAttn<320, 256, 3, step_k> fa(scale, softcap, sinkf); + FlashAttn<320, 256, 3, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 3, nk1, stride_q, stride_m, stride_qkv, q, mask, qkv, M, S); } else if (nq1 == 2) { - FlashAttn<320, 256, 2, step_k> fa(scale, softcap, sinkf); + FlashAttn<320, 256, 2, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 2, nk1, stride_q, stride_m, stride_qkv, q, mask, qkv, M, S); } else { - FlashAttn<320, 256, 1, step_k> fa(scale, softcap, sinkf); + FlashAttn<320, 256, 1, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 1, nk1, stride_q, stride_m, stride_qkv, q, mask, qkv, M, S); } } @@ -55,55 +56,55 @@ template inline bool iqk_deepseek_helper(ggml_type type_k, int nq1, int nk1, int stride_q, int stride_k, int stride_v, int stride_m, int stride_qkv, const float * q, const char * k, const char * v, const char * mask, - float scale, float softcap, float * qkv, const float * sinkf, float * M, float * S) { + float scale, float softcap, float * qkv, const float * sinkf, int sink_stride, float * M, float * S) { if (type_k == GGML_TYPE_Q8_0) { HelperQ80 kh((const char *)k, stride_k); HelperQ80 vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } if (type_k == GGML_TYPE_Q8_0_R8) { HelperQ80R8<320> kh((const char *)k, stride_k); HelperQ80 vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } if (type_k == GGML_TYPE_Q6_0) { HelperQ60 kh((const char *)k, stride_k); HelperQ60 vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } #if GGML_IQK_FA_ALL_QUANTS if (type_k == GGML_TYPE_Q8_KV) { HelperQ8KV<320> kh((const char *)k, stride_k); HelperQ8KV<256> vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } if (type_k == GGML_TYPE_Q4_0) { HelperQ40 kh((const char *)k, stride_k); HelperQ40 vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } if (type_k == GGML_TYPE_Q4_1) { HelperQ41 kh((const char *)k, stride_k); HelperQ41 vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } if (type_k == GGML_TYPE_IQ4_NL) { HelperIQ4nl kh((const char *)k, stride_k); HelperIQ4nl vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } #endif if (type_k == GGML_TYPE_F16) { HelperF16 kh((const char *)k, stride_k); HelperF16 vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } #ifdef __AVX512BF16__ @@ -111,10 +112,10 @@ inline bool iqk_deepseek_helper(ggml_type type_k, HelperBF16<320, step_k> kh((const char *)k, stride_k); HelperBF16<256, step_k> vh((const char *)v, stride_v); if (nq1 % 8 == 0) { - FlashAttnBF16<320, 256, 8, step_k> fa(scale, softcap, sinkf); + FlashAttnBF16<320, 256, 8, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); } else { - FlashAttnBF16<320, 256, 1, step_k> fa(scale, softcap, sinkf); + FlashAttnBF16<320, 256, 1, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); } return true; @@ -135,7 +136,7 @@ IQK_FA_CASE(iqk_fa_320_256) { } stride_q /= sizeof(float); // q stride as float return iqk_deepseek_helper<32>(type_k, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, (const char *)k, (const char *)v, (const char *)mask, scale, softcap, qkv, sinkf, M, S); + q, (const char *)k, (const char *)v, (const char *)mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } diff --git a/ggml/src/iqk/fa/iqk_fa_512_512.cpp b/ggml/src/iqk/fa/iqk_fa_512_512.cpp index 862a3a93..4a7c06a0 100644 --- a/ggml/src/iqk/fa/iqk_fa_512_512.cpp +++ b/ggml/src/iqk/fa/iqk_fa_512_512.cpp @@ -19,26 +19,26 @@ IQK_FA_CASE(iqk_fa_512_512) { if (type_v != GGML_TYPE_BF16) return false; // we do not support mixing bf16 k-cache with other types if (nk%64 == 0) { iqk_flash_helper_T<512, 512, 64>(nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } iqk_flash_helper_T<512, 512, 32>(nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } #endif if (nk%128 == 0) { return iqk_flash_helper_T<512, 512, 128>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } if (nk%64 == 0) { return iqk_flash_helper_T<512, 512, 64>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } return iqk_flash_helper_T<512, 512, 32>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } diff --git a/ggml/src/iqk/fa/iqk_fa_576_512.cpp b/ggml/src/iqk/fa/iqk_fa_576_512.cpp index 132d4350..3aa6cf61 100644 --- a/ggml/src/iqk/fa/iqk_fa_576_512.cpp +++ b/ggml/src/iqk/fa/iqk_fa_576_512.cpp @@ -10,43 +10,44 @@ template inline void iqk_deepseek_helper(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q, int stride_m, int stride_qkv, const float * q, const char * mask, float scale, float softcap, float * qkv, - const float * sinkf, float * M, float * S) { - auto update = [&nq1, &mask, &q, &qkv, &M, &S, stride_q, stride_m, stride_qkv] (int n) { + const float * sinkf, int sink_stride, float * M, float * S) { + auto update = [&nq1, &mask, &q, &qkv, &M, &S, stride_q, stride_m, stride_qkv, &sinkf, sink_stride] (int n) { nq1 -= n; if (nq1 == 0) return true; q += n*stride_q; mask += n*stride_m; qkv += n*stride_qkv; if (M && S) { M += n; S += n; } + if (sinkf) sinkf += n*sink_stride; return false; }; if (nq1 >= 16) { int n_step = nq1/16; - FlashAttn<576, 512, 16, step_k> fa(scale, softcap, sinkf); + FlashAttn<576, 512, 16, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 16*n_step, nk1, stride_q, stride_m, stride_qkv, q, mask, qkv, M, S); if (update(16*n_step)) return; } if (nq1 >= 8) { int n_step = nq1/8; - FlashAttn<576, 512, 8, step_k> fa(scale, softcap, sinkf); + FlashAttn<576, 512, 8, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 8*n_step, nk1, stride_q, stride_m, stride_qkv, q, mask, qkv, M, S); if (update(8*n_step)) return; } if (nq1 >= 4) { int n_step = nq1/4; - FlashAttn<576, 512, 4, step_k> fa(scale, softcap, sinkf); + FlashAttn<576, 512, 4, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 4*n_step, nk1, stride_q, stride_m, stride_qkv, q, mask, qkv, M, S); if (update(4*n_step)) return; } if (nq1 == 3) { - FlashAttn<576, 512, 3, step_k> fa(scale, softcap, sinkf); + FlashAttn<576, 512, 3, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 3, nk1, stride_q, stride_m, stride_qkv, q, mask, qkv, M, S); } else if (nq1 == 2) { - FlashAttn<576, 512, 2, step_k> fa(scale, softcap, sinkf); + FlashAttn<576, 512, 2, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 2, nk1, stride_q, stride_m, stride_qkv, q, mask, qkv, M, S); } else { - FlashAttn<576, 512, 1, step_k> fa(scale, softcap, sinkf); + FlashAttn<576, 512, 1, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 1, nk1, stride_q, stride_m, stride_qkv, q, mask, qkv, M, S); } } @@ -55,55 +56,55 @@ template inline bool iqk_deepseek_helper(ggml_type type_k, int nq1, int nk1, int stride_q, int stride_k, int stride_v, int stride_m, int stride_qkv, const float * q, const char * k, const char * v, const char * mask, - float scale, float softcap, float * qkv, const float * sinkf, float * M, float * S) { + float scale, float softcap, float * qkv, const float * sinkf, int sink_stride, float * M, float * S) { if (type_k == GGML_TYPE_Q8_0) { HelperQ80 kh((const char *)k, stride_k); HelperQ80 vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } if (type_k == GGML_TYPE_Q8_0_R8) { HelperQ80R8<576> kh((const char *)k, stride_k); HelperQ80 vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } if (type_k == GGML_TYPE_Q6_0) { HelperQ60 kh((const char *)k, stride_k); HelperQ60 vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } #if GGML_IQK_FA_ALL_QUANTS if (type_k == GGML_TYPE_Q8_KV) { HelperQ8KV<576> kh((const char *)k, stride_k); HelperQ8KV<512> vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } if (type_k == GGML_TYPE_Q4_0) { HelperQ40 kh((const char *)k, stride_k); HelperQ40 vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } if (type_k == GGML_TYPE_Q4_1) { HelperQ41 kh((const char *)k, stride_k); HelperQ41 vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } if (type_k == GGML_TYPE_IQ4_NL) { HelperIQ4nl kh((const char *)k, stride_k); HelperIQ4nl vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } #endif if (type_k == GGML_TYPE_F16) { HelperF16 kh((const char *)k, stride_k); HelperF16 vh((const char *)v, stride_v); - iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_deepseek_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } #ifdef __AVX512BF16__ @@ -111,10 +112,10 @@ inline bool iqk_deepseek_helper(ggml_type type_k, HelperBF16<576, step_k> kh((const char *)k, stride_k); HelperBF16<512, step_k> vh((const char *)v, stride_v); if (nq1 % 8 == 0) { - FlashAttnBF16<576, 512, 8, step_k> fa(scale, softcap, sinkf); + FlashAttnBF16<576, 512, 8, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); } else { - FlashAttnBF16<576, 512, 1, step_k> fa(scale, softcap, sinkf); + FlashAttnBF16<576, 512, 1, step_k> fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); } return true; @@ -135,7 +136,7 @@ IQK_FA_CASE(iqk_fa_576_512) { } stride_q /= sizeof(float); // q stride as float return iqk_deepseek_helper<32>(type_k, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, (const char *)k, (const char *)v, (const char *)mask, scale, softcap, qkv, sinkf, M, S); + q, (const char *)k, (const char *)v, (const char *)mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } diff --git a/ggml/src/iqk/fa/iqk_fa_64_64.cpp b/ggml/src/iqk/fa/iqk_fa_64_64.cpp index 84b9bad0..0f1dbaf4 100644 --- a/ggml/src/iqk/fa/iqk_fa_64_64.cpp +++ b/ggml/src/iqk/fa/iqk_fa_64_64.cpp @@ -19,26 +19,26 @@ IQK_FA_CASE(iqk_fa_64_64) { if (type_v != GGML_TYPE_BF16) return false; // we do not support mixing bf16 k-cache with other types if (nk%64 == 0) { iqk_flash_helper_T<64, 64, 64>(nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } iqk_flash_helper_T<64, 64, 32>(nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } #endif if (nk%128 == 0) { return iqk_flash_helper_T<64, 64, 128>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } if (nk%64 == 0) { return iqk_flash_helper_T<64, 64, 64>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } return iqk_flash_helper_T<64, 64, 32>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } diff --git a/ggml/src/iqk/fa/iqk_fa_96_96.cpp b/ggml/src/iqk/fa/iqk_fa_96_96.cpp index 44544f8b..345f2905 100644 --- a/ggml/src/iqk/fa/iqk_fa_96_96.cpp +++ b/ggml/src/iqk/fa/iqk_fa_96_96.cpp @@ -19,26 +19,26 @@ IQK_FA_CASE(iqk_fa_96_96) { if (type_v != GGML_TYPE_BF16) return false; // we do not support mixing bf16 k-cache with other types if (nk%64 == 0) { iqk_flash_helper_T<96, 96, 64>(nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } iqk_flash_helper_T<96, 96, 32>(nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); return true; } #endif if (nk%128 == 0) { return iqk_flash_helper_T<96, 96, 128>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } if (nk%64 == 0) { return iqk_flash_helper_T<96, 96, 64>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } return iqk_flash_helper_T<96, 96, 32>(type_k, type_v, nq, nk, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, ck, cv, cm, scale, softcap, qkv, sinkf, M, S); + q, ck, cv, cm, scale, softcap, qkv, sinkf, sink_stride, M, S); } diff --git a/ggml/src/iqk/fa/iqk_fa_templates.h b/ggml/src/iqk/fa/iqk_fa_templates.h index f1776804..68dbbec3 100644 --- a/ggml/src/iqk/fa/iqk_fa_templates.h +++ b/ggml/src/iqk/fa/iqk_fa_templates.h @@ -1156,11 +1156,11 @@ struct FlashQKV { } template - inline void normalize_and_store_1row(const FMS& fms, int j, qkv_cache_t * R, float * qkv, const float * sinkf) const { + inline void normalize_and_store_1row(const FMS& fms, int j, qkv_cache_t * R, float * qkv, const float * sinkf, int sink_stride) const { static_assert(q_step == FMS::q_step); float S = fms.S[j]; if (sinkf) { - float s = *sinkf; + float s = sinkf[j*sink_stride]; if (s > fms.M[j]) { float m = expf(fms.M[j] - s); auto vm = F16::set1(m); @@ -1187,7 +1187,7 @@ struct FlashQKV { } template - inline void normalize_and_store(const FMS& fms, int nq1, int stride_qkv, float * qkv, const float * sinkf, float * M, float * S) { + inline void normalize_and_store(const FMS& fms, int nq1, int stride_qkv, float * qkv, const float * sinkf, int sink_stride, float * M, float * S) { static_assert(q_step == FMS::q_step); if (M && S) { std::memcpy(M, fms.M, nq1*sizeof(float)); @@ -1207,7 +1207,7 @@ struct FlashQKV { } else { auto R = qkv_cache; for (int j = 0; j < nq1; ++j) { - normalize_and_store_1row(fms, j, R, qkv, sinkf); + normalize_and_store_1row(fms, j, R, qkv, sinkf, sink_stride); qkv += stride_qkv; R += D; } @@ -1215,7 +1215,7 @@ struct FlashQKV { } template - inline void normalize_and_store(const FMS& fms, int stride_qkv, float * qkv, const float * sinkf, float * M, float * S) { + inline void normalize_and_store(const FMS& fms, int stride_qkv, float * qkv, const float * sinkf, int sink_stride, float * M, float * S) { static_assert(q_step == FMS::q_step); if (M && S) { std::memcpy(M, fms.M, q_step*sizeof(float)); @@ -1235,7 +1235,7 @@ struct FlashQKV { } else { auto R = qkv_cache; for (int j = 0; j < q_step; ++j) { - normalize_and_store_1row(fms, j, R, qkv, sinkf); + normalize_and_store_1row(fms, j, R, qkv, sinkf, sink_stride); qkv += stride_qkv; R += D; } @@ -1369,7 +1369,7 @@ void compute_helper(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q, in FlashMS& fms, FlashQKV& fqkv, const float * q, const char * mask, float * qkv, - const float * sinkf, float * M, float * S) { + const float * sinkf, int sink_stride, float * M, float * S) { #ifdef __aarch64__ float16_t q_f16[Dk*q_step]; #endif @@ -1394,12 +1394,13 @@ void compute_helper(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q, in vh.next_block(k_step); mr += k_step*sizeof(ggml_half); } - fqkv.normalize_and_store(fms, stride_qkv, qkv, sinkf, M, S); + fqkv.normalize_and_store(fms, stride_qkv, qkv, sinkf, sink_stride, M, S); q += q_step*stride_q; mask += q_step*stride_m; qkv += q_step*stride_qkv; if (M && S) { M += q_step; S += q_step; } + if (sinkf) sinkf += q_step*sink_stride; } int n_left = nq1 - q_step*(nq1/q_step); if (n_left > 0) { @@ -1421,7 +1422,7 @@ void compute_helper(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q, in vh.next_block(k_step); mr += k_step*sizeof(ggml_half); } - fqkv.normalize_and_store(fms, n_left, stride_qkv, qkv, sinkf, M, S); + fqkv.normalize_and_store(fms, n_left, stride_qkv, qkv, sinkf, sink_stride, M, S); } } @@ -1430,7 +1431,7 @@ void compute_helper_q(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q, FlashMS& fms, FlashQKV& fqkv, const float * q, const char * mask, float * qkv, - const float * sinkf, float * M, float * S, char * qptr) { + const float * sinkf, int sink_stride, float * M, float * S, char * qptr) { auto q8 = (typename KHelper::block_q8 *)qptr; // This optimization fails under certain conditions (see https://github.com/ikawrakow/ik_llama.cpp/issues/1205) // => disabling until I figure out what goes wrong @@ -1453,7 +1454,7 @@ void compute_helper_q(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q, vh.next_block(k_step); mr += k_step*sizeof(ggml_half); } - fqkv.normalize_and_store(fms, stride_qkv, qkv, sinkf, M, S); + fqkv.normalize_and_store(fms, stride_qkv, qkv, sinkf, sink_stride, M, S); return; } } @@ -1491,16 +1492,17 @@ void compute_helper_q(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q, } #if FA_TIMING t1 = Perf::cur_time(); - fqkv.normalize_and_store(fms, stride_qkv, qkv, sinkf, M, S); + fqkv.normalize_and_store(fms, stride_qkv, qkv, sinkf, sink_stride, M, S); perf.accum_nolock(3, t1); #else - fqkv.normalize_and_store(fms, stride_qkv, qkv, sinkf, M, S); + fqkv.normalize_and_store(fms, stride_qkv, qkv, sinkf, sink_stride, M, S); #endif q += q_step*stride_q; mask += q_step*stride_m; qkv += q_step*stride_qkv; if (M && S) { M += q_step; S += q_step; } + if (sinkf) sinkf += q_step*sink_stride; } int n_left = nq1 - q_step*(nq1/q_step); if (n_left > 0) { @@ -1516,7 +1518,7 @@ void compute_helper_q(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q, vh.next_block(k_step); mr += k_step*sizeof(ggml_half); } - fqkv.normalize_and_store(fms, n_left, stride_qkv, qkv, sinkf, M, S); + fqkv.normalize_and_store(fms, n_left, stride_qkv, qkv, sinkf, sink_stride, M, S); } #if FA_TIMING Perf::instance().add(perf); @@ -1546,7 +1548,8 @@ struct FlashAttn { static_assert(k_step%F16::block_size == 0); static_assert(q_step <= 4 || q_step%4 == 0); - FlashAttn(float scale, float softcap, const float * sinkf) : fms(scale, softcap), sinkf(sinkf) {} + FlashAttn(float scale, float softcap, const float * sinkf, int sink_stride) : fms(scale, softcap), + sinkf(sinkf), sink_stride(sink_stride) {} template void compute(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q, int stride_m, int stride_qkv, @@ -1575,7 +1578,7 @@ struct FlashAttn { HelperQ80R8 khr4(nk1, kh); #endif compute_helper_q, VHelper, FlashQKfp32>( - khr4, vh, nq1, nk1, stride_q, stride_m, stride_qkv, fms, fqkv, q, mask, qkv, sinkf, M, S, qptr); + khr4, vh, nq1, nk1, stride_q, stride_m, stride_qkv, fms, fqkv, q, mask, qkv, sinkf, sink_stride, M, S, qptr); return; } @@ -1589,30 +1592,31 @@ struct FlashAttn { HelperQ8KVR8 khr4(nk1, kh); #endif compute_helper_q, VHelper, FlashQKfp32>( - khr4, vh, nq1, nk1, stride_q, stride_m, stride_qkv, fms, fqkv, q, mask, qkv, sinkf, M, S, qptr); + khr4, vh, nq1, nk1, stride_q, stride_m, stride_qkv, fms, fqkv, q, mask, qkv, sinkf, sink_stride, M, S, qptr); return; } #endif } compute_helper_q>( - kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, fms, fqkv, q, mask, qkv, sinkf, M, S, qptr); + kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, fms, fqkv, q, mask, qkv, sinkf, sink_stride, M, S, qptr); } else { typename KHelper::block_q8 q8[q_step*(Dk/KHelper::block_size_q)]; compute_helper_q>( - kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, fms, fqkv, q, mask, qkv, sinkf, M, S, (char *)q8); + kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, fms, fqkv, q, mask, qkv, sinkf, sink_stride, M, S, (char *)q8); } } else { compute_helper>( - kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, fms, fqkv, q, mask, qkv, sinkf, M, S); + kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, fms, fqkv, q, mask, qkv, sinkf, sink_stride, M, S); } } FlashMS fms; FlashQKV fqkv; const float * sinkf; + int sink_stride; }; @@ -1970,7 +1974,7 @@ struct FlashAttnBF16 { static_assert(k_step%32 == 0); static_assert(q_step <= 4 || q_step%4 == 0); - FlashAttnBF16(float scale, float softcap, const float * sinkf) : fms(scale, softcap), sinkf(sinkf) {} + FlashAttnBF16(float scale, float softcap, const float * sinkf, int sink_stride) : fms(scale, softcap), sinkf(sinkf), sink_stride(sink_stride) {} template void compute(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q, int stride_m, int stride_qkv, @@ -2011,7 +2015,7 @@ struct FlashAttnBF16 { #if FA_TIMING t1 = Perf::cur_time(); #endif - fqkv.normalize_and_store(fms, stride_qkv, qkv, sinkf, M, S); + fqkv.normalize_and_store(fms, stride_qkv, qkv, sinkf, sink_stride, M, S); #if FA_TIMING perf.accum_nolock(4, t1); #endif @@ -2020,6 +2024,7 @@ struct FlashAttnBF16 { mask += q_step*stride_m; qkv += q_step*stride_qkv; if (M && S) { M += q_step; S += q_step; } + if (sinkf) sinkf += q_step*sink_stride; } int n_left = nq1 - q_step*(nq1/q_step); if (n_left > 0) { @@ -2035,7 +2040,7 @@ struct FlashAttnBF16 { vh.next_block(k_step); mr += k_step*sizeof(ggml_half); } - fqkv.normalize_and_store(fms, n_left, stride_qkv, qkv, sinkf, M, S); + fqkv.normalize_and_store(fms, n_left, stride_qkv, qkv, sinkf, sink_stride, M, S); } #if FA_TIMING Perf::instance().add(perf); @@ -2045,74 +2050,76 @@ struct FlashAttnBF16 { FlashMS fms; FlashQKV fqkv; const float * sinkf; + int sink_stride; }; #endif template inline void iqk_flash_helper(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q, int stride_m, int stride_qkv, const float * q, const char * mask, float scale, float softcap, float * qkv, - const float * sinkf, float * M, float * S) { + const float * sinkf, int sink_stride, float * M, float * S) { - auto update = [&nq1, &mask, &q, &qkv, &M, &S, stride_q, stride_m, stride_qkv] (int n) { + auto update = [&nq1, &mask, &q, &qkv, &M, &S, stride_q, stride_m, stride_qkv, &sinkf, sink_stride] (int n) { nq1 -= n; if (nq1 == 0) return true; q += n*stride_q; mask += n*stride_m; qkv += n*stride_qkv; if (M && S) { M += n; S += n; } + if (sinkf) sinkf += n*sink_stride; return false; }; if (nk1 >= 512) { if (nq1 >= 128) { int n_step = nq1/128; - FlashAttn fa(scale, softcap, sinkf); + FlashAttn fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 128*n_step, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); if (update(128*n_step)) return; } if (nq1 >= 64) { int n_step = nq1/64; - FlashAttn fa(scale, softcap, sinkf); + FlashAttn fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 64*n_step, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); if (update(64*n_step)) return; } if (nq1 >= 32) { int n_step = nq1/32; - FlashAttn fa(scale, softcap, sinkf); + FlashAttn fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 32*n_step, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); if (update(32*n_step)) return; } if (nq1 >= 16) { int n_step = nq1/16; - FlashAttn fa(scale, softcap, sinkf); + FlashAttn fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 16*n_step, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); if (update(16*n_step)) return; } } if (nq1 == 12) { // Special case: TG for GLM-4.5/4.6 - FlashAttn fa(scale, softcap, sinkf); + FlashAttn fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 12, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); return; } if (nq1 >= 8) { int n_step = nq1/8; - FlashAttn fa(scale, softcap, sinkf); + FlashAttn fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 8*n_step, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); if (update(8*n_step)) return; } else if (nq1 >= 4) { int n_step = nq1/4; - FlashAttn fa(scale, softcap, sinkf); + FlashAttn fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 4*n_step, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); if (update(4*n_step)) return; } else if (nq1 >= 2) { int n_step = nq1/2; - FlashAttn fa(scale, softcap, sinkf); + FlashAttn fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, 2*n_step, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); if (update(2*n_step)) return; } - FlashAttn fa(scale, softcap, sinkf); + FlashAttn fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); } @@ -2120,26 +2127,26 @@ inline void iqk_flash_helper(KHelper& kh, VHelper& vh, int nq1, int nk1, int str template inline void iqk_flash_helper_T(int nq1, int nk1, int stride_q, int stride_k, int stride_v, int stride_m, int stride_qkv, const float * q, const char * k, const char * v, const char * mask, - float scale, float softcap, float * qkv, const float * sinkf, float * M, float * S) { + float scale, float softcap, float * qkv, const float * sinkf, int sink_stride, float * M, float * S) { HelperBF16 kh(k, stride_k); HelperBF16 vh(v, stride_v); if (nk1 >= 4096) { if (nq1 >= 64) { - FlashAttnBF16 fa(scale, softcap, sinkf); + FlashAttnBF16 fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); return; } else if (nq1 >= 16) { - FlashAttnBF16 fa(scale, softcap, sinkf); + FlashAttnBF16 fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); return; } } if (nq1 >= 8) { - FlashAttnBF16 fa(scale, softcap, sinkf); + FlashAttnBF16 fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); } else { - FlashAttnBF16 fa(scale, softcap, sinkf); + FlashAttnBF16 fa(scale, softcap, sinkf, sink_stride); fa.compute(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, (const char *)mask, qkv, M, S); } } @@ -2149,43 +2156,43 @@ template inline bool iqk_flash_helper_T(KHelper& kh, ggml_type type_v, int nq1, int nk1, int stride_q, int stride_v, int stride_m, int stride_qkv, const float * q, const char * v, const char * mask, - float scale, float softcap, float * qkv, const float * sinkf, float * M, float * S) { + float scale, float softcap, float * qkv, const float * sinkf, int sink_stride, float * M, float * S) { switch (type_v) { case GGML_TYPE_F16: { HelperF16 vh(v, stride_v); - iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; #ifdef __AVX512BF16__ case GGML_TYPE_BF16: { HelperBF16 vh(v, stride_v); - iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; #endif case GGML_TYPE_Q8_0: { HelperQ80 vh(v, stride_v); - iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; case GGML_TYPE_Q8_KV: { HelperQ8KV vh(v, stride_v); - iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; case GGML_TYPE_Q6_0: { HelperQ60 vh(v, stride_v); - iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; #if GGML_IQK_FA_ALL_QUANTS case GGML_TYPE_Q4_0: { HelperQ40 vh(v, stride_v); - iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; case GGML_TYPE_Q4_1: { HelperQ41 vh(v, stride_v); - iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; case GGML_TYPE_IQ4_NL: { HelperIQ4nl vh(v, stride_v); - iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S); + iqk_flash_helper(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; #endif default: return false; @@ -2197,42 +2204,42 @@ template inline bool iqk_flash_helper_T(ggml_type type_k, ggml_type type_v, int nq1, int nk1, int stride_q, int stride_k, int stride_v, int stride_m, int stride_qkv, const float * q, const char * k, const char * v, const char * mask, - float scale, float softcap, float * qkv, const float * sinkf, float * M, float * S) { + float scale, float softcap, float * qkv, const float * sinkf, int sink_stride, float * M, float * S) { bool result = false; switch (type_k) { case GGML_TYPE_F16: { HelperF16 kh(k, stride_k); - result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, M, S); + result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; case GGML_TYPE_Q8_0: { HelperQ80 kh(k, stride_k); - result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, M, S); + result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; case GGML_TYPE_Q8_0_R8: { HelperQ80R8 kh(k, stride_k); - result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, M, S); + result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; case GGML_TYPE_Q6_0: { HelperQ60 kh(k, stride_k); - result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, M, S); + result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; #if GGML_IQK_FA_ALL_QUANTS case GGML_TYPE_Q8_KV: { HelperQ8KV kh(k, stride_k); - result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, M, S); + result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; case GGML_TYPE_Q4_0: { HelperQ40 kh(k, stride_k); - result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, M, S); + result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; case GGML_TYPE_Q4_1: { HelperQ41 kh(k, stride_k); - result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, M, S); + result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; case GGML_TYPE_IQ4_NL: { HelperIQ4nl kh(k, stride_k); - result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, M, S); + result = iqk_flash_helper_T(kh, type_v, nq1, nk1, stride_q, stride_v, stride_m, stride_qkv, q, v, mask, scale, softcap, qkv, sinkf, sink_stride, M, S); } break; #endif default: break; @@ -2247,7 +2254,7 @@ inline bool iqk_flash_helper_T(ggml_type type_k, ggml_type type_v, int stride_q, int stride_k, int stride_v, int stride_m, int stride_qkv,\ const float * q, const void * k, const void * v, const void * mask,\ float scale, float softcap,\ - float * qkv, const float * sinkf, float * M, float * S) + float * qkv, const float * sinkf, int sink_stride, float * M, float * S) IQK_FA_CASE(iqk_fa_576_512); IQK_FA_CASE(iqk_fa_512_512); diff --git a/ggml/src/iqk/iqk_cpu_ops.cpp b/ggml/src/iqk/iqk_cpu_ops.cpp index e50f51ea..8fd51f1c 100644 --- a/ggml/src/iqk/iqk_cpu_ops.cpp +++ b/ggml/src/iqk/iqk_cpu_ops.cpp @@ -1066,7 +1066,7 @@ inline void iqk_add_f16(int n, const ggml_half * x, ggml_half * y) { void iqk_mask_topk(struct ggml_tensor * dst, int ith, int nth) { auto mask = dst->src[0]; auto topk = dst->src[1]; - GGML_ASSERT(mask->ne[0] >= topk->ne[1]); + GGML_ASSERT(mask->ne[0] >= topk->ne[0]); GGML_ASSERT(mask->ne[1] >= topk->ne[1] && mask->ne[2] == topk->ne[2] && mask->ne[3] == topk->ne[3]); GGML_ASSERT(ggml_are_same_shape(mask, dst)); GGML_ASSERT(mask->type == GGML_TYPE_F16 || mask->type == GGML_TYPE_F32); diff --git a/ggml/src/iqk/iqk_flash_attn.cpp b/ggml/src/iqk/iqk_flash_attn.cpp index e1ea5177..20095d0f 100644 --- a/ggml/src/iqk/iqk_flash_attn.cpp +++ b/ggml/src/iqk/iqk_flash_attn.cpp @@ -52,12 +52,15 @@ size_t iqk_fa_work_buffer_size(const struct ggml_tensor * dst, int nth) { auto K = dst->src[1]; auto V = dst->src[2]; auto indexer = dst->src[5]; - if (indexer && indexer->type == GGML_TYPE_I32 && indexer->ne[0] < K->ne[1] && Q->ne[1] >= nth && + if (indexer && indexer->type == GGML_TYPE_I32 && indexer->ne[0] < K->ne[1] && Q->ne[3] == 1 && K->ne[3] == 1 && V->ne[3] == 1 && K->ne[2] == 1) { auto row_size_k = ggml_row_size(K->type, K->ne[0]); auto row_size_v = ggml_row_size(V->type, V->ne[0]); auto work_size = (row_size_k + row_size_v + 64) * indexer->ne[0]; - return work_size * nth; + size_t result = work_size * nth; + if (Q->ne[1]== 1) result += 512*sizeof(float); + return result; + //return work_size * nth; } int rk2 = Q->ne[2]/K->ne[2]; size_t size = 0; @@ -176,7 +179,66 @@ extern "C" IQK_API bool iqk_flash_attn_noalibi(int type_q, int type_mask, float if (type_q != 0 || type_mask != 1 || max_bias > 0) return false; if (indexer && indexer->type == GGML_TYPE_I32) { - if (indexer->ne[0] < nek1 && neq1 >= nth && neq3 == 1 && nek3 == 1 && nev3 == 1 && nek2 == 1) { + //if (indexer->ne[0] < nek1 && neq1 >= nth && neq3 == 1 && nek3 == 1 && nev3 == 1 && nek2 == 1) { + if (indexer->ne[0] < nek1 && neq3 == 1 && nek3 == 1 && nev3 == 1 && nek2 == 1) { + // Workbuffer: we need + // * indexer->ne[0] * sizeof(ggml_half) to extract the mask for a row + // * indexer->ne[0] * ggml_row_size(int_type_k_in, Dk) to extract the selected K cache entries + // * indexer->ne[0] * ggml_row_size(int_type_v, Dv) to extract the selected V cache entries + auto row_size_k = ggml_row_size(ggml_type(int_type_k_in), Dk); + auto row_size_v = ggml_row_size(ggml_type(int_type_v ), Dv); + auto work_size = (row_size_k + row_size_v + 64) * indexer->ne[0]; + ggml_fp16_t h_inf = ggml_fp32_to_fp16(-INFINITY); + int nkv = indexer->ne[0]; + if (neq1 == 1) { + GGML_ASSERT(neq2 <= 256); + int npt = (neq2 + nth - 1)/nth; + int ith_mid = nth; + int neq2_this_thread = npt; + int first = ith*npt; + if (npt*nth > neq2) { + ith_mid = neq2 - nth*(npt - 1); + if (ith >= ith_mid) { + --neq2_this_thread; + //if (neq2_this_thread < 1) return true; + first = ith_mid*npt + (ith - ith_mid)*neq2_this_thread; + } + } + auto idx = (const int *)indexer->data; + auto M = (const ggml_fp16_t *)mask; + auto work_k = (char *)work_buffer_in; + auto work_v = work_k + row_size_k*indexer->ne[0]; + auto work_m = (ggml_fp16_t *)(work_v + row_size_v*indexer->ne[0]) + indexer->ne[0]*ith; + int last_found = -1; + for (int j = 0; j < nkv; ++j) { + if (idx[j] >= 0) { + last_found = j; + work_m[j] = M[idx[j]]; + if (j % nth == ith) { + std::memcpy(work_k + row_size_k*j, ((const char *)k + idx[j]*stride_k), row_size_k); + std::memcpy(work_v + row_size_v*j, ((const char *)v + idx[j]*stride_v), row_size_v); + } + } else { + work_m[j] = h_inf; + if (j % nth == ith) { + std::memset(work_k + row_size_k*j, 0, row_size_k); + std::memset(work_v + row_size_v*j, 0, row_size_v); + } + } + } + barrier(barrier_data); + if (last_found < 0 || neq2_this_thread < 1) return true; + ++last_found; + int this_nkv = 32*((last_found + 31)/32); + auto this_q = (const char *)q + first*nbq2; + auto this_qkv = qkv + first*nb1/sizeof(float); + if (!iqk_flash_attn_impl(int_type_k_in, int_type_v, + Dk, Dv, neq2_this_thread, this_nkv, nbq2, row_size_k, row_size_v, 0, Dv, + (const float *)this_q, work_k, work_v, work_m, (const float *)sinks, 1, + scale, softcap, + this_qkv, nullptr, nullptr)) return false; + return true; + } int npt = (neq1 + nth - 1)/nth; int ith_mid = nth; int neq1_this_thread = npt; @@ -185,33 +247,37 @@ extern "C" IQK_API bool iqk_flash_attn_noalibi(int type_q, int type_mask, float ith_mid = neq1 - nth*(npt - 1); if (ith >= ith_mid) { --neq1_this_thread; + if (neq1_this_thread < 1) return true; first = ith_mid*npt + (ith - ith_mid)*neq1_this_thread; } } - // Workbuffer: we need - // * indexer->ne[0] * sizeof(ggml_half) to extract the mask for a row - // * indexer->ne[0] * ggml_row_size(int_type_k_in, Dk) to extract the selected K cache entries - // * indexer->ne[0] * ggml_row_size(int_type_v, Dv) to extract the selected V cache entries - auto row_size_k = ggml_row_size(ggml_type(int_type_k_in), Dk); - auto row_size_v = ggml_row_size(ggml_type(int_type_v ), Dv); - auto work_size = (row_size_k + row_size_v + 64) * indexer->ne[0]; auto work_k = (char *)work_buffer_in + ith*work_size; auto work_v = work_k + row_size_k*indexer->ne[0]; - auto work_m = (uint16_t *)(work_v + row_size_v*indexer->ne[0]); - int nkv = indexer->ne[0]; + auto work_m = (ggml_fp16_t *)(work_v + row_size_v*indexer->ne[0]); for (int iq = first; iq < first + neq1_this_thread; ++iq) { auto idx = (const int *)((const char *)indexer->data + iq*indexer->nb[1]); - auto M = (const uint16_t *)((const char *)mask + iq*stride_m); + auto M = (const ggml_fp16_t *)((const char *)mask + iq*stride_m); + int last_found = -1; for (int j = 0; j < nkv; ++j) { - std::memcpy(work_k + row_size_k*j, ((const char *)k + idx[j]*stride_k), row_size_k); - std::memcpy(work_v + row_size_v*j, ((const char *)v + idx[j]*stride_v), row_size_v); - work_m[j] = M[idx[j]]; + if (idx[j] >= 0) { + std::memcpy(work_k + row_size_k*j, ((const char *)k + idx[j]*stride_k), row_size_k); + std::memcpy(work_v + row_size_v*j, ((const char *)v + idx[j]*stride_v), row_size_v); + work_m[j] = M[idx[j]]; + last_found = j; + } else { + std::memset(work_k + row_size_k*j, 0, row_size_k); + std::memset(work_v + row_size_v*j, 0, row_size_v); + work_m[j] = h_inf; + } } + if (last_found < 0) continue; + ++last_found; + int this_nkv = 32*((last_found + 31)/32); auto this_q = (const char *)q + iq*stride_q; auto this_qkv = qkv + iq*ne1*nb1/sizeof(float); if (!iqk_flash_attn_impl(int_type_k_in, int_type_v, - Dk, Dv, neq2, nkv, nbq2, row_size_k, row_size_v, 0, Dv, - (const float *)this_q, work_k, work_v, work_m, nullptr, 0, + Dk, Dv, neq2, this_nkv, nbq2, row_size_k, row_size_v, 0, Dv, + (const float *)this_q, work_k, work_v, work_m, (const float *)sinks, 1, scale, softcap, this_qkv, nullptr, nullptr)) return false; } @@ -536,7 +602,7 @@ extern "C" IQK_API bool iqk_flash_attn_noalibi(int type_q, int type_mask, float (const float *)((const char *)q + iq2*nbq2 + iq3*nbq3 + iq1*stride_q), (const void *)((const char *)k + iq2/rk2*nbk2 + iq3/rk3*nbk3), (const void *)((const char *)v + iq2/rv2*nbv2 + iq3/rv3*nbv3), - mask ? (const void *)((const char *)mask + iq1*stride_m) : nullptr, sinksf, 1, + mask ? (const void *)((const char *)mask + iq1*stride_m) : nullptr, sinksf, 0, scale, softcap, (float *)((char *)qkv + (iq3*ne2*ne1 + iq2 + iq1*ne1)*nb1), nullptr, nullptr)) return false; } diff --git a/ggml/src/iqk/iqk_flash_impl.h b/ggml/src/iqk/iqk_flash_impl.h index 8db97731..211b337c 100644 --- a/ggml/src/iqk/iqk_flash_impl.h +++ b/ggml/src/iqk/iqk_flash_impl.h @@ -24,7 +24,7 @@ bool iqk_flash_attn_impl(int type_k, // type of k const void * v, // v matrix. Assumed to be fp16, nq x nk elements const void * mask, // mask. If not null, assumed to be fp16. nq x nk elements const float * sinksf, // attention sinks - int nsinks, // number of sinks + int sink_stride, // stride between sinks float scale, // scale applied before softmax float softcap, // if > 0, a "soft-cap" operation is applied before softmax float * qkv, // v*softmax(scale*(k*q)) diff --git a/ggml/src/iqk/iqk_mul_mat.cpp b/ggml/src/iqk/iqk_mul_mat.cpp index 5158c6c5..21f68e2d 100644 --- a/ggml/src/iqk/iqk_mul_mat.cpp +++ b/ggml/src/iqk/iqk_mul_mat.cpp @@ -1377,7 +1377,7 @@ bool iqk_flash_attn_impl(int int_type_k, // type of k const void * v, // v matrix. Assumed to be fp16, nq x nk elements const void * mask, // mask. If not null, assumed to be fp16. nq x nk elements const float * sinksf, // mask. If not null, assumed to be fp16. nq x nk elements - [[maybe_unused]] int nsinks, + int sink_stride, // stride between sinks (used if sinksf is not null) float scale, // scale applied before softmax float softcap, // if > 0, a "soft-cap" operation is applied before softmax float * qkv, // v*softmax(scale*(k*q)) @@ -1387,45 +1387,45 @@ bool iqk_flash_attn_impl(int int_type_k, // type of k if (Dk == 576 && Dv == 512) { return iqk_fa_576_512(int_type_k, int_type_v, nq1, nk1, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, k, v, mask, scale, softcap, qkv, sinksf, M, S); + q, k, v, mask, scale, softcap, qkv, sinksf, sink_stride, M, S); } if (Dk == 512 && Dv == 512) { return iqk_fa_512_512(int_type_k, int_type_v, nq1, nk1, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, k, v, mask, scale, softcap, qkv, sinksf, M, S); + q, k, v, mask, scale, softcap, qkv, sinksf, sink_stride, M, S); } if (Dk == 320 && Dv == 256) { return iqk_fa_320_256(int_type_k, int_type_v, nq1, nk1, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, k, v, mask, scale, softcap, qkv, sinksf, M, S); + q, k, v, mask, scale, softcap, qkv, sinksf, sink_stride, M, S); } if (Dk == 192 && Dv == 128) { return iqk_fa_192_128(int_type_k, int_type_v, nq1, nk1, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, k, v, mask, scale, softcap, qkv, sinksf, M, S); + q, k, v, mask, scale, softcap, qkv, sinksf, sink_stride, M, S); } if (Dk == 192 && Dv == 192) { return iqk_fa_192_192(int_type_k, int_type_v, nq1, nk1, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, k, v, mask, scale, softcap, qkv, sinksf, M, S); + q, k, v, mask, scale, softcap, qkv, sinksf, sink_stride, M, S); } if (Dk == 256 && Dv == 256) { return iqk_fa_256_256(int_type_k, int_type_v, nq1, nk1, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, k, v, mask, scale, softcap, qkv, sinksf, M, S); + q, k, v, mask, scale, softcap, qkv, sinksf, sink_stride, M, S); } if (Dk == 128 && Dv == 128) { return iqk_fa_128_128(int_type_k, int_type_v, nq1, nk1, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, k, v, mask, scale, softcap, qkv, sinksf, M, S); + q, k, v, mask, scale, softcap, qkv, sinksf, sink_stride, M, S); } if (Dk == 96 && Dv == 96) { return iqk_fa_96_96(int_type_k, int_type_v, nq1, nk1, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, k, v, mask, scale, softcap, qkv, sinksf, M, S); + q, k, v, mask, scale, softcap, qkv, sinksf, sink_stride, M, S); } if (Dk == 64 && Dv == 64) { return iqk_fa_64_64(int_type_k, int_type_v, nq1, nk1, stride_q, stride_k, stride_v, stride_m, stride_qkv, - q, k, v, mask, scale, softcap, qkv, sinksf, M, S); + q, k, v, mask, scale, softcap, qkv, sinksf, sink_stride, M, S); } return false; diff --git a/gguf-py/gguf/constants.py b/gguf-py/gguf/constants.py index 0433d5ce..413403b5 100644 --- a/gguf-py/gguf/constants.py +++ b/gguf-py/gguf/constants.py @@ -122,6 +122,9 @@ class Keys: SHARED_KV_LAYERS = "{arch}.attention.shared_kv_layers" KEY_LENGTH_SWA = "{arch}.attention.key_length_swa" VALUE_LENGTH_SWA = "{arch}.attention.value_length_swa" + INDEXER_HEAD_COUNT = "{arch}.attention.indexer.head_count" + INDEXER_KEY_LENGTH = "{arch}.attention.indexer.key_length" + INDEXER_TOP_K = "{arch}.attention.indexer.top_k" VALUE_SCALE = "{arch}.attention.value_scale" OUTPUT_SCALE = "{arch}.attention.output_scale" TEMPERATURE_LENGTH = "{arch}.attention.temperature_length" @@ -262,6 +265,7 @@ class MODEL_ARCH(IntEnum): OPENELM = auto() ARCTIC = auto() DEEPSEEK2 = auto() + DEEPSEEK4 = auto() GLM4_MOE = auto() OPENPANGU = auto() CHATGLM = auto() @@ -384,6 +388,10 @@ class MODEL_TENSOR(IntEnum): NEXTN_HNORM = auto() # nextn tensors (glm4moe) NEXTN_SHARED_HEAD_HEAD = auto() # nextn tensors (glm4moe) NEXTN_SHARED_HEAD_NORM = auto() # nextn tensors (glm4moe) + INDEXER_K_NORM = auto() + INDEXER_PROJ = auto() + INDEXER_ATTN_K = auto() + INDEXER_ATTN_Q_B = auto() MTP_PRE_PROJ = auto() MTP_POST_PROJ = auto() MTP_TOKEN_ORDERING = auto() @@ -466,6 +474,7 @@ MODEL_ARCH_NAMES: dict[MODEL_ARCH, str] = { MODEL_ARCH.OPENELM: "openelm", MODEL_ARCH.ARCTIC: "arctic", MODEL_ARCH.DEEPSEEK2: "deepseek2", + MODEL_ARCH.DEEPSEEK4: "deepseek4", MODEL_ARCH.CHATGLM: "chatglm", MODEL_ARCH.GLM4_MOE: "glm4moe", MODEL_ARCH.OPENPANGU: "openpangu", @@ -589,6 +598,10 @@ TENSOR_NAMES: dict[MODEL_TENSOR, str] = { MODEL_TENSOR.NEXTN_HNORM: "blk.{bid}.nextn.hnorm", MODEL_TENSOR.NEXTN_SHARED_HEAD_HEAD: "blk.{bid}.nextn.shared_head_head", MODEL_TENSOR.NEXTN_SHARED_HEAD_NORM: "blk.{bid}.nextn.shared_head_norm", + MODEL_TENSOR.INDEXER_K_NORM: "blk.{bid}.indexer.k_norm", + MODEL_TENSOR.INDEXER_PROJ: "blk.{bid}.indexer.proj", + MODEL_TENSOR.INDEXER_ATTN_K: "blk.{bid}.indexer.attn_k", + MODEL_TENSOR.INDEXER_ATTN_Q_B: "blk.{bid}.indexer.attn_q_b", MODEL_TENSOR.MTP_PRE_PROJ: "mtp_pre_proj", MODEL_TENSOR.MTP_POST_PROJ: "mtp_post_proj", MODEL_TENSOR.MTP_TOKEN_ORDERING: "mtp_token_ordering", @@ -1310,6 +1323,47 @@ MODEL_TENSORS: dict[MODEL_ARCH, list[MODEL_TENSOR]] = { MODEL_TENSOR.FFN_UP_SHEXP, MODEL_TENSOR.FFN_EXP_PROBS_B ], + MODEL_ARCH.DEEPSEEK4: [ + MODEL_TENSOR.TOKEN_EMBD, + MODEL_TENSOR.OUTPUT_NORM, + MODEL_TENSOR.OUTPUT, + MODEL_TENSOR.ROPE_FREQS, + MODEL_TENSOR.ATTN_NORM, + MODEL_TENSOR.ATTN_Q, + MODEL_TENSOR.ATTN_Q_A, + MODEL_TENSOR.ATTN_Q_B, + MODEL_TENSOR.ATTN_KV_A_MQA, + MODEL_TENSOR.ATTN_KV_B, + MODEL_TENSOR.ATTN_K_B, + MODEL_TENSOR.ATTN_V_B, + MODEL_TENSOR.ATTN_Q_A_NORM, + MODEL_TENSOR.ATTN_KV_A_NORM, + MODEL_TENSOR.ATTN_OUT, + MODEL_TENSOR.ATTN_ROT_EMBD, + MODEL_TENSOR.FFN_GATE_INP, + MODEL_TENSOR.FFN_NORM, + MODEL_TENSOR.FFN_GATE, + MODEL_TENSOR.FFN_DOWN, + MODEL_TENSOR.FFN_UP, + MODEL_TENSOR.FFN_GATE_EXP, + MODEL_TENSOR.FFN_DOWN_EXP, + MODEL_TENSOR.FFN_UP_EXP, + MODEL_TENSOR.FFN_GATE_INP_SHEXP, + MODEL_TENSOR.FFN_GATE_SHEXP, + MODEL_TENSOR.FFN_DOWN_SHEXP, + MODEL_TENSOR.FFN_UP_SHEXP, + MODEL_TENSOR.FFN_EXP_PROBS_B, + MODEL_TENSOR.INDEXER_K_NORM, + MODEL_TENSOR.INDEXER_PROJ, + MODEL_TENSOR.INDEXER_ATTN_K, + MODEL_TENSOR.INDEXER_ATTN_Q_B, + MODEL_TENSOR.NEXTN_EH_PROJ, + MODEL_TENSOR.NEXTN_EMBED_TOKENS, + MODEL_TENSOR.NEXTN_ENORM, + MODEL_TENSOR.NEXTN_HNORM, + MODEL_TENSOR.NEXTN_SHARED_HEAD_HEAD, + MODEL_TENSOR.NEXTN_SHARED_HEAD_NORM, + ], MODEL_ARCH.CHATGLM : [ MODEL_TENSOR.TOKEN_EMBD, MODEL_TENSOR.ROPE_FREQS, @@ -1753,6 +1807,10 @@ MODEL_TENSOR_SKIP: dict[MODEL_ARCH, list[MODEL_TENSOR]] = { MODEL_TENSOR.ROPE_FREQS, MODEL_TENSOR.ATTN_ROT_EMBD, ], + MODEL_ARCH.DEEPSEEK4: [ + MODEL_TENSOR.ROPE_FREQS, + MODEL_TENSOR.ATTN_ROT_EMBD, + ], MODEL_ARCH.CHATGLM: [ MODEL_TENSOR.ROPE_FREQS, ], diff --git a/gguf-py/gguf/gguf_writer.py b/gguf-py/gguf/gguf_writer.py index 09fc6aba..4f606c64 100644 --- a/gguf-py/gguf/gguf_writer.py +++ b/gguf-py/gguf/gguf_writer.py @@ -782,6 +782,15 @@ class GGUFWriter: def add_nextn_predict_layers(self, count: int) -> None: self.add_uint32(Keys.LLM.NEXTN_PREDICT_LAYERS.format(arch=self.arch), count) + def add_attention_indexer_head_count(self, count: int) -> None: + self.add_uint32(Keys.Attention.INDEXER_HEAD_COUNT.format(arch=self.arch), count) + + def add_attention_indexer_key_length(self, length: int) -> None: + self.add_uint32(Keys.Attention.INDEXER_KEY_LENGTH.format(arch=self.arch), length) + + def add_attention_indexer_top_k(self, top_k: int) -> None: + self.add_uint32(Keys.Attention.INDEXER_TOP_K.format(arch=self.arch), top_k) + def add_swin_norm(self, value: bool) -> None: self.add_bool(Keys.LLM.SWIN_NORM.format(arch=self.arch), value) diff --git a/gguf-py/gguf/tensor_mapping.py b/gguf-py/gguf/tensor_mapping.py index a4cba86f..2b80d249 100644 --- a/gguf-py/gguf/tensor_mapping.py +++ b/gguf-py/gguf/tensor_mapping.py @@ -699,6 +699,22 @@ class TensorNameMap: MODEL_TENSOR.NEXTN_SHARED_HEAD_NORM: ( "model.layers.{bid}.shared_head.norm", ), + + MODEL_TENSOR.INDEXER_K_NORM: ( + "model.layers.{bid}.indexer.k_norm", + ), + + MODEL_TENSOR.INDEXER_PROJ: ( + "model.layers.{bid}.indexer.proj", + ), + + MODEL_TENSOR.INDEXER_ATTN_K: ( + "model.layers.{bid}.indexer.attn_k", + ), + + MODEL_TENSOR.INDEXER_ATTN_Q_B: ( + "model.layers.{bid}.indexer.attn_q_b", + ), } # architecture-specific block mappings diff --git a/models/templates/README.md b/models/templates/README.md index 3a649b8f..53792da4 100644 --- a/models/templates/README.md +++ b/models/templates/README.md @@ -23,4 +23,5 @@ These templates can be updated with the following commands: ./scripts/get_chat_template.py Qwen/Qwen3-0.6B > models/templates/Qwen-Qwen3-0.6B.jinja ./scripts/get_chat_template.py zai-org/GLM-4.5 > models/templates/zai-org-GLM-4.5.jinja ./scripts/get_chat_template.py deepseek-ai/DeepSeek-V3.1 > models/templates/deepseek-ai-DeepSeek-V3.1.jinja +./scripts/get_chat_template.py deepseek-ai/DeepSeek-V4 > models/templates/deepseek-ai-DeepSeek-V4.jinja ``` diff --git a/models/templates/deepseek-ai-DeepSeek-V4.jinja b/models/templates/deepseek-ai-DeepSeek-V4.jinja new file mode 100644 index 00000000..d4b0165d --- /dev/null +++ b/models/templates/deepseek-ai-DeepSeek-V4.jinja @@ -0,0 +1,112 @@ +{%- if not add_generation_prompt is defined -%} + {%- set add_generation_prompt = false -%} +{%- endif -%} +{%- if not thinking is defined -%} + {%- if enable_thinking is defined -%} + {%- set thinking = enable_thinking -%} + {%- else -%} + {%- set thinking = false -%} + {%- endif -%} +{%- endif -%} +{%- set dsml_token = '|DSML|' -%} +{%- set thinking_start_token = '' -%} +{%- set thinking_end_token = '' -%} +{%- set tools_header = '## Tools\n\nYou have access to a set of tools to help answer the user\'s question. You can invoke tools by writing a "<' + dsml_token + 'tool_calls>" block like the following:\n\n<' + dsml_token + 'tool_calls>\n<' + dsml_token + 'invoke name="$TOOL_NAME">\n<' + dsml_token + 'parameter name="$PARAMETER_NAME" string="true|false">$PARAMETER_VALUE\n...\n\n<' + dsml_token + 'invoke name="$TOOL_NAME2">\n...\n\n\n\nString parameters should be specified as is and set `string="true"`. For all other types (numbers, booleans, arrays, objects), pass the value in JSON format and set `string="false"`.\n\nIf thinking_mode is enabled (triggered by ' + thinking_start_token + '), you MUST output your complete reasoning inside ' + thinking_start_token + '...' + thinking_end_token + ' BEFORE any tool calls or final response.\n\nOtherwise, output directly after ' + thinking_end_token + ' with tool calls or final response.\n\n### Available Tool Schemas\n\n' -%} +{%- set tools_footer = '\nYou MUST strictly follow the above defined tool name and parameter schemas to invoke tool calls.\n' -%} +{%- set ns = namespace(system_prompt = '', is_first_sp = true) -%} +{%- for message in messages -%} + {%- if message['role'] == 'system' -%} + {%- if ns.is_first_sp -%} + {%- set ns.system_prompt = ns.system_prompt + (message['content'] or '') -%} + {%- set ns.is_first_sp = false -%} + {%- else -%} + {%- set ns.system_prompt = ns.system_prompt + '\n\n' + (message['content'] or '') -%} + {%- endif -%} + {%- endif -%} +{%- endfor -%} +{%- if tools is defined and tools -%} + {%- set ts = namespace(schemas = '') -%} + {%- for tool in tools -%} + {%- if tool['type'] == 'function' -%} + {%- set ts.schemas = ts.schemas + (tool['function'] | tojson) + '\n' -%} + {%- endif -%} + {%- endfor -%} + {%- if ns.system_prompt -%} + {%- set ns.system_prompt = ns.system_prompt + '\n\n' + tools_header + ts.schemas + tools_footer -%} + {%- else -%} + {%- set ns.system_prompt = tools_header + ts.schemas + tools_footer -%} + {%- endif -%} +{%- endif -%} +{{- bos_token -}} +{{- ns.system_prompt -}} +{%- set last_user_idx = namespace(value = -1) -%} +{%- for message in messages -%} + {%- if message['role'] == 'user' or message['role'] == 'developer' or message['role'] == 'tool' -%} + {%- set last_user_idx.value = loop.index0 -%} + {%- endif -%} +{%- endfor -%} +{%- set state = namespace(in_user = false) -%} +{%- for message in messages -%} + {%- if message['role'] == 'user' or message['role'] == 'developer' -%} + {%- if state.in_user -%} + {{- '\n\n' -}} + {%- else -%} + {{- '<|User|>' -}} + {%- set state.in_user = true -%} + {%- endif -%} + {{- message['content'] or '' -}} + {%- elif message['role'] == 'tool' -%} + {%- if state.in_user -%} + {{- '\n\n' -}} + {%- else -%} + {{- '<|User|>' -}} + {%- set state.in_user = true -%} + {%- endif -%} + {{- '' + (message['content'] or '') + '' -}} + {%- elif message['role'] == 'assistant' -%} + {%- set state.in_user = false -%} + {{- '<|Assistant|>' -}} + {%- set is_after_last_user = loop.index0 > last_user_idx.value -%} + {%- if is_after_last_user and thinking -%} + {{- thinking_start_token -}} + {%- if message['reasoning_content'] is defined and message['reasoning_content'] -%} + {{- message['reasoning_content'] -}} + {%- endif -%} + {{- thinking_end_token -}} + {%- else -%} + {{- thinking_end_token -}} + {%- endif -%} + {%- if message['content'] is defined and message['content'] -%} + {{- message['content'] -}} + {%- endif -%} + {%- if message['tool_calls'] -%} + {{- '\n\n<' + dsml_token + 'tool_calls>\n' -}} + {%- for tool in message['tool_calls'] -%} + {%- set func = tool['function'] -%} + {{- '<' + dsml_token + 'invoke name="' + func['name'] + '">\n' -}} + {%- set args = func['arguments'] -%} + {%- if args is string -%} + {%- set args = args | from_json -%} + {%- endif -%} + {%- for key, val in args.items() -%} + {%- if val is string -%} + {{- '<' + dsml_token + 'parameter name="' + key + '" string="true">' + val + '\n' -}} + {%- else -%} + {{- '<' + dsml_token + 'parameter name="' + key + '" string="false">' + (val | tojson) + '\n' -}} + {%- endif -%} + {%- endfor -%} + {{- '\n' -}} + {%- endfor -%} + {{- '' -}} + {%- endif -%} + {{- '<|end▁of▁sentence|>' -}} + {%- endif -%} +{%- endfor -%} +{%- if add_generation_prompt -%} + {{- '<|Assistant|>' -}} + {%- if thinking -%} + {{- thinking_start_token -}} + {%- else -%} + {{- thinking_end_token -}} + {%- endif -%} +{%- endif -%} diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt index 0ee65e18..4d8e7c31 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -43,6 +43,7 @@ add_library(llama llama-spec-features.cpp llama-spec-features-dflash.cpp llama-dflash.cpp + llama-dsv4.cpp llama-vocab.cpp llama-grammar.cpp llama-sampling.cpp @@ -111,6 +112,7 @@ add_library(llama graphs/build_gptneox.cpp graphs/build_arctic.cpp graphs/build_deepseek2.cpp + graphs/build_deepseek4.cpp graphs/build_openpangu.cpp graphs/build_glm4.cpp graphs/build_bitnet.cpp diff --git a/src/graphs/build_deepseek4.cpp b/src/graphs/build_deepseek4.cpp new file mode 100644 index 00000000..9c88fa54 --- /dev/null +++ b/src/graphs/build_deepseek4.cpp @@ -0,0 +1,1701 @@ +#include "../llama-model.h" +#include "../llama-context.h" +#include "../llama-build-context.h" +#include "../llama-dsv4.h" + +#include +#include +#include +#include + +static float dsv4_rope_attn_factor(float freq_scale, float ext_factor) { + if (ext_factor == 0.0f) { + return 1.0f; + } + + return 1.0f / (1.0f + 0.1f*logf(1.0f/freq_scale)); +} + +static size_t dsv4_elem_offset(const ggml_tensor * t, int64_t i) { + return ggml_row_size(t->type, i); +} + +static ggml_tensor * dsv4_view_1d(ggml_context * ctx, ggml_tensor * t, int64_t ne0, int64_t i0) { + return ggml_view_1d(ctx, t, ne0, dsv4_elem_offset(t, i0)); +} + +static ggml_tensor * dsv4_view_2d( + ggml_context * ctx, + ggml_tensor * t, + int64_t ne0, + int64_t ne1, + int64_t i0) { + return ggml_view_2d(ctx, t, ne0, ne1, t->nb[1], dsv4_elem_offset(t, i0)); +} + +static ggml_tensor * dsv4_concat_named( + ggml_context * ctx, + ggml_tensor * a, + ggml_tensor * b, + int dim, + const char * name) { + ggml_tensor * r = ggml_concat(ctx, a, b, dim); + ggml_set_name(r, name); + return r; +} + +static ggml_tensor * dsv4_hc_affine( + ggml_context * ctx, + ggml_tensor * x, + ggml_tensor * scale, + ggml_tensor * base) { + x = ggml_mul(ctx, x, scale); + x = ggml_add(ctx, x, base); + return x; +} + +static ggml_tensor * dsv4_new_i32_input(ggml_context * ctx, ggml_tensor ** dst, int64_t n, const char * name) { + *dst = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, std::max(1, n)); + ggml_set_input(*dst); + ggml_set_name(*dst, name); + return *dst; +} + +static ggml_tensor * dsv4_new_i64_input(ggml_context * ctx, ggml_tensor ** dst, int64_t n, const char * name) { + *dst = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, std::max(1, n)); + ggml_set_input(*dst); + ggml_set_name(*dst, name); + return *dst; +} + +static ggml_tensor * dsv4_new_mask_input(ggml_context * ctx, ggml_tensor ** dst, int64_t n_kv, int64_t n_tokens, const char * name, + ggml_type mask_type) { + //*dst = ggml_new_tensor_2d(ctx, mask_type, std::max(1, n_kv), GGML_PAD(std::max(1, n_tokens), GGML_KQ_MASK_PAD)); + *dst = ggml_new_tensor_2d(ctx, mask_type, std::max(1, n_kv), std::max(1, n_tokens)); + ggml_set_input(*dst); + ggml_set_name(*dst, name); + return *dst; +} + +static void dsv4_build_plan_inputs( + ggml_context * ctx, + llama_context::dsv4_runtime::comp_inputs & inputs, + const llama_context::dsv4_runtime::comp_plan & plan, + const char * tag, + int64_t n_tokens, + bool create_mask = true, bool flash_attn = true) { + //printf("%s(%s): n_tokens = %ld\n", __func__, tag, n_tokens); + dsv4_new_i32_input(ctx, &inputs.state_pos, (int64_t) plan.state_pos.size(), (std::string(tag) + "_state_pos").c_str()); + dsv4_new_i32_input(ctx, &inputs.state_persist_src_idxs, (int64_t) plan.state_persist_src_idxs.size(), (std::string(tag) + "_persist_src").c_str()); + dsv4_new_i32_input(ctx, &inputs.state_persist_dst_idxs, (int64_t) plan.state_persist_dst_idxs.size(), (std::string(tag) + "_persist_dst").c_str()); + dsv4_new_i32_input(ctx, &inputs.state_read_idxs, (int64_t) plan.state_read_idxs.size(), (std::string(tag) + "_state_read").c_str()); + dsv4_new_i64_input(ctx, &inputs.state_write_idxs, (int64_t) plan.state_write_idxs.size(), (std::string(tag) + "_state_write").c_str()); + dsv4_new_i32_input(ctx, &inputs.state_write_pos, (int64_t) plan.state_write_pos.size(), (std::string(tag) + "_write_pos").c_str()); + if (create_mask) { + auto type = flash_attn ? GGML_TYPE_F16 : GGML_TYPE_F32; + dsv4_new_mask_input(ctx, &inputs.kq_mask, std::max(1, plan.n_kv), n_tokens, (std::string(tag) + "_kq_mask").c_str(), type); + } else { + inputs.kq_mask = nullptr; + } +} + +static ggml_tensor * dsv4_append_zero_row(ggml_context * ctx, ggml_tensor * t, bool neg_inf) { + ggml_tensor * row = ggml_view_1d(ctx, t, t->ne[0], 0); + row = neg_inf ? ggml_scale_bias(ctx, row, 0.0f, -INFINITY) : ggml_scale(ctx, row, 0.0f); + row = ggml_reshape_2d(ctx, row, t->ne[0], 1); + return dsv4_concat_named(ctx, t, row, 1, "dsv4_append_zero_row"); +} + +static ggml_tensor * dsv4_cache_view_2d( + ggml_context * ctx, + ggml_tensor * cache, + int64_t dim0, + int64_t dim1) { + return ggml_view_2d(ctx, cache, dim0, dim1, ggml_row_size(cache->type, dim0), 0); +} + +static ggml_tensor * dsv4_build_mask_stream_view( + ggml_context * ctx, + ggml_tensor * mask, + int64_t n_stream, + int64_t n_tokens) { + if (n_stream <= 1) { + return mask; + } + + GGML_ASSERT(n_tokens % n_stream == 0); + const int64_t n_tokens_stream = n_tokens/n_stream; + return ggml_view_4d(ctx, mask, mask->ne[0], n_tokens_stream, 1, n_stream, + mask->nb[1], mask->nb[1]*n_tokens_stream, mask->nb[1]*n_tokens_stream, 0); +} + +static ggml_tensor * dsv4_build_raw_mask_view( + ggml_context * ctx, + ggml_tensor * mask, + ggml_tensor * raw_k_read_idxs, + int64_t n_kv, + int64_t n_tokens, + int64_t n_stream, + const llm_build_cb & cb, int il) { + const int64_t n_tokens_stream = n_stream > 0 ? n_tokens/n_stream : n_tokens; + const int64_t n_rows_stream = GGML_PAD(n_kv, 256); + + if (raw_k_read_idxs == nullptr) { + auto base = ggml_view_2d(ctx, mask, n_kv, n_tokens, mask->nb[1], 0); + if (!ggml_is_contiguous(base)) { + base = ggml_cont(ctx, base); + cb(base, "mask_base", il); + } + return n_stream == 1 ? base : dsv4_build_mask_stream_view(ctx, base, n_stream, n_tokens); + //auto base = mask->ne[0] == n_kv && mask->ne[1] == n_tokens ? mask + // : ggml_cont(ctx, ggml_view_2d(ctx, mask, n_kv, n_tokens, mask->nb[1], 0)); + //return n_stream == 1 ? base : dsv4_build_mask_stream_view(ctx, base, n_stream, n_tokens); + ////ggml_tensor * base = ggml_cont(ctx, ggml_view_2d(ctx, mask, n_kv, n_tokens, mask->nb[1], 0)); + ////return dsv4_build_mask_stream_view(ctx, base, n_stream, n_tokens); + } + + if (n_stream <= 0 || n_tokens % n_stream != 0 || raw_k_read_idxs->ne[0] < n_rows_stream*n_stream) { + //printf("%s(Oops): %d, %d, %d\n", __func__, n_stream <= 0, n_tokens % n_stream != 0, raw_k_read_idxs->ne[0] < n_rows_stream*n_stream); + ggml_tensor * base = ggml_cont(ctx, ggml_view_2d(ctx, mask, n_kv, n_tokens, mask->nb[1], 0)); + cb(base, "mask_base1", il); + return dsv4_build_mask_stream_view(ctx, base, std::max(1, n_stream), n_tokens); + } + + if (n_stream == 1 && mask->ne[0] == raw_k_read_idxs->ne[0]) { + return mask; + } + + printf("%s: Oops(%s). mask is %ld x %ld x %ld x %ld. n_stream = %ld, n_tokens = %ld, raw_k_read_idxs = %ld x %ld x %ld x %ld\n", + __func__, mask->name, mask->ne[0], mask->ne[1], mask->ne[2], mask->ne[3], n_stream, n_tokens, + raw_k_read_idxs->ne[0], raw_k_read_idxs->ne[1], raw_k_read_idxs->ne[2], raw_k_read_idxs->ne[3]); + + ggml_tensor * mask_t = ggml_cont(ctx, ggml_transpose(ctx, mask)); + ggml_tensor * result = nullptr; + for (int64_t s = 0; s < n_stream; ++s) { + ggml_tensor * idxs = ggml_view_1d(ctx, raw_k_read_idxs, n_kv, + s*n_rows_stream*ggml_element_size(raw_k_read_idxs)); + ggml_tensor * mask_s = ggml_view_2d(ctx, mask_t, n_tokens_stream, mask->ne[0], mask_t->nb[1], + s*n_tokens_stream*mask_t->nb[0]); + ggml_tensor * rows = ggml_get_rows(ctx, mask_s, idxs); + ggml_tensor * stream = ggml_reshape_4d(ctx, ggml_cont(ctx, ggml_transpose(ctx, rows)), + n_kv, n_tokens_stream, 1, 1); + result = result == nullptr ? stream : ggml_concat(ctx, result, stream, 3); + cb(result, "raw_mask_view", s); + } + return result; +} + +static ggml_tensor * dsv4_pad_raw_k_to( + ggml_context * ctx, + ggml_tensor * raw_k, + int64_t n_kv_target) { + const int64_t n_kv_cur = raw_k->ne[2]; + if (n_kv_target <= n_kv_cur) { + return raw_k; + } + + //printf("Oops: padding KV cache\n"); + + const int64_t n_pad = n_kv_target - n_kv_cur; + ggml_tensor * row0 = ggml_view_4d(ctx, raw_k, + raw_k->ne[0], raw_k->ne[1], 1, raw_k->ne[3], + raw_k->nb[1], raw_k->nb[2], raw_k->nb[3], 0); + ggml_tensor * zero_row = ggml_cont(ctx, row0); + if (zero_row->type != GGML_TYPE_F32) { + zero_row = ggml_cast(ctx, zero_row, GGML_TYPE_F32); + } + zero_row = ggml_scale(ctx, zero_row, 0.0f); + if (raw_k->type != zero_row->type && !ggml_is_quantized(raw_k->type)) { + zero_row = ggml_cast(ctx, zero_row, raw_k->type); + } + ggml_tensor * zeros = ggml_repeat_4d(ctx, zero_row, zero_row->ne[0], zero_row->ne[1], n_pad, zero_row->ne[3]); + return dsv4_concat_named(ctx, raw_k, zeros, 2, "dsv4_raw_k_pad"); +} + +static ggml_tensor * dsv4_pad_raw_mask_to( + ggml_context * ctx, + ggml_tensor * raw_mask, + int64_t n_kv_target, + int64_t n_tokens) { + const int64_t n_kv_cur = raw_mask->ne[0]; + if (n_kv_target <= n_kv_cur) { + return raw_mask; + } + + printf("%s: Oops, padding mask\n", __func__); + + const int64_t n_pad = n_kv_target - n_kv_cur; + GGML_UNUSED(n_tokens); + ggml_tensor * pad = ggml_new_tensor_4d(ctx, raw_mask->type, n_pad, raw_mask->ne[1], raw_mask->ne[2], raw_mask->ne[3]); + pad = ggml_fill(ctx, pad, -INFINITY); + return dsv4_concat_named(ctx, raw_mask, pad, 0, "dsv4_raw_mask_pad"); +} + +static ggml_tensor * dsv4_pad_mask_tokens( + ggml_context * ctx, + ggml_tensor * mask, + int64_t n_tokens) { + const int64_t n_stream = std::max(1, mask->ne[3]); + GGML_ASSERT(n_tokens % n_stream == 0); + const int64_t n_tokens_pad = GGML_PAD(n_tokens/n_stream, GGML_KQ_MASK_PAD); + if (mask->ne[1] >= n_tokens_pad) { + return mask; + } + + ggml_tensor * pad = ggml_new_tensor_4d(ctx, mask->type, mask->ne[0], n_tokens_pad - mask->ne[1], mask->ne[2], mask->ne[3]); + pad = ggml_fill(ctx, pad, -INFINITY); + auto new_mask = dsv4_concat_named(ctx, mask, pad, 1, "dsv4_mask_tokens_pad"); + + //printf("%s: Oops: padding mask %s from %ld x %ld x %ld x %ld to %ld x %ld x %ld x %ld\n", __func__, mask->name, + // mask->ne[0], mask->ne[1], mask->ne[2], mask->ne[3], + // new_mask->ne[0], new_mask->ne[1], new_mask->ne[2], new_mask->ne[3]); + return new_mask; +} + +static ggml_tensor * dsv4_cache_view_3d( + ggml_context * ctx, + ggml_tensor * cache, + int64_t n_embd_head, + int64_t n_kv) { + return ggml_view_3d(ctx, cache, + n_embd_head, 1, n_kv, + ggml_row_size(cache->type, n_embd_head), + ggml_row_size(cache->type, n_embd_head), + 0); +} + +static ggml_tensor * dsv4_slice_1d( + ggml_context * ctx, + ggml_tensor * t, + int64_t offset, + int64_t size) { + return ggml_view_1d(ctx, t, size, offset*ggml_element_size(t)); +} + +static ggml_tensor * dsv4_require_f32_rows( + ggml_context * ctx, + ggml_tensor * t) { + if (t == nullptr || t->type == GGML_TYPE_F32) { + return t; + } + + return ggml_cast(ctx, t, GGML_TYPE_F32); +} + +static ggml_tensor * dsv4_cache_read_f32(ggml_context * ctx, ggml_tensor * t) { + return t != nullptr && ggml_is_quantized(t->type) ? ggml_cast(ctx, t, GGML_TYPE_F32) : t; +} + +static ggml_tensor * dsv4_cache_stream_view_3d( + ggml_context * ctx, + ggml_tensor * cache, + int64_t n_embd_head, + int64_t n_kv, + int64_t kv_size, + int64_t stream) { + const size_t offset = (size_t) ggml_row_size(cache->type, n_embd_head) * (size_t) kv_size * (size_t) stream; + return ggml_view_3d(ctx, cache, + n_embd_head, 1, n_kv, + ggml_row_size(cache->type, n_embd_head), + ggml_row_size(cache->type, n_embd_head), + offset); +} + +static ggml_tensor * dsv4_cache_stream_view_4d( + ggml_context * ctx, + ggml_tensor * cache, + int64_t n_embd_head, + int64_t n_kv, + int64_t kv_size, + int64_t s0, + int64_t n_stream) { + const size_t row_size = ggml_row_size(cache->type, n_embd_head); + const size_t offset = row_size*(size_t) kv_size*(size_t) s0; + return ggml_view_4d(ctx, cache, + n_embd_head, 1, n_kv, n_stream, + row_size, + row_size, + row_size*(size_t) kv_size, + offset); +} + +static ggml_tensor * dsv4_raw_get_k( + llama_context * lctx, + ggml_context * ctx, + ggml_tensor * cache, + ggml_tensor * raw_k_read_idxs, + int64_t n_embd_head, const llm_build_cb & cb, [[maybe_unused]] int il) { + if (cache == nullptr) { + return nullptr; + } + + const auto & raw = lctx->dsv4.raw; + const int64_t n_kv_visible = raw.n_kv; + if (n_kv_visible <= 0) { + return nullptr; + } + + // Keep the visible row count for masks, but expose a stable 256-row cache + // view to attention. + const int64_t n_kv = std::max(256, GGML_PAD(n_kv_visible, 256)); + + const auto & sinfo = raw.sinfo_read; + const int64_t n_stream = (int64_t) sinfo.n_stream(); + if (n_stream <= 0) { + return nullptr; + } + + const int64_t n_embd_gqa = cache->ne[0]; + GGML_ASSERT(n_embd_head > 0); + GGML_ASSERT(n_embd_gqa % n_embd_head == 0); + + const int64_t n_head_kv = n_embd_gqa/n_embd_head; + + if (n_stream == 1 && lctx->kv_self.n == raw_k_read_idxs->ne[0]) { + return ggml_view_3d(ctx, cache, n_embd_head, n_head_kv, n_kv, + ggml_row_size(cache->type, n_embd_head), + ggml_row_size(cache->type, n_embd_head)*n_head_kv, 0); + } + + GGML_ASSERT(raw_k_read_idxs != nullptr); + GGML_ASSERT(raw_k_read_idxs->ne[0] >= n_kv*n_stream); + + // Gather controller-owned slots into the attention layout. + ggml_tensor * cache_2d = dsv4_cache_view_2d(ctx, cache, n_embd_gqa, cache->ne[1]); + ggml_tensor * idxs = raw_k_read_idxs->type == GGML_TYPE_I32 + ? raw_k_read_idxs : ggml_cast(ctx, raw_k_read_idxs, GGML_TYPE_I32); + ggml_tensor * rows = ggml_get_rows(ctx, cache_2d, idxs); + cb(rows, "raw_k", il); + if (ggml_is_quantized(cache->type) && rows->type != GGML_TYPE_F32) { + rows = ggml_cast(ctx, rows, GGML_TYPE_F32); + } else if (rows->type != cache->type && !ggml_is_quantized(cache->type)) { + rows = ggml_cast(ctx, rows, cache->type); + } + + ggml_tensor * raw_k = ggml_reshape_4d(ctx, rows, n_embd_head, n_head_kv, n_kv, n_stream); + + return raw_k; +} + +static ggml_tensor * dsv4_raw_cpy_k( + llama_context * lctx, + ggml_context * ctx, + ggml_tensor * cache, + ggml_tensor * k_cur, + ggml_tensor * raw_k_write_src_idxs, + ggml_tensor * raw_k_write_idxs, + ggml_cgraph * gf, + int64_t n_embd_head, + const llm_build_cb & cb, + int64_t il) { + if (cache == nullptr || k_cur == nullptr || raw_k_write_idxs == nullptr || raw_k_write_src_idxs == nullptr) { + return nullptr; + } + + GGML_ASSERT(2*il + 1 < (int64_t) lctx->cache_copies.size()); + GGML_ASSERT(k_cur->ne[1] == 1); + + ggml_tensor * cache_2d = dsv4_cache_view_2d(ctx, cache, n_embd_head, cache->ne[1]); + ggml_tensor * cur_2d = ggml_view_2d(ctx, k_cur, n_embd_head, k_cur->ne[2], k_cur->nb[2], 0); + ggml_tensor * write = nullptr; + + const auto & sinfo = lctx->dsv4.raw.sinfo_write; + if (sinfo.n_stream() <= 1 && cur_2d->ne[1] == raw_k_write_idxs->ne[0]) { + cur_2d = dsv4_require_f32_rows(ctx, cur_2d); + write = ggml_set_rows(ctx, cache_2d, cur_2d, raw_k_write_idxs); + } else if (sinfo.n_stream() <= 1) { + ggml_tensor * src_idxs = raw_k_write_src_idxs->type == GGML_TYPE_I32 ? raw_k_write_src_idxs : ggml_cast(ctx, raw_k_write_src_idxs, GGML_TYPE_I32); + ggml_tensor * cur_sel = ggml_get_rows(ctx, cur_2d, src_idxs); + cb(cur_sel, "sel", il); + cur_sel = dsv4_require_f32_rows(ctx, cur_sel); + write = ggml_set_rows(ctx, cache_2d, cur_sel, raw_k_write_idxs); + } else { + const int64_t n_fanout = (int64_t) sinfo.size()*(int64_t) sinfo.n_stream(); + + GGML_ASSERT(sinfo.n_stream() > 1); + GGML_ASSERT(raw_k_write_idxs->ne[0] == n_fanout); + GGML_ASSERT(raw_k_write_src_idxs->ne[0] == n_fanout); + + for (uint32_t s = 0; s < sinfo.n_stream(); ++s) { + ggml_tensor * src_idxs_s = ggml_view_1d(ctx, raw_k_write_src_idxs, sinfo.size(), + s*sinfo.size()*ggml_element_size(raw_k_write_src_idxs)); + ggml_tensor * k_idxs_s = ggml_view_1d(ctx, raw_k_write_idxs, sinfo.size(), s*sinfo.size()*ggml_element_size(raw_k_write_idxs)); + ggml_tensor * cur_sel = ggml_get_rows(ctx, cur_2d, src_idxs_s); + cb(cur_sel, "sel", il); + ggml_tensor * cur_f32 = dsv4_require_f32_rows(ctx, cur_sel); + ggml_tensor * cur = ggml_set_rows(ctx, cache_2d, cur_f32, k_idxs_s); + if (write == nullptr) { + write = cur; + } else { + ggml_build_forward_expand(gf, cur); + } + } + } + + lctx->cache_copies[2*il + 0].cpy = write; + lctx->cache_copies[2*il + 0].step = ggml_row_size(cache->type, n_embd_head); + ggml_build_forward_expand(gf, write); + + return write; +} + +static ggml_tensor * dsv4_comp_get_k( + ggml_context * ctx, + ggml_tensor * cache, + const llama_context::dsv4_runtime::comp_context & comp, + int64_t n_embd_head, + int64_t kv_size) { + const int64_t n_kv = comp.n_kv; + if (cache == nullptr || n_kv <= 0) { + return nullptr; + } + + if (comp.sinfo.n_stream() == 0) { + return dsv4_cache_read_f32(ctx, ggml_reshape_4d(ctx, dsv4_cache_view_3d(ctx, cache, n_embd_head, n_kv), n_embd_head, 1, n_kv, 1)); + } + + return dsv4_cache_read_f32(ctx, dsv4_cache_stream_view_4d(ctx, cache, n_embd_head, n_kv, kv_size, comp.sinfo.s0, (int64_t) comp.sinfo.n_stream())); +} + +static ggml_tensor * dsv4_comp_cpy_k( + ggml_context * ctx, + ggml_tensor * cache, + ggml_tensor * cur, + ggml_tensor * idxs, + int64_t n_embd_head) { + ggml_tensor * cache_2d = dsv4_cache_view_2d(ctx, cache, n_embd_head, cache->ne[1]); + cur = dsv4_require_f32_rows(ctx, cur); + return ggml_set_rows(ctx, cache_2d, cur, idxs); +} + +static ggml_tensor * dsv4_comp_state_cpy( + ggml_context * ctx, + ggml_tensor * cache, + ggml_tensor * cur, + ggml_tensor * idxs) { + cur = dsv4_require_f32_rows(ctx, cur); + return ggml_set_rows(ctx, cache, cur, idxs); +} + +static ggml_tensor * dsv4_repeat_streams(ggml_context * ctx, ggml_tensor * t, int64_t n_stream) { + if (t->ne[3] == n_stream) { + return t; + } + + GGML_ASSERT(t->ne[3] == 1); + return ggml_repeat_4d(ctx, t, t->ne[0], t->ne[1], t->ne[2], n_stream); +} + +static ggml_tensor * dsv4_build_kq_zero_bias( + ggml_context * ctx, + const llama_cparams & cparams, + ggml_tensor * kq_mask, + int64_t n_head) { + GGML_UNUSED(ctx); + GGML_UNUSED(n_head); + + if (!cparams.flash_attn || kq_mask->ne[3] == 1) { + return nullptr; + } + + // The zero-bias fallback is only needed for unified multi-stream KV. + // The DSV4 cache/controller is non-unified, so keep the direct FA path. + return nullptr; +} + +static ggml_tensor * dsv4_build_attn( + ggml_context * ctx, + const llama_hparams & hparams, + const llama_cparams & cparams, + ggml_tensor * q, + ggml_tensor * k, + ggml_tensor * v, + ggml_tensor * kq_b, + ggml_tensor * kq_mask, + ggml_tensor * sinks, + float kq_scale, + const llm_build_cb & cb, + int il, + int n_compressed, + ggml_cgraph * gf) { + const bool v_trans = v->nb[1] > v->nb[2]; + const int64_t n_stream = k->ne[3]; + + if (!cparams.flash_attn && n_stream > 1) { + GGML_ASSERT(kq_b == nullptr); + GGML_ASSERT(q->ne[2] % n_stream == 0); + const int64_t n_tokens_stream = q->ne[2]/n_stream; + ggml_tensor * result = nullptr; + + for (int64_t s = 0; s < n_stream; ++s) { + ggml_tensor * q_s = ggml_view_3d(ctx, q, q->ne[0], q->ne[1], n_tokens_stream, + q->nb[1], q->nb[2], s*n_tokens_stream*q->nb[2]); + ggml_tensor * k_s = ggml_view_4d(ctx, k, k->ne[0], k->ne[1], k->ne[2], 1, + k->nb[1], k->nb[2], k->nb[3], s*k->nb[3]); + ggml_tensor * v_s = ggml_view_4d(ctx, v, v->ne[0], v->ne[1], v->ne[2], 1, + v->nb[1], v->nb[2], v->nb[3], s*v->nb[3]); + ggml_tensor * mask_s = kq_mask; + if (ggml_is_matrix(kq_mask)) { + mask_s = ggml_view_2d(ctx, kq_mask, kq_mask->ne[0], n_tokens_stream, + kq_mask->nb[1], s*n_tokens_stream*kq_mask->nb[1]); + } else { + mask_s = ggml_view_2d(ctx, kq_mask, kq_mask->ne[0], kq_mask->ne[1], + kq_mask->nb[1], s*kq_mask->nb[3]); + } + + ggml_tensor * cur_s = dsv4_build_attn(ctx, hparams, cparams, + q_s, k_s, v_s, nullptr, mask_s, sinks, kq_scale, cb, il, n_compressed, gf); + result = result == nullptr ? cur_s : ggml_concat(ctx, result, cur_s, 1); + } + return result; + } + + q = ggml_view_4d(ctx, q, q->ne[0], q->ne[1], q->ne[2] / n_stream, n_stream, + q->nb[1], q->nb[2], q->nb[3] / n_stream, 0); + q = ggml_permute(ctx, q, 0, 2, 1, 3); + k = ggml_permute(ctx, k, 0, 2, 1, 3); + v = ggml_permute(ctx, v, 0, 2, 1, 3); + + // The DSV4 cache/controller is non-unified. Keep the eligibility rule + // explicit so a future unified cache cannot route multi-stream masks + // through Flash Attention accidentally. + constexpr bool kv_unified = false; + const bool use_flash_attn = cparams.flash_attn && + (!kv_unified || kq_mask->ne[3] == 1) && + kq_b == nullptr; + if (use_flash_attn) { + GGML_ASSERT(kq_b == nullptr && "Flash attention does not support KQ bias yet"); + + if (v_trans) { + v = ggml_transpose(ctx, v); + } + + if (k->type == GGML_TYPE_F32) { + k = ggml_cast(ctx, k, GGML_TYPE_F16); + } + + if (v->type == GGML_TYPE_F32) { + v = ggml_cast(ctx, v, GGML_TYPE_F16); + } + + if (kq_mask->type == GGML_TYPE_F32) { + kq_mask = ggml_cast(ctx, kq_mask, GGML_TYPE_F16); + } + + ggml_tensor * selected = nullptr; + if (n_compressed > 0) { + int n_compressed_padded = GGML_PAD(n_compressed, 256); + if (n_compressed_padded < kq_mask->ne[0]) { + selected = ggml_mask_to_index(ctx, kq_mask, n_compressed_padded); + cb(selected, "mask_to_idx", il); + ggml_build_forward_expand(gf, selected); + //if (q->ne[1] == 1) { + // selected = ggml_view_1d(ctx, selected, selected->ne[0], 0); + // kq_mask = ggml_view_1d(ctx, kq_mask, kq_mask->ne[0], 0); + // kq_mask = ggml_reshape_2d(ctx, kq_mask, 1, kq_mask->ne[0]); + // kq_mask = ggml_get_rows(ctx, kq_mask, selected); + // kq_mask = ggml_reshape_1d(ctx, kq_mask, kq_mask->ne[1]); + // k = ggml_get_rows(ctx, k, selected); + // auto kq = ggml_mul_mat(ctx, k, q); + // if (kq_b != nullptr) { + // kq = ggml_add(ctx, kq, kq_b); + // } + // kq = ggml_soft_max_ext(ctx, kq, kq_mask, kq_scale, 0.0f); + // ggml_soft_max_add_sinks(kq, sinks); + // v = ggml_cont(ctx, ggml_transpose(ctx, k)); + // auto kqv = ggml_mul_mat(ctx, v, kq); + // kqv = ggml_permute(ctx, kqv, 0, 2, 1, 3); + // kqv = ggml_reshape_2d(ctx, kqv, kqv->ne[0]*kqv->ne[1], kqv->ne[2]*kqv->ne[3]); + // return kqv; + //} + } + } + + ggml_tensor * cur = ggml_flash_attn_ext(ctx, q, k, v, kq_mask, kq_scale, hparams.f_max_alibi_bias, + hparams.attn_soft_cap ? hparams.f_attn_logit_softcapping : 0.0f); + cb(cur, "fattn", il); + // DSV4 uses the generic CPU FA path here for numerical correctness. + if (selected) { + cur->src[5] = selected; + } else { + cur->op_params[4] = GGML_FLASH_ATTN_EXT_IQK_DISABLED; + } + ggml_flash_attn_ext_add_sinks(cur, sinks); + ggml_flash_attn_ext_set_prec(cur, GGML_PREC_F32); + ggml_build_forward_expand(gf, cur); + return ggml_reshape_2d(ctx, cur, cur->ne[0] * cur->ne[1], cur->ne[2] * cur->ne[3]); + } + + ggml_tensor * kq = ggml_mul_mat(ctx, k, q); + cb(kq, "kq", il); + ggml_mul_mat_set_prec(kq, GGML_PREC_F32); + + if (kq_b != nullptr) { + kq = ggml_add(ctx, kq, kq_b); + cb(kq, "kq_plus_kq_b", il); + } + + if (kq->type != GGML_TYPE_F32) { + kq = ggml_cast(ctx, kq, GGML_TYPE_F32); + } + + if (hparams.attn_soft_cap) { + kq = ggml_scale(ctx, kq, 1.0f / hparams.f_attn_logit_softcapping); + kq = ggml_tanh(ctx, kq); + kq = ggml_scale(ctx, kq, hparams.f_attn_logit_softcapping); + kq = ggml_soft_max_ext(ctx, kq, kq_mask, kq_scale, hparams.f_max_alibi_bias); + ggml_soft_max_add_sinks(kq, sinks); + } else { + kq = ggml_soft_max_ext(ctx, kq, kq_mask, kq_scale, hparams.f_max_alibi_bias); + ggml_soft_max_add_sinks(kq, sinks); + } + cb(kq, "kq_soft_max", il); + + if (!v_trans) { + v = ggml_cont(ctx, ggml_transpose(ctx, v)); + cb(v, "v_cont", il); + } + + ggml_tensor * kqv = ggml_mul_mat(ctx, v, kq); + cb(kqv, "kqv", il); + ggml_tensor * cur = ggml_permute(ctx, kqv, 0, 2, 1, 3); + return ggml_cont_2d(ctx, cur, cur->ne[0] * cur->ne[1], cur->ne[2] * cur->ne[3]); +} + +static ggml_tensor * build_hc_sinkhorn( + ggml_context * ctx0, + const llama_hparams & hparams, + ggml_tensor * comb) { + comb = ggml_soft_max(ctx0, comb); + + ggml_tensor * eps = ggml_new_tensor_1d(ctx0, GGML_TYPE_F32, 1); + eps = ggml_fill(ctx0, eps, hparams.dsv4_hc_eps); + comb = ggml_add(ctx0, comb, eps); + + auto norm_cols = [&]() { + ggml_tensor * comb_src_dst = ggml_cont(ctx0, ggml_permute(ctx0, comb, 1, 0, 2, 3)); + ggml_tensor * col_sum = ggml_sum_rows(ctx0, comb_src_dst); + col_sum = ggml_add(ctx0, col_sum, eps); + col_sum = ggml_permute(ctx0, col_sum, 1, 0, 2, 3); + comb = ggml_div(ctx0, comb, col_sum); + }; + + auto norm_rows = [&]() { + ggml_tensor * row_sum = ggml_sum_rows(ctx0, comb); + row_sum = ggml_add(ctx0, row_sum, eps); + comb = ggml_div(ctx0, comb, row_sum); + }; + + norm_cols(); + for (uint32_t i = 1; i < hparams.dsv4_hc_sinkhorn_iters; ++i) { + norm_rows(); + norm_cols(); + } + + return comb; +} + +static ggml_tensor * build_hc_pre( + ggml_context * ctx0, + llm_build_context & llm, + const llama_hparams & hparams, + int64_t n_embd, + float norm_rms_eps, + ggml_tensor * x, + ggml_tensor * hc_fn, + ggml_tensor * hc_scale, + ggml_tensor * hc_base, + ggml_tensor ** post_out, + ggml_tensor ** comb_out, + const llm_build_cb & cb, int il) { + const int64_t hc = hparams.dsv4_hc_mult; + const int64_t nt = x->ne[2]; + + if (!ggml_is_contiguous(x)) { + x = ggml_cont(ctx0, x); + } + auto flat = ggml_reshape_2d(ctx0, x, n_embd * hc, nt); + auto normed = ggml_rms_norm(ctx0, flat, norm_rms_eps); + cb(normed, "hc_pre", il); + auto mixes = ggml_mul_mat(ctx0, hc_fn, normed); + cb(mixes, "hc_pre_mixes", il); + + auto all = ggml_hc_pre(ctx0, mixes, hc_scale, hc_base, hc, hparams.dsv4_hc_sinkhorn_iters, hparams.dsv4_hc_eps); + + auto pre = ggml_view_2d(ctx0, all, hc, nt, hc*sizeof(float), 0); + auto post = ggml_view_2d(ctx0, all, hc, nt, hc*sizeof(float), hc*nt*sizeof(float)); + auto comb = ggml_view_3d(ctx0, all, hc, hc, nt, hc*sizeof(float), hc*hc*sizeof(float), 2*hc*nt*sizeof(float)); + + ////ggml_tensor * mixes = llm.build_mhc_pre_projection(x, hc_fn, nullptr, + //// n_embd, hc, norm_rms_eps, false); + + ////printf("hc_scale: %ld x %ld x %ld x %ld, hc_base: %ld x %ld x %ld x %ld\n", + //// hc_scale->ne[0], hc_scale->ne[1], hc_scale->ne[2], hc_scale->ne[3], + //// hc_base->ne[0], hc_base->ne[1], hc_base->ne[2], hc_base->ne[3]); + //ggml_tensor * scale_pre = dsv4_view_1d(ctx0, hc_scale, 1, 0); + //ggml_tensor * scale_post = dsv4_view_1d(ctx0, hc_scale, 1, 1); + //ggml_tensor * scale_comb = dsv4_view_1d(ctx0, hc_scale, 1, 2); + + //ggml_tensor * base_pre = dsv4_view_1d(ctx0, hc_base, hc, 0); + //ggml_tensor * base_post = dsv4_view_1d(ctx0, hc_base, hc, hc); + //ggml_tensor * base_comb = dsv4_view_1d(ctx0, hc_base, hc*hc, 2*hc); + + //ggml_tensor * pre = ggml_cont(ctx0, dsv4_view_2d(ctx0, mixes, hc, nt, 0)); + //pre = dsv4_hc_affine(ctx0, pre, scale_pre, base_pre); + //pre = ggml_sigmoid(ctx0, pre); + //pre = ggml_scale_bias(ctx0, pre, 1.0f, hparams.dsv4_hc_eps); + + //auto post = ggml_cont(ctx0, dsv4_view_2d(ctx0, mixes, hc, nt, hc)); + //post = dsv4_hc_affine(ctx0, post, scale_post, base_post); + //post = ggml_sigmoid(ctx0, post); + //post = ggml_scale(ctx0, post, 2.0f); + + //auto comb = ggml_cont(ctx0, dsv4_view_2d(ctx0, mixes, hc*hc, nt, 2*hc)); + //comb = dsv4_hc_affine(ctx0, comb, scale_comb, base_comb); + //comb = ggml_sinkhorn(ctx0, comb, hc, hparams.dsv4_hc_sinkhorn_iters, hparams.dsv4_hc_eps, false); + ////*comb = ggml_reshape_3d(ctx0, *comb, hc, hc, nt); + ////*comb = build_hc_sinkhorn(ctx0, hparams, *comb); + //printf("pre: %ld x %ld x %ld x %ld, post: %ld x %ld x %ld x %ld, comb: %ld x %ld x %ld x %ld\n", + // pre->ne[0], pre->ne[1], pre->ne[2], pre->ne[3], post->ne[0], post->ne[1], post->ne[2], post->ne[3], + // comb->ne[0], comb->ne[1], comb->ne[2], comb->ne[3]); + + *post_out = post; + *comb_out = comb; + + return llm.build_mhc_weighted_sum(x, pre, n_embd, hc); +} + +static ggml_tensor * build_hc_head( + ggml_context * ctx0, + llm_build_context & llm, + const llama_hparams & hparams, + int64_t n_embd, + float norm_rms_eps, + ggml_tensor * x, + ggml_tensor * hc_fn, + ggml_tensor * hc_scale, + ggml_tensor * hc_base) { + const int64_t hc = hparams.dsv4_hc_mult; + + ggml_tensor * mixes = llm.build_mhc_pre_projection(x, hc_fn, nullptr, + n_embd, hc, norm_rms_eps, false); + ggml_tensor * pre = dsv4_hc_affine(ctx0, mixes, hc_scale, hc_base); + pre = ggml_sigmoid(ctx0, pre); + pre = ggml_scale_bias(ctx0, pre, 1.0f, hparams.dsv4_hc_eps); + + return llm.build_mhc_weighted_sum(x, pre, n_embd, hc); +} + +static ggml_tensor * build_hca_compressed_kv_from_state( + ggml_context * ctx0, + llm_build_context & llm, + ggml_tensor * kv_state, + ggml_tensor * score_state, + ggml_tensor * state_read_idxs, + ggml_tensor * comp_pos, + ggml_tensor * norm, + int64_t n_embd_head, + const llm_build_cb & cb, + int il) { + const int64_t n_embd_head_rope = llm.hparams.n_rot; + const int64_t n_embd_head_nope = n_embd_head - n_embd_head_rope; + const int64_t n_blocks = comp_pos ? comp_pos->ne[0] : 0; + + GGML_ASSERT(n_blocks > 0); + GGML_ASSERT(state_read_idxs != nullptr); + + ggml_tensor * kv = ggml_get_rows(ctx0, kv_state, state_read_idxs); + cb(kv, "hca_kv", il); + kv = ggml_reshape_3d(ctx0, kv, n_embd_head, llama_context::dsv4_runtime::HCA_RATIO, n_blocks); + llm.cb(kv, "hca_comp_kv_rows", il); + + ggml_tensor * score = ggml_get_rows(ctx0, score_state, state_read_idxs); + cb(score, "hca_score", il); + score = ggml_reshape_3d(ctx0, score, n_embd_head, llama_context::dsv4_runtime::HCA_RATIO, n_blocks); + llm.cb(score, "hca_comp_score_rows", il); + + ggml_tensor * values = ggml_cont(ctx0, ggml_permute(ctx0, kv, 1, 0, 2, 3)); + ggml_tensor * scores = ggml_cont(ctx0, ggml_permute(ctx0, score, 1, 0, 2, 3)); + ggml_tensor * weights = ggml_soft_max(ctx0, scores); + ggml_tensor * comp = ggml_mul(ctx0, values, weights); + comp = ggml_sum_rows(ctx0, comp); + comp = ggml_reshape_3d(ctx0, comp, comp->ne[1], comp->ne[2], comp->ne[3]); + //comp = ggml_cont(ctx0, ggml_permute(ctx0, comp, 1, 0, 2, 3)); + llm.cb(comp, "hca_comp_merge", il); + + comp = llm.llm_build_norm(ctx0, comp, llm.hparams, norm, nullptr, LLM_NORM_RMS, llm.cb, il); + llm.cb(comp, "hca_comp_norm", il); + + ggml_tensor * comp_nope = ggml_view_3d(ctx0, comp, n_embd_head_nope, 1, n_blocks, + ggml_row_size(comp->type, n_embd_head), + ggml_row_size(comp->type, n_embd_head), + 0); + ggml_tensor * comp_pe = ggml_view_3d(ctx0, comp, n_embd_head_rope, 1, n_blocks, + ggml_row_size(comp->type, n_embd_head), + ggml_row_size(comp->type, n_embd_head), + ggml_row_size(comp->type, n_embd_head_nope)); + comp_pe = ggml_rope_ext(ctx0, comp_pe, comp_pos, nullptr, n_embd_head_rope, llm.rope_type, llm.n_ctx_orig, + llm.hparams.dsv4_compress_rope_base, llm.freq_scale, llm.ext_factor, + dsv4_rope_attn_factor(llm.freq_scale, llm.ext_factor), llm.beta_fast, llm.beta_slow); + comp = ggml_concat(ctx0, comp_nope, comp_pe, 0); + llm.cb(comp, "hca_comp_out", il); + + return comp; +} + +static ggml_tensor * build_overlap_compressed_kv_from_state( + ggml_context * ctx0, + llm_build_context & llm, + ggml_tensor * kv_state, + ggml_tensor * score_state, + ggml_tensor * state_read_idxs, + ggml_tensor * comp_pos, + ggml_tensor * norm, + int64_t ratio, + int64_t n_embd_head, + int il, + const char * tag, const llm_build_cb & cb) { + const int64_t n_embd_head_rope = llm.hparams.n_rot; + const int64_t n_embd_head_nope = n_embd_head - n_embd_head_rope; + const int64_t n_blocks = comp_pos ? comp_pos->ne[0] : 0; + + GGML_ASSERT(n_blocks > 0); + GGML_ASSERT(state_read_idxs != nullptr); + + // Why do we need this? + kv_state = dsv4_append_zero_row(ctx0, kv_state, false); + score_state = dsv4_append_zero_row(ctx0, score_state, true); + + auto kv_state_prev = ggml_view_4d(ctx0, kv_state, n_embd_head, kv_state->ne[1], kv_state->ne[2], kv_state->ne[3], + kv_state->nb[1], kv_state->nb[2], kv_state->nb[3], 0); + auto kv_state_cur = ggml_view_4d(ctx0, kv_state, n_embd_head, kv_state->ne[1], kv_state->ne[2], kv_state->ne[3], + kv_state->nb[1], kv_state->nb[2], kv_state->nb[3], ggml_row_size(kv_state->type, n_embd_head)); + auto score_state_prev = ggml_view_4d(ctx0, score_state, n_embd_head, score_state->ne[1], score_state->ne[2], score_state->ne[3], + score_state->nb[1], score_state->nb[2], score_state->nb[3], 0); + auto score_state_cur = ggml_view_4d(ctx0, score_state, n_embd_head, score_state->ne[1], score_state->ne[2], score_state->ne[3], + score_state->nb[1], score_state->nb[2], score_state->nb[3], ggml_row_size(score_state->type, n_embd_head)); + + ggml_tensor * prev_idxs = dsv4_view_1d(ctx0, state_read_idxs, ratio * n_blocks, 0); + ggml_tensor * cur_idxs = dsv4_view_1d(ctx0, state_read_idxs, ratio * n_blocks, ratio * n_blocks); + + //ggml_tensor * kv_prev = ggml_get_rows(ctx0, kv_state, prev_idxs); + //kv_prev = ggml_cont(ctx0, ggml_view_2d(ctx0, kv_prev, n_embd_head, ratio * n_blocks, kv_prev->nb[1], 0)); + ggml_tensor * kv_prev = ggml_get_rows(ctx0, kv_state_prev, prev_idxs); + cb(kv_prev, tag, il); + kv_prev = ggml_reshape_3d(ctx0, kv_prev, n_embd_head, ratio, n_blocks); + + //ggml_tensor * score_prev = ggml_get_rows(ctx0, score_state, prev_idxs); + //score_prev = ggml_cont(ctx0, ggml_view_2d(ctx0, score_prev, n_embd_head, ratio * n_blocks, score_prev->nb[1], 0)); + ggml_tensor * score_prev = ggml_get_rows(ctx0, score_state_prev, prev_idxs); + cb(score_prev, tag, il); + score_prev = ggml_reshape_3d(ctx0, score_prev, n_embd_head, ratio, n_blocks); + + //ggml_tensor * kv_cur = ggml_get_rows(ctx0, kv_state, cur_idxs); + //kv_cur = ggml_cont(ctx0, ggml_view_2d(ctx0, kv_cur, n_embd_head, ratio * n_blocks, kv_cur->nb[1], + // ggml_row_size(kv_cur->type, n_embd_head))); + ggml_tensor * kv_cur = ggml_get_rows(ctx0, kv_state_cur, cur_idxs); + cb(kv_cur, tag, il); + kv_cur = ggml_reshape_3d(ctx0, kv_cur, n_embd_head, ratio, n_blocks); + + //ggml_tensor * score_cur = ggml_get_rows(ctx0, score_state, cur_idxs); + //score_cur = ggml_cont(ctx0, ggml_view_2d(ctx0, score_cur, n_embd_head, ratio * n_blocks, score_cur->nb[1], + // ggml_row_size(score_cur->type, n_embd_head))); + ggml_tensor * score_cur = ggml_get_rows(ctx0, score_state_cur, cur_idxs); + cb(score_cur, tag, il); + score_cur = ggml_reshape_3d(ctx0, score_cur, n_embd_head, ratio, n_blocks); + + ggml_tensor * values = dsv4_concat_named(ctx0, kv_prev, kv_cur, 1, "dsv4_comp_values"); + ggml_tensor * scores = dsv4_concat_named(ctx0, score_prev, score_cur, 1, "dsv4_comp_scores"); + values = ggml_cont(ctx0, ggml_permute(ctx0, values, 1, 0, 2, 3)); + scores = ggml_cont(ctx0, ggml_permute(ctx0, scores, 1, 0, 2, 3)); + + ggml_tensor * weights = ggml_soft_max(ctx0, scores); + ggml_tensor * comp = ggml_mul(ctx0, values, weights); + comp = ggml_sum_rows(ctx0, comp); + //comp = ggml_cont(ctx0, ggml_permute(ctx0, comp, 1, 0, 2, 3)); + comp = ggml_reshape_3d(ctx0, comp, comp->ne[1], comp->ne[2], comp->ne[3]); + llm.cb(comp, tag, il); + + comp = llm.llm_build_norm(ctx0, comp, llm.hparams, norm, nullptr, LLM_NORM_RMS, llm.cb, il); + llm.cb(comp, tag, il); + + ggml_tensor * comp_nope = ggml_view_3d(ctx0, comp, n_embd_head_nope, 1, n_blocks, + ggml_row_size(comp->type, n_embd_head), + ggml_row_size(comp->type, n_embd_head), + 0); + ggml_tensor * comp_pe = ggml_view_3d(ctx0, comp, n_embd_head_rope, 1, n_blocks, + ggml_row_size(comp->type, n_embd_head), + ggml_row_size(comp->type, n_embd_head), + ggml_row_size(comp->type, n_embd_head_nope)); + comp_pe = ggml_rope_ext(ctx0, comp_pe, comp_pos, nullptr, n_embd_head_rope, llm.rope_type, llm.n_ctx_orig, + llm.hparams.dsv4_compress_rope_base, llm.freq_scale, llm.ext_factor, + dsv4_rope_attn_factor(llm.freq_scale, llm.ext_factor), llm.beta_fast, llm.beta_slow); + comp = ggml_concat(ctx0, comp_nope, comp_pe, 0); + llm.cb(comp, tag, il); + + return comp; +} + +static ggml_tensor * build_top_k_mask( + ggml_context * ctx0, + ggml_tensor * kq_mask, + ggml_tensor * top_k) { + if (!ggml_is_contiguous(kq_mask)) { + kq_mask = ggml_cont(ctx0, kq_mask); + } + if (top_k->ne[0] <= kq_mask->ne[0] && top_k->ne[1] <= kq_mask->ne[1] && top_k->ne[2] == kq_mask->ne[2] && top_k->ne[3] == kq_mask->ne[3]) { + return ggml_indexer_mask(ctx0, kq_mask, top_k); + } + ggml_tensor * kq_mask_all = ggml_fill(ctx0, kq_mask, -INFINITY); + //ggml_tensor * kq_mask_top_k = ggml_blend(ctx0, kq_mask_all, top_k, 0.0f); + // Fo siome reason the above is not faster than this + kq_mask_all = ggml_view_4d(ctx0, kq_mask_all, 1, kq_mask_all->ne[0], kq_mask_all->ne[1], kq_mask_all->ne[3], + kq_mask_all->nb[0], kq_mask_all->nb[1], kq_mask_all->nb[2], 0); + + ggml_tensor * top_k_3d = ggml_view_4d(ctx0, top_k, top_k->ne[0], top_k->ne[1], top_k->ne[3], 1, + top_k->nb[1], top_k->nb[2], top_k->ne[3]*top_k->nb[3], 0); + + ggml_tensor * zeros = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, 1, top_k_3d->ne[0], top_k_3d->ne[1], top_k_3d->ne[2]); + zeros = ggml_fill(ctx0, zeros, 0.0f); + + ggml_tensor * kq_mask_top_k = ggml_set_rows(ctx0, kq_mask_all, zeros, top_k_3d); + kq_mask_top_k = ggml_view_4d(ctx0, kq_mask_top_k, + kq_mask_top_k->ne[1], kq_mask_top_k->ne[2], 1, kq_mask_top_k->ne[3], + kq_mask_top_k->nb[2], kq_mask_top_k->nb[3], kq_mask_top_k->nb[3], 0); + return ggml_add(ctx0, kq_mask_top_k, kq_mask); +} + +static ggml_tensor * dsv4_build_lid_top_k_shared( + ggml_context * ctx0, + ggml_tensor * indexer_k, + ggml_tensor * indexer_q, + ggml_tensor * indexer_weights, + ggml_tensor * indexer_mask, + int n_top_k, const llm_build_cb & cb) { + const int64_t n_stream = indexer_k->ne[3]; + const int64_t n_tokens = indexer_q->ne[1]; + + if (n_stream <= 0 || indexer_k->ne[2] != 1 || indexer_q->ne[3] != n_stream || + indexer_weights->ne[3] != n_stream || indexer_mask->ne[1] < n_tokens || + indexer_mask->ne[3] < n_stream) { + return nullptr; + } + + ggml_tensor * selected = nullptr; + for (int64_t s = 0; s < n_stream; ++s) { + ggml_tensor * k = ggml_view_2d(ctx0, indexer_k, + indexer_k->ne[0], indexer_k->ne[1], indexer_k->nb[1], s*indexer_k->nb[3]); + ggml_tensor * q = ggml_view_3d(ctx0, indexer_q, + indexer_q->ne[0], indexer_q->ne[1], indexer_q->ne[2], + indexer_q->nb[1], indexer_q->nb[2], s*indexer_q->nb[3]); + q = ggml_permute(ctx0, q, 0, 2, 1, 3); + + ggml_tensor * w = ggml_view_2d(ctx0, indexer_weights, + indexer_weights->ne[0], indexer_weights->ne[1], indexer_weights->nb[1], + s*indexer_weights->nb[3]); + ggml_tensor * mask = ggml_view_2d(ctx0, indexer_mask, + indexer_mask->ne[0], n_tokens, indexer_mask->nb[1], + s*n_tokens*indexer_mask->nb[1]); + + ggml_tensor * cur = ggml_indexer_topk(ctx0, k, q, w, mask, + GGML_UNARY_OP_RELU, n_top_k); + if (selected) { + selected = ggml_concat(ctx0, selected, cur, 1); + cb(selected, "top_k", s); + } else { + selected = cur; + } + //selected = selected == nullptr ? cur : ggml_concat(ctx0, selected, cur, 1); + } + + return selected == nullptr ? nullptr : ggml_cont(ctx0, selected); +} + +static ggml_tensor * dsv4_build_lid_top_k( + ggml_context * ctx0, + llm_build_context & llm, + ggml_tensor * qr, + ggml_tensor * cur, + ggml_tensor * inp_pos, + int il, ggml_cgraph * gf, const llm_build_cb & cb) { + const auto & hparams = llm.hparams; + const auto & layer = llm.model.layers[il]; + const int64_t n_embd_indexer_head = hparams.indexer_head_size; + const int64_t n_embd_indexer_head_rope = hparams.n_rot; + const int64_t n_embd_indexer_head_nope = n_embd_indexer_head - n_embd_indexer_head_rope; + const int64_t n_indexer_head = hparams.indexer_n_head; + const int64_t n_tokens = cur->ne[1]; + const int64_t n_lid = llm.lctx.dsv4.lid_plan.n_kv; + const int hadamard_block = llama_model::hadamard_size((int) n_embd_indexer_head); + + GGML_ASSERT(n_embd_indexer_head >= n_embd_indexer_head_rope); + GGML_ASSERT(n_lid > 0); + GGML_ASSERT(hadamard_block > 0); + GGML_ASSERT(n_embd_indexer_head % hadamard_block == 0); + + ggml_tensor * indexer_q = llm.llm_build_lora_mm(llm.lctx, ctx0, layer.indexer_attn_q_b, qr); + llm.cb(indexer_q, "lid_q", il); + indexer_q = ggml_reshape_3d(ctx0, indexer_q, n_embd_indexer_head, n_indexer_head, n_tokens); + + ggml_tensor * indexer_q_nope = ggml_view_3d(ctx0, indexer_q, n_embd_indexer_head_nope, n_indexer_head, n_tokens, + ggml_row_size(indexer_q->type, n_embd_indexer_head), + ggml_row_size(indexer_q->type, n_embd_indexer_head) * n_indexer_head, + 0); + ggml_tensor * indexer_q_pe = ggml_view_3d(ctx0, indexer_q, n_embd_indexer_head_rope, n_indexer_head, n_tokens, + ggml_row_size(indexer_q->type, n_embd_indexer_head), + ggml_row_size(indexer_q->type, n_embd_indexer_head) * n_indexer_head, + ggml_row_size(indexer_q->type, n_embd_indexer_head_nope)); + indexer_q_pe = ggml_rope_ext(ctx0, indexer_q_pe, inp_pos, nullptr, n_embd_indexer_head_rope, + llm.rope_type, llm.n_ctx_orig, + hparams.dsv4_compress_rope_base, llm.freq_scale, + llm.ext_factor, dsv4_rope_attn_factor(llm.freq_scale, llm.ext_factor), llm.beta_fast, llm.beta_slow); + indexer_q = ggml_concat(ctx0, indexer_q_nope, indexer_q_pe, 0); + llm.cb(indexer_q, "indexer_q", il); + GGML_ASSERT(indexer_q->ne[0] % hadamard_block == 0); + indexer_q = ggml_hadamard(ctx0, indexer_q, hadamard_block); + llm.cb(indexer_q, "lid_q_hadamard", il); + + ggml_tensor * indexer_weights = llm.llm_build_lora_mm(llm.lctx, ctx0, layer.indexer_proj, cur); + llm.cb(indexer_weights, "lid_weights", il); + indexer_weights = ggml_scale(ctx0, indexer_weights, 1.0f / std::sqrt(float(n_embd_indexer_head * n_indexer_head))); + + ggml_tensor * indexer_k = dsv4_comp_get_k(ctx0, + llm.lctx.dsv4.cache.lid_k[(size_t) il], + llm.lctx.dsv4.lid_ctx, + n_embd_indexer_head, + llm.lctx.dsv4.cache.lid_k[(size_t) il]->ne[1]/std::max(1, llm.lctx.dsv4.cache.n_stream)); + llm.cb(indexer_k, "lid_k", il); + + const int64_t n_stream = std::max(1, indexer_k->ne[3]); + indexer_q = ggml_view_4d(ctx0, indexer_q, + indexer_q->ne[0], indexer_q->ne[1], indexer_q->ne[2] / n_stream, n_stream, + indexer_q->nb[1], indexer_q->nb[2], indexer_q->nb[3] / n_stream, 0); + indexer_weights = ggml_view_4d(ctx0, indexer_weights, + indexer_weights->ne[0], indexer_weights->ne[1] / n_stream, indexer_weights->ne[2], n_stream, + indexer_weights->nb[1], indexer_weights->nb[2] / n_stream, indexer_weights->nb[3] / n_stream, 0); + + indexer_q = ggml_permute(ctx0, indexer_q, 0, 2, 1, 3); + llm.cb(indexer_q, "lid_q_stream", il); + indexer_k = ggml_permute(ctx0, indexer_k, 0, 2, 1, 3); + llm.cb(indexer_k, "lid_k_stream", il); + + GGML_ASSERT(llm.lctx.dsv4.inputs.csa.kq_mask != nullptr); + ggml_tensor * lid_mask = dsv4_build_raw_mask_view(ctx0, + llm.lctx.dsv4.inputs.csa.kq_mask, nullptr, n_lid, n_tokens, n_stream, cb, il); + const uint32_t n_top_k = (uint32_t) std::min(n_lid, hparams.indexer_top_k); + if (llm.cparams.fused_idx_topk && n_lid > n_top_k) { + if (ggml_tensor * selected = dsv4_build_lid_top_k_shared(ctx0, + indexer_k, indexer_q, indexer_weights, lid_mask, (int) n_top_k, cb)) { + if (selected) { + ggml_build_forward_expand(gf, selected); + llm.cb(selected, "lid_top_k", il); + return selected; + } + } + } + + ggml_tensor * indexer_kq = ggml_mul_mat(ctx0, indexer_k, indexer_q); + llm.cb(indexer_kq, "lid_kq", il); + + indexer_kq = ggml_cont(ctx0, ggml_permute(ctx0, indexer_kq, 2, 1, 0, 3)); + llm.cb(indexer_kq, "lid_kq_perm", il); + + ggml_tensor * indexer_score = ggml_relu(ctx0, indexer_kq); + indexer_score = ggml_mul(ctx0, indexer_score, indexer_weights); + indexer_score = ggml_sum_rows(ctx0, indexer_score); + indexer_score = ggml_cont(ctx0, ggml_permute(ctx0, indexer_score, 2, 1, 0, 3)); + llm.cb(indexer_score, "lid_score", il); + + indexer_score = ggml_add(ctx0, indexer_score, lid_mask); + llm.cb(indexer_score, "lid_score_masked", il); + + ggml_tensor * top_k = ggml_cont(ctx0, ggml_top_k(ctx0, indexer_score, n_top_k)); + llm.cb(top_k, "lid_top_k", il); + + return top_k; +} + +ggml_cgraph * llm_build_context::build_deepseek4() { + ggml_cgraph * gf = new_graph_custom(); + + if (lctx.cparams.mtp_op_type != MTP_OP_NONE) { + GGML_ABORT("DeepSeek4 MTP execution is not implemented"); + } + + //printf("================================================================= %s\n", __func__); + + const int64_t n_embd_head = hparams.n_embd_head_k(0); + const int64_t n_embd_head_rope = hparams.n_rot; + const int64_t n_embd_head_nope = n_embd_head - n_embd_head_rope; + const int64_t hc = hparams.dsv4_hc_mult; + + GGML_ASSERT(n_embd_head == hparams.n_embd_head_v(0)); + GGML_ASSERT(n_embd_head_nope > 0); + + dsv4_new_i32_input(ctx0, &lctx.dsv4.inputs.raw_k_write_src_idxs, (int64_t) lctx.dsv4.raw.write_src_idxs.size(), "dsv4_raw_k_write_src_idxs"); + dsv4_new_i32_input(ctx0, &lctx.dsv4.inputs.raw_k_write_idxs, (int64_t) lctx.dsv4.raw.write_dst_idxs.size(), "dsv4_raw_k_write_idxs"); + dsv4_new_i32_input(ctx0, &lctx.dsv4.inputs.raw_k_read_idxs, (int64_t) lctx.dsv4.raw.read_dst_idxs.size(), "dsv4_raw_k_read_idxs"); + dsv4_build_plan_inputs(ctx0, lctx.dsv4.inputs.csa, lctx.dsv4.csa_plan, "dsv4_csa", n_tokens, true, lctx.cparams.flash_attn); + dsv4_build_plan_inputs(ctx0, lctx.dsv4.inputs.hca, lctx.dsv4.hca_plan, "dsv4_hca", n_tokens, true, lctx.cparams.flash_attn); + dsv4_build_plan_inputs(ctx0, lctx.dsv4.inputs.lid, lctx.dsv4.lid_plan, "dsv4_lid", n_tokens, false, lctx.cparams.flash_attn); + + ggml_tensor * inp = llm_build_inp_embd(ctx0, lctx, hparams, batch, model.tok_embd, cb); + ggml_tensor * inp_pos = build_inp_pos(); + ggml_tensor * KQ_mask = hparams.n_swa > 0 ? build_inp_KQ_mask_swa() : build_inp_KQ_mask(); + ggml_tensor * inpL = ggml_reshape_3d(ctx0, inp, n_embd, 1, n_tokens); + inpL = ggml_repeat_4d(ctx0, inpL, n_embd, hc, n_tokens, 1); + cb(inpL, "hc_init", -1); + + for (int il = 0; il < n_layer; ++il) { + ggml_tensor * residual = inpL; + ggml_tensor * post = nullptr; + ggml_tensor * comb = nullptr; + + ggml_tensor * cur = build_hc_pre(ctx0, *this, hparams, n_embd, hparams.f_norm_rms_eps, + inpL, + model.layers[il].hc_attn_fn, + model.layers[il].hc_attn_scale, + model.layers[il].hc_attn_base, + &post, &comb, cb, il); + cb(cur, "hc_attn_pre", il); + + cur = llm_build_norm(ctx0, cur, hparams, model.layers[il].attn_norm, nullptr, LLM_NORM_RMS, cb, il); + cb(cur, "attn_norm", il); + + ggml_tensor * qr = llm_build_lora_mm(lctx, ctx0, model.layers[il].wq_a, cur); + cb(qr, "qr", il); + + qr = llm_build_norm(ctx0, qr, hparams, model.layers[il].attn_q_a_norm, nullptr, LLM_NORM_RMS, cb, il); + cb(qr, "qr_norm", il); + + const int64_t ratio = hparams.dsv4_compress_ratios[(size_t) il]; + const bool use_compress_rope = ratio != 0; + const float freq_base_l = use_compress_rope ? hparams.dsv4_compress_rope_base : freq_base; + const float freq_scale_l = use_compress_rope ? freq_scale : 1.0f; + const float ext_factor_l = use_compress_rope ? ext_factor : 0.0f; + const float attn_factor_l = dsv4_rope_attn_factor(freq_scale_l, ext_factor_l); + const float beta_fast_l = use_compress_rope ? beta_fast : 0.0f; + const float beta_slow_l = use_compress_rope ? beta_slow : 0.0f; + const int32_t n_ctx_orig_l = use_compress_rope ? n_ctx_orig : 0; + + ggml_tensor * q = llm_build_lora_mm(lctx, ctx0, model.layers[il].wq_b, qr); + cb(q, "q_b", il); + q = ggml_reshape_3d(ctx0, q, n_embd_head, n_head, n_tokens); + q = ggml_rms_norm(ctx0, q, hparams.f_norm_rms_eps); + cb(q, "q_b", il); + + ggml_tensor * q_nope = ggml_view_3d(ctx0, q, n_embd_head_nope, n_head, n_tokens, + ggml_row_size(q->type, n_embd_head), + ggml_row_size(q->type, n_embd_head) * n_head, + 0); + ggml_tensor * q_pe = ggml_view_3d(ctx0, q, n_embd_head_rope, n_head, n_tokens, + ggml_row_size(q->type, n_embd_head), + ggml_row_size(q->type, n_embd_head) * n_head, + ggml_row_size(q->type, n_embd_head_nope)); + q_pe = ggml_rope_ext(ctx0, q_pe, inp_pos, nullptr, n_embd_head_rope, rope_type, n_ctx_orig_l, + freq_base_l, freq_scale_l, ext_factor_l, attn_factor_l, beta_fast_l, beta_slow_l); + cb(q_pe, "q_pe", il); + q = ggml_concat(ctx0, q_nope, q_pe, 0); + cb(q, "q", il); + + ggml_tensor * kv = llm_build_lora_mm(lctx, ctx0, model.layers[il].wkv_latent, cur); + cb(kv, "wkv", il); + kv = llm_build_norm(ctx0, kv, hparams, model.layers[il].attn_kv_norm, nullptr, LLM_NORM_RMS, cb, il); + kv = ggml_reshape_3d(ctx0, kv, n_embd_head, 1, n_tokens); + cb(kv, "kv_norm", il); + + ggml_tensor * kv_nope = ggml_view_3d(ctx0, kv, n_embd_head_nope, 1, n_tokens, + ggml_row_size(kv->type, n_embd_head), + ggml_row_size(kv->type, n_embd_head), + 0); + ggml_tensor * kv_pe = ggml_view_3d(ctx0, kv, n_embd_head_rope, 1, n_tokens, + ggml_row_size(kv->type, n_embd_head), + ggml_row_size(kv->type, n_embd_head), + ggml_row_size(kv->type, n_embd_head_nope)); + kv_pe = ggml_rope_ext(ctx0, kv_pe, inp_pos, nullptr, n_embd_head_rope, rope_type, n_ctx_orig_l, + freq_base_l, freq_scale_l, ext_factor_l, attn_factor_l, beta_fast_l, beta_slow_l); + cb(kv_pe, "kv_pe", il); + kv = ggml_concat(ctx0, kv_nope, kv_pe, 0); + cb(kv, "kv", il); + + if (cparams.k_cache_hadamard) { + if (int block_size = lctx.model.hadamard_size_k(il); block_size > 0) { + q = ggml_hadamard(ctx0, q, block_size); + kv = ggml_hadamard(ctx0, kv, block_size); + cb(q, "q_hadamard", il); + cb(kv, "kv_hadamard", il); + } + } + cb(kv, "dsv4_raw_k_before_write", il); + const float kq_scale = 1.0f / std::sqrt(float(n_embd_head)); + + ggml_tensor * hca_state_kv = nullptr; + ggml_tensor * hca_state_score = nullptr; + if (ratio == llama_context::dsv4_runtime::HCA_RATIO && lctx.dsv4.inputs.hca.state_pos != nullptr && lctx.dsv4.hca_plan.state_pos.size() > 0) { + hca_state_kv = llm_build_lora_mm(lctx, ctx0, model.layers[il].attn_comp_wkv, cur); + cb(hca_state_kv, "hca_state_kv", il); + hca_state_score = llm_build_lora_mm(lctx, ctx0, model.layers[il].attn_comp_wgate, cur); + cb(hca_state_score, "hca_state_score", il); + ggml_tensor * ape_rows = ggml_get_rows(ctx0, model.layers[il].attn_comp_ape, lctx.dsv4.inputs.hca.state_pos); + cb(ape_rows, "ape", il); + hca_state_score = ggml_add(ctx0, hca_state_score, ape_rows); + cb(hca_state_kv, "hca_state_kv", il); + cb(hca_state_score, "hca_state_score", il); + } + + if (ratio == llama_context::dsv4_runtime::CSA_RATIO && lctx.dsv4.inputs.csa.state_pos != nullptr && lctx.dsv4.csa_plan.state_pos.size() > 0) { + ggml_tensor * csa_state_kv = llm_build_lora_mm(lctx, ctx0, model.layers[il].attn_comp_wkv, cur); + cb(csa_state_kv, "csa_state_kv", il); + ggml_tensor * csa_state_score = llm_build_lora_mm(lctx, ctx0, model.layers[il].attn_comp_wgate, cur); + cb(csa_state_score, "csa_state_score", il); + ggml_tensor * csa_ape_rows = ggml_get_rows(ctx0, model.layers[il].attn_comp_ape, lctx.dsv4.inputs.csa.state_pos); + cb(csa_ape_rows, "csa_ape", il); + csa_state_score = ggml_add(ctx0, csa_state_score, csa_ape_rows); + ggml_tensor * csa_dep = nullptr; + + if (lctx.dsv4.inputs.csa.state_write_idxs != nullptr && lctx.dsv4.csa_plan.state_write_idxs.size() > 0) { + ggml_tensor * csa_source_kv = dsv4_concat_named(ctx0, lctx.dsv4.cache.csa_state_kv[(size_t) il], csa_state_kv, 1, "dsv4_csa_source_kv"); + ggml_tensor * csa_source_score = dsv4_concat_named(ctx0, lctx.dsv4.cache.csa_state_score[(size_t) il], csa_state_score, 1, "dsv4_csa_source_score"); + ggml_tensor * csa_comp = build_overlap_compressed_kv_from_state( + ctx0, *this, + csa_source_kv, csa_source_score, + lctx.dsv4.inputs.csa.state_read_idxs, + lctx.dsv4.inputs.csa.state_write_pos, + model.layers[il].attn_comp_norm, + llama_context::dsv4_runtime::CSA_RATIO, + n_embd_head, + il, + "csa_state_compress", cb); + ggml_tensor * csa_comp_2d = ggml_reshape_2d(ctx0, csa_comp, n_embd_head, lctx.dsv4.inputs.csa.state_write_idxs->ne[0]); + ggml_tensor * csa_write = dsv4_comp_cpy_k(ctx0, lctx.dsv4.cache.csa_k[(size_t) il], csa_comp_2d, lctx.dsv4.inputs.csa.state_write_idxs, n_embd_head); + ggml_build_forward_expand(gf, csa_write); + cb(csa_write, "dsv4_csa_k_write", il); + csa_dep = csa_comp; + } + + if (csa_dep) { + ggml_build_forward_expand(gf, csa_dep); + } + ggml_tensor * csa_persist_kv = ggml_get_rows(ctx0, csa_state_kv, lctx.dsv4.inputs.csa.state_persist_src_idxs); + cb(csa_persist_kv, "csa_persist_kv", il); + ggml_tensor * csa_persist_score = ggml_get_rows(ctx0, csa_state_score, lctx.dsv4.inputs.csa.state_persist_src_idxs); + cb(csa_persist_score, "csa_persist_score", il); + ggml_tensor * csa_state_kv_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.csa_state_kv[(size_t) il], csa_persist_kv, lctx.dsv4.inputs.csa.state_persist_dst_idxs); + ggml_tensor * csa_state_score_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.csa_state_score[(size_t) il], csa_persist_score, lctx.dsv4.inputs.csa.state_persist_dst_idxs); + ggml_build_forward_expand(gf, csa_state_kv_write); + ggml_build_forward_expand(gf, csa_state_score_write); + cb(csa_state_kv_write, "dsv4_csa_k_state_persist", il); + cb(csa_state_score_write, "dsv4_csa_score_state_persist", il); + + ggml_tensor * lid_state_kv = llm_build_lora_mm(lctx, ctx0, model.layers[il].indexer_comp_wkv, cur); + cb(lid_state_kv, "lid_state_kv", il); + ggml_tensor * lid_state_score = llm_build_lora_mm(lctx, ctx0, model.layers[il].indexer_comp_wgate, cur); + cb(lid_state_score, "lid_state_score", il); + ggml_tensor * lid_ape_rows = ggml_get_rows(ctx0, model.layers[il].indexer_comp_ape, lctx.dsv4.inputs.lid.state_pos); + cb(lid_ape_rows, "lid_ape", il); + lid_state_score = ggml_add(ctx0, lid_state_score, lid_ape_rows); + ggml_tensor * lid_dep = nullptr; + + if (lctx.dsv4.inputs.lid.state_write_idxs != nullptr && lctx.dsv4.lid_plan.state_write_idxs.size() > 0) { + ggml_tensor * lid_source_kv = dsv4_concat_named(ctx0, lctx.dsv4.cache.lid_state_kv[(size_t) il], lid_state_kv, 1, "dsv4_lid_source_kv"); + ggml_tensor * lid_source_score = dsv4_concat_named(ctx0, lctx.dsv4.cache.lid_state_score[(size_t) il], lid_state_score, 1, "dsv4_lid_source_score"); + ggml_tensor * lid_comp = build_overlap_compressed_kv_from_state( + ctx0, *this, + lid_source_kv, lid_source_score, + lctx.dsv4.inputs.lid.state_read_idxs, + lctx.dsv4.inputs.lid.state_write_pos, + model.layers[il].indexer_comp_norm, + llama_context::dsv4_runtime::CSA_RATIO, + hparams.indexer_head_size, + il, + "lid_state_compress", cb); + const int hadamard_block = llama_model::hadamard_size((int) hparams.indexer_head_size); + GGML_ASSERT(hadamard_block > 0); + GGML_ASSERT(lid_comp->ne[0] % hadamard_block == 0); + lid_comp = ggml_hadamard(ctx0, lid_comp, hadamard_block); + cb(lid_comp, "lid_state_compress_hadamard", il); + ggml_tensor * lid_comp_2d = ggml_reshape_2d(ctx0, lid_comp, hparams.indexer_head_size, lctx.dsv4.inputs.lid.state_write_idxs->ne[0]); + ggml_tensor * lid_write = dsv4_comp_cpy_k(ctx0, lctx.dsv4.cache.lid_k[(size_t) il], lid_comp_2d, lctx.dsv4.inputs.lid.state_write_idxs, hparams.indexer_head_size); + ggml_build_forward_expand(gf, lid_write); + cb(lid_write, "dsv4_lid_k_write", il); + lid_dep = lid_comp; + } + + if (lid_dep) { + ggml_build_forward_expand(gf, lid_dep); + } + ggml_tensor * lid_persist_kv = ggml_get_rows(ctx0, lid_state_kv, lctx.dsv4.inputs.lid.state_persist_src_idxs); + cb(lid_persist_kv, "lid_persist_kv", il); + ggml_tensor * lid_persist_score = ggml_get_rows(ctx0, lid_state_score, lctx.dsv4.inputs.lid.state_persist_src_idxs); + cb(lid_persist_score, "lid_persist_score", il); + ggml_tensor * lid_state_kv_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.lid_state_kv[(size_t) il], lid_persist_kv, lctx.dsv4.inputs.lid.state_persist_dst_idxs); + ggml_tensor * lid_state_score_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.lid_state_score[(size_t) il], lid_persist_score, lctx.dsv4.inputs.lid.state_persist_dst_idxs); + ggml_build_forward_expand(gf, lid_state_kv_write); + ggml_build_forward_expand(gf, lid_state_score_write); + cb(lid_state_kv_write, "dsv4_lid_k_state_persist", il); + cb(lid_state_score_write, "dsv4_lid_score_state_persist", il); + } + + if (ratio == llama_context::dsv4_runtime::HCA_RATIO && hca_state_kv != nullptr && hca_state_score != nullptr) { + ggml_tensor * hca_dep = nullptr; + if (lctx.dsv4.inputs.hca.state_write_idxs != nullptr && lctx.dsv4.hca_plan.state_write_idxs.size() > 0) { + ggml_tensor * hca_source_kv = dsv4_concat_named(ctx0, lctx.dsv4.cache.hca_state_kv[(size_t) il], hca_state_kv, 1, "dsv4_hca_source_kv"); + ggml_tensor * hca_source_score = dsv4_concat_named(ctx0, lctx.dsv4.cache.hca_state_score[(size_t) il], hca_state_score, 1, "dsv4_hca_source_score"); + ggml_tensor * hca_comp = build_hca_compressed_kv_from_state( + ctx0, *this, + hca_source_kv, hca_source_score, + lctx.dsv4.inputs.hca.state_read_idxs, + lctx.dsv4.inputs.hca.state_write_pos, + model.layers[il].attn_comp_norm, + n_embd_head, + cb, il); + ggml_tensor * hca_comp_2d = ggml_reshape_2d(ctx0, hca_comp, n_embd_head, lctx.dsv4.inputs.hca.state_write_idxs->ne[0]); + ggml_tensor * hca_write = dsv4_comp_cpy_k(ctx0, lctx.dsv4.cache.hca_k[(size_t) il], hca_comp_2d, lctx.dsv4.inputs.hca.state_write_idxs, n_embd_head); + ggml_build_forward_expand(gf, hca_write); + cb(hca_write, "dsv4_hca_k_write", il); + hca_dep = hca_comp; + } + + if (hca_dep) { + ggml_build_forward_expand(gf, hca_dep); + } + ggml_tensor * hca_persist_kv = ggml_get_rows(ctx0, hca_state_kv, lctx.dsv4.inputs.hca.state_persist_src_idxs); + ggml_tensor * hca_persist_score = ggml_get_rows(ctx0, hca_state_score, lctx.dsv4.inputs.hca.state_persist_src_idxs); + cb(hca_persist_kv, "hca_persist_kv", il); + cb(hca_persist_score, "hca_persist_score", il); + ggml_tensor * hca_state_kv_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.hca_state_kv[(size_t) il], hca_persist_kv, lctx.dsv4.inputs.hca.state_persist_dst_idxs); + ggml_tensor * hca_state_score_write = dsv4_comp_state_cpy(ctx0, lctx.dsv4.cache.hca_state_score[(size_t) il], hca_persist_score, lctx.dsv4.inputs.hca.state_persist_dst_idxs); + ggml_build_forward_expand(gf, hca_state_kv_write); + ggml_build_forward_expand(gf, hca_state_score_write); + cb(hca_state_kv_write, "dsv4_hca_k_state_persist", il); + cb(hca_state_score_write, "dsv4_hca_score_state_persist", il); + } + + ggml_tensor * raw_k_write = nullptr; + if (hparams.n_head_kv(il) == 1 && lctx.dsv4.inputs.raw_k_write_idxs != nullptr) { + raw_k_write = dsv4_raw_cpy_k(&lctx, ctx0, kv_self.k_l[il], kv, lctx.dsv4.inputs.raw_k_write_src_idxs, lctx.dsv4.inputs.raw_k_write_idxs, gf, n_embd_head, cb, il); + if (raw_k_write != nullptr) { + cb(raw_k_write, "dsv4_raw_k_write", il); + } + } + if (raw_k_write == nullptr) { + llm_build_kv_store(lctx, ctx0, hparams, cparams, kv_self, gf, kv, nullptr, n_tokens, kv_head, cb, il); + } + if (il < (int64_t) kv_self.v_l.size() && kv_self.v_l[(size_t) il] != nullptr) { + llm_build_kv_store(lctx, ctx0, hparams, cparams, kv_self, gf, nullptr, kv, n_tokens, kv_head, cb, il); + } + + ggml_tensor * raw_k = nullptr; + if (hparams.n_head_kv(il) == 1 && lctx.dsv4.inputs.raw_k_read_idxs != nullptr) { + raw_k = dsv4_raw_get_k(&lctx, ctx0, kv_self.k_l[il], lctx.dsv4.inputs.raw_k_read_idxs, n_embd_head, cb, il); + } + if (raw_k == nullptr) { + raw_k = ggml_view_3d(ctx0, kv_self.k_l[il], + n_embd_head, hparams.n_head_kv(il), n_kv, + ggml_row_size(kv_self.k_l[il]->type, n_embd_head), + ggml_row_size(kv_self.k_l[il]->type, n_embd_head) * hparams.n_head_kv(il), + 0); + } + cb(raw_k, "raw_k", il); + + const int64_t raw_kq_n_kv = raw_k != nullptr && lctx.dsv4.raw.n_kv > 0 + ? lctx.dsv4.raw.n_kv + : (raw_k != nullptr ? raw_k->ne[2] * raw_k->ne[3] : n_kv); + const int64_t raw_attn_n_kv = raw_kq_n_kv > 0 ? std::max(256, GGML_PAD(raw_kq_n_kv, 256)) : raw_kq_n_kv; + if (raw_k != nullptr && raw_k->ne[3] == 1) { + raw_k = dsv4_pad_raw_k_to(ctx0, raw_k, raw_attn_n_kv); + } + ggml_tensor * raw_mask = dsv4_build_raw_mask_view(ctx0, KQ_mask, + lctx.dsv4.inputs.raw_k_read_idxs, raw_kq_n_kv, n_tokens, raw_k->ne[3], cb, il); + cb(raw_mask, "raw_mask_view", il); + raw_mask = dsv4_pad_mask_tokens(ctx0, raw_mask, n_tokens); + raw_mask = dsv4_pad_raw_mask_to(ctx0, raw_mask, raw_attn_n_kv, n_tokens); + cb(raw_mask, "dsv4_raw_mask_padded", il); + ggml_tensor * attn = nullptr; + + if (ratio == llama_context::dsv4_runtime::CSA_RATIO && + lctx.dsv4.inputs.csa.kq_mask != nullptr && + lctx.dsv4.csa_plan.n_kv > 0 && + lctx.dsv4.lid_plan.n_kv > 0 && + !cparams.k_cache_hadamard) { + ggml_tensor * csa_k = dsv4_comp_get_k(ctx0, + lctx.dsv4.cache.csa_k[(size_t) il], + lctx.dsv4.csa_ctx, + n_embd_head, + lctx.dsv4.cache.csa_k[(size_t) il]->ne[1]/std::max(1, lctx.dsv4.cache.n_stream)); + ggml_tensor * top_k = dsv4_build_lid_top_k(ctx0, *this, qr, cur, inp_pos, il, gf, cb); + ggml_tensor * csa_mask = build_top_k_mask(ctx0, + dsv4_build_raw_mask_view(ctx0, lctx.dsv4.inputs.csa.kq_mask, nullptr, + lctx.dsv4.csa_plan.n_kv, n_tokens, csa_k->ne[3], cb, il), + top_k); + cb(csa_mask, "csa_mask", il); + const bool use_fattn = cparams.flash_attn; + if (use_fattn) { + csa_mask = dsv4_pad_mask_tokens(ctx0, csa_mask, n_tokens); + } + raw_k = dsv4_repeat_streams(ctx0, raw_k, csa_k->ne[3]); + if (!use_fattn) { + raw_mask = dsv4_build_raw_mask_view(ctx0, KQ_mask, + lctx.dsv4.inputs.raw_k_read_idxs, raw_kq_n_kv, n_tokens, csa_k->ne[3], cb, il); + raw_mask = dsv4_pad_raw_mask_to(ctx0, raw_mask, raw_attn_n_kv, n_tokens); + } + if (use_fattn && csa_mask->type != GGML_TYPE_F16) { + csa_mask = ggml_cast(ctx0, csa_mask, GGML_TYPE_F16); + } + if (raw_mask->type != csa_mask->type) { + raw_mask = ggml_cast(ctx0, raw_mask, csa_mask->type); + } + { + constexpr int k_fa_chunk = 256; + int n_swa = hparams.n_swa; + int ntokens = std::max(k_fa_chunk, int(q->ne[2])); + int nton = k_fa_chunk*((ntokens + n_swa + k_fa_chunk - 1)/k_fa_chunk); + int first = raw_k->ne[2] - nton; + if (first > 0) { + raw_k = ggml_view_4d(ctx0, raw_k, raw_k->ne[0], raw_k->ne[1], nton, raw_k->ne[3], + raw_k->nb[1], raw_k->nb[2], raw_k->nb[3], raw_k->nb[2]*first); + raw_mask = ggml_view_4d(ctx0, raw_mask, nton, raw_mask->ne[1], raw_mask->ne[2], raw_mask->ne[3], + raw_mask->nb[1], raw_mask->nb[2], raw_mask->nb[3], raw_mask->nb[0]*first); + } + } + ggml_tensor * k_all = ggml_concat(ctx0, raw_k, csa_k, 2); + //printf("k_all: %ld x %ld x %ld x %ld, raw_k: %ld x %ld x %ld x %ld, csa_k = %ld x %ld x %ld x %ld, q = %ld x %ld x %ld x %ld\n", + // k_all->ne[0], k_all->ne[1], k_all->ne[2], k_all->ne[3], + // raw_k->ne[0], raw_k->ne[1], raw_k->ne[2], raw_k->ne[3], + // csa_k->ne[0], csa_k->ne[1], csa_k->ne[2], csa_k->ne[3], + // q->ne[0], q->ne[1], q->ne[2], q->ne[3]); + ggml_tensor * kq_mask = ggml_concat(ctx0, raw_mask, csa_mask, 0); + ggml_tensor * kq_b = dsv4_build_kq_zero_bias(ctx0, cparams, kq_mask, q->ne[1]); + cb(csa_k, "csa_k", il); + cb(k_all, "csa_k_all", il); + cb(kq_mask, "csa_kq_mask", il); + int n_csa = hparams.n_swa + hparams.indexer_top_k; + attn = dsv4_build_attn(ctx0, hparams, cparams, q, k_all, k_all, kq_b, kq_mask, model.layers[il].attn_sinks, kq_scale, cb, il, n_csa, gf); + cb(attn, "attn_csa", il); + } else if (ratio == llama_context::dsv4_runtime::HCA_RATIO && + lctx.dsv4.inputs.hca.kq_mask != nullptr && + lctx.dsv4.hca_plan.n_kv > 0 && + std::any_of(lctx.dsv4.hca_plan.n_visible.begin(), lctx.dsv4.hca_plan.n_visible.end(), + [](int32_t n_visible) { return n_visible > 0; }) && + !cparams.k_cache_hadamard) { + ggml_tensor * hca_k = dsv4_comp_get_k(ctx0, + lctx.dsv4.cache.hca_k[(size_t) il], + lctx.dsv4.hca_ctx, + n_embd_head, + lctx.dsv4.cache.hca_k[(size_t) il]->ne[1]/std::max(1, lctx.dsv4.cache.n_stream)); + const bool use_fattn = cparams.flash_attn; + ggml_tensor * hca_mask = dsv4_build_raw_mask_view(ctx0, lctx.dsv4.inputs.hca.kq_mask, nullptr, + lctx.dsv4.hca_plan.n_kv, n_tokens, hca_k->ne[3], cb, il); + hca_mask = dsv4_pad_mask_tokens(ctx0, hca_mask, n_tokens); + if (use_fattn && hca_mask->type != GGML_TYPE_F16) { + hca_mask = ggml_cast(ctx0, hca_mask, GGML_TYPE_F16); + } + raw_k = dsv4_repeat_streams(ctx0, raw_k, hca_k->ne[3]); + if (raw_mask->type != hca_mask->type) { + raw_mask = ggml_cast(ctx0, raw_mask, hca_mask->type); + } + { + constexpr int k_fa_chunk = 256; + int n_swa = hparams.n_swa; + int ntokens = std::max(k_fa_chunk, int(q->ne[2])); + int nton = k_fa_chunk*((ntokens + n_swa + k_fa_chunk - 1)/k_fa_chunk); + int first = raw_k->ne[2] - nton; + if (first > 0) { + raw_k = ggml_view_4d(ctx0, raw_k, raw_k->ne[0], raw_k->ne[1], nton, raw_k->ne[3], + raw_k->nb[1], raw_k->nb[2], raw_k->nb[3], raw_k->nb[2]*first); + raw_mask = ggml_view_4d(ctx0, raw_mask, nton, raw_mask->ne[1], raw_mask->ne[2], raw_mask->ne[3], + raw_mask->nb[1], raw_mask->nb[2], raw_mask->nb[3], raw_mask->nb[0]*first); + } + } + ggml_tensor * k_all = ggml_concat(ctx0, raw_k, hca_k, 2); + ggml_tensor * kq_mask = ggml_concat(ctx0, raw_mask, hca_mask, 0); + ggml_tensor * kq_b = dsv4_build_kq_zero_bias(ctx0, cparams, kq_mask, q->ne[1]); + cb(hca_k, "hca_k", il); + cb(k_all, "hca_k_all", il); + cb(kq_mask, "hca_kq_mask", il); + int n_hca = (n_kv + llama_context::dsv4_runtime::HCA_RATIO - 1)/llama_context::dsv4_runtime::HCA_RATIO; + n_hca += hparams.n_swa; + attn = dsv4_build_attn(ctx0, hparams, cparams, q, k_all, k_all, kq_b, kq_mask, model.layers[il].attn_sinks, kq_scale, cb, il, n_hca, gf); + cb(attn, "attn_hca", il); + } else { + ggml_tensor * kq_b = dsv4_build_kq_zero_bias(ctx0, cparams, raw_mask, q->ne[1]); + attn = dsv4_build_attn(ctx0, hparams, cparams, q, raw_k, raw_k, kq_b, raw_mask, model.layers[il].attn_sinks, kq_scale, cb, il, -1, gf); + cb(attn, "attn_raw", il); + } + + attn = ggml_reshape_3d(ctx0, attn, n_embd_head, n_head, n_tokens); + ggml_tensor * attn_nope = ggml_view_3d(ctx0, attn, n_embd_head_nope, n_head, n_tokens, + ggml_row_size(attn->type, n_embd_head), + ggml_row_size(attn->type, n_embd_head) * n_head, + 0); + ggml_tensor * attn_pe = ggml_view_3d(ctx0, attn, n_embd_head_rope, n_head, n_tokens, + ggml_row_size(attn->type, n_embd_head), + ggml_row_size(attn->type, n_embd_head) * n_head, + ggml_row_size(attn->type, n_embd_head_nope)); + attn_pe = ggml_rope_back(ctx0, attn_pe, inp_pos, nullptr, n_embd_head_rope, rope_type, n_ctx_orig_l, + freq_base_l, freq_scale_l, ext_factor_l, attn_factor_l, beta_fast_l, beta_slow_l); + cb(attn_pe, "attn_derope", il); + attn = ggml_concat(ctx0, attn_nope, attn_pe, 0); + cb(attn, "attn", il); + + const int64_t o_group_dim = model.layers[il].wo_a->ne[0]; + const int64_t n_groups = (n_head * n_embd_head) / o_group_dim; + const int64_t o_lora_rank = model.layers[il].wo_b->ne[0] / n_groups; + + GGML_ASSERT((n_head * n_embd_head) % o_group_dim == 0); + GGML_ASSERT(model.layers[il].wo_b->ne[0] % n_groups == 0); + + attn = ggml_reshape_3d(ctx0, attn, o_group_dim, n_groups, n_tokens); + attn = ggml_permute(ctx0, attn, 0, 2, 1, 3); + + ggml_tensor * oa = ggml_mul_mat(ctx0, + ggml_reshape_3d(ctx0, model.layers[il].wo_a, model.layers[il].wo_a->ne[0], o_lora_rank, n_groups), + attn); + cb(oa, "attn_wo_a", il); + oa = ggml_permute(ctx0, oa, 0, 2, 1, 3); + if (n_tokens == 1) { + oa = ggml_reshape_2d(ctx0, oa, o_lora_rank * n_groups, n_tokens); + } else { + oa = ggml_cont_2d(ctx0, oa, o_lora_rank * n_groups, n_tokens); + } + + cur = llm_build_lora_mm(lctx, ctx0, model.layers[il].wo_b, oa); + cb(cur, "attn_out", il); + + inpL = build_mhc_post(cur, post, residual, comb, n_embd, hc, true); + cb(inpL, "hc_attn_post", il); + + residual = inpL; + cur = build_hc_pre(ctx0, *this, hparams, n_embd, hparams.f_norm_rms_eps, + inpL, + model.layers[il].hc_ffn_fn, + model.layers[il].hc_ffn_scale, + model.layers[il].hc_ffn_base, + &post, &comb, cb, il); + cb(cur, "hc_ffn_pre", il); + + cur = llm_build_norm(ctx0, cur, hparams, model.layers[il].ffn_norm, nullptr, LLM_NORM_RMS, cb, il); + cb(cur, "ffn_norm", il); + + if ((uint32_t) il < hparams.n_layer_dense_lead) { + cur = llm_build_ffn(ctx0, lctx, nullptr, cur, + model.layers[il].ffn_up, nullptr, nullptr, + model.layers[il].ffn_gate, nullptr, nullptr, + model.layers[il].ffn_down, nullptr, nullptr, + nullptr, + LLM_FFN_SILU, LLM_FFN_PAR, cb, il); + } else { + // DSV4 uses separate up and gate expert tensors. Do not silently + // select the fork-only merged gate path for another GGUF. + GGML_ASSERT(model.layers[il].ffn_up_gate_exps == nullptr && + "merged DSV4 MoE gate tensors use an unsupported layout"); + ggml_tensor * selected_experts = nullptr; + ggml_tensor * exp_probs_b = model.layers[il].ffn_exp_probs_b; + if ((uint32_t) il < hparams.dsv4_hash_layer_count) { + selected_experts = ggml_get_rows(ctx0, model.layers[il].ffn_gate_tid2eid, lctx.inp_tokens); + cb(selected_experts, "hashed_exps", il); + exp_probs_b = nullptr; + } + + // Hash layers carry an explicit fixed-width expert map. During + // warmup the generic graph reserves all experts, but this input + // still contains only the model's active expert IDs. + const int64_t moe_n_expert_used = selected_experts != nullptr + ? selected_experts->ne[0] + : n_expert_used; + + const int64_t dsv4_n_stream = std::max(1, lctx.dsv4.csa_ctx.graph_n_stream); + // Wide packed DSV4 fused/IQK MoE diverges above 1024 total tokens. + // Evaluate each active stream independently to preserve packed parity. + constexpr int64_t dsv4_moe_max_tokens = 1024; + + auto build_dsv4_moe = [&](ggml_tensor * moe_cur, + ggml_tensor * moe_exp_probs_b, + ggml_tensor * moe_selected_experts) { + return llm_build_moe_ffn(ctx0, lctx, moe_cur, + model.layers[il].ffn_gate_inp, + nullptr, + model.layers[il].ffn_up_exps, + nullptr, + model.layers[il].ffn_gate_exps, + nullptr, + model.layers[il].ffn_down_exps, + nullptr, + moe_exp_probs_b, + n_expert, moe_n_expert_used, + LLM_FFN_SILU, hparams.expert_weights_norm, + true, hparams.expert_weights_scale, + (enum llm_expert_gating_func_type) hparams.expert_gating_func, + cb, il, gf, false, model.layers[il].ffn_up_gate_exps, nullptr, nullptr, nullptr, + moe_selected_experts); + }; + + ggml_tensor * moe_out = nullptr; + if (dsv4_n_stream > 1 && cur->ne[1] > dsv4_moe_max_tokens && + cur->ne[1] % dsv4_n_stream == 0) { + const int64_t n_tokens_stream = cur->ne[1]/dsv4_n_stream; + auto stream_view = [&](ggml_tensor * tensor, int64_t stream) { + if (tensor == nullptr || tensor->ne[1] != cur->ne[1]) { + return tensor; + } + return ggml_view_2d(ctx0, tensor, tensor->ne[0], n_tokens_stream, + tensor->nb[1], stream*n_tokens_stream*tensor->nb[1]); + }; + + for (int64_t stream = 0; stream < dsv4_n_stream; ++stream) { + ggml_tensor * stream_result = build_dsv4_moe( + stream_view(cur, stream), + stream_view(exp_probs_b, stream), + stream_view(selected_experts, stream)); + moe_out = moe_out == nullptr ? stream_result : ggml_concat(ctx0, moe_out, stream_result, 1); + } + } else { + moe_out = build_dsv4_moe(cur, exp_probs_b, selected_experts); + } + cb(moe_out, "ffn_moe_out", il); + + ggml_tensor * ffn_shexp = llm_build_ffn(ctx0, lctx, nullptr, cur, + model.layers[il].ffn_up_shexp, nullptr, nullptr, + model.layers[il].ffn_gate_shexp, nullptr, nullptr, + model.layers[il].ffn_down_shexp, nullptr, nullptr, + nullptr, + LLM_FFN_SILU, LLM_FFN_PAR, cb, il); + cb(ffn_shexp, "ffn_shexp", il); + + cur = ggml_add(ctx0, moe_out, ffn_shexp); + } + + cb(cur, "ffn_out", il); + + inpL = build_mhc_post(cur, post, residual, comb, n_embd, hc, true); + inpL = lctx.cvec.apply_to(ctx0, inpL, il); + cb(inpL, "l_out", il); + } + + if (n_outputs != n_tokens) { + ggml_tensor * inp_out_ids = build_inp_out_ids(); + ggml_tensor * flat = ggml_reshape_2d(ctx0, inpL, n_embd*hc, n_tokens); + flat = ggml_get_rows(ctx0, flat, inp_out_ids); + inpL = ggml_reshape_3d(ctx0, flat, n_embd, hc, n_outputs); + } + + ggml_tensor * out = build_hc_head(ctx0, *this, hparams, n_embd, hparams.f_norm_rms_eps, + inpL, + model.hc_head_fn, + model.hc_head_scale, + model.hc_head_base); + cb(out, "hc_head", -1); + + if (model.output_norm != nullptr) { + out = llm_build_norm(ctx0, out, hparams, model.output_norm, nullptr, LLM_NORM_RMS, cb, -1); + cb(out, "result_norm", -1); + out = build_output(lctx, ctx0, out, model.output, nullptr, cb); + } else { + out = build_output(lctx, ctx0, out, model.output, nullptr, cb); + } + cb(out, "result_output", -1); + + ggml_build_forward_expand(gf, out); + + return gf; +} diff --git a/src/graphs/build_openpangu.cpp b/src/graphs/build_openpangu.cpp index b0d9ca20..6fb6a116 100644 --- a/src/graphs/build_openpangu.cpp +++ b/src/graphs/build_openpangu.cpp @@ -1046,10 +1046,11 @@ ggml_cgraph * llm_build_context::build_openpangu() { h_pre = ggml_add(ctx0, ggml_mul(ctx0, ggml_cont(ctx0, h_pre), a_pre), b_pre); // broadcast scalar + [S] h_pre = ggml_sigmoid(ctx0, h_pre); // [S,T] (+eps omitted, inert) - // combine: x[h,t] = sum_s h_pre[s,t] * R[h,s,t] + //// combine: x[h,t] = sum_s h_pre[s,t] * R[h,s,t] ggml_tensor * hpre3 = ggml_reshape_3d(ctx0, h_pre, 1, S, n_tokens); - ggml_tensor * weighted = ggml_mul(ctx0, Rin, hpre3); // [H,S,T] - ggml_tensor * x = ggml_reshape_2d(ctx0, ggml_sum_rows_ext(ctx0, weighted, 1), n_embd, n_tokens); + auto x = ggml_mul_multi_add(ctx0, Rin, hpre3); + //ggml_tensor * weighted = ggml_mul(ctx0, Rin, hpre3); // [H,S,T] + //ggml_tensor * x = ggml_reshape_2d(ctx0, ggml_sum_rows_ext(ctx0, weighted, 1), n_embd, n_tokens); ggml_build_forward_expand(gf, x); *h_post_out = ggml_cont(ctx0, h_post); diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp index d9f6f804..bf07387f 100644 --- a/src/llama-arch.cpp +++ b/src/llama-arch.cpp @@ -53,6 +53,7 @@ static const std::map LLM_ARCH_NAMES = { { LLM_ARCH_OPENELM, "openelm" }, { LLM_ARCH_ARCTIC, "arctic" }, { LLM_ARCH_DEEPSEEK2, "deepseek2" }, + { LLM_ARCH_DEEPSEEK4, "deepseek4" }, { LLM_ARCH_CHATGLM, "chatglm" }, { LLM_ARCH_GLM4, "glm4" }, { LLM_ARCH_GLM4_MOE, "glm4moe" }, @@ -191,12 +192,22 @@ static const std::map LLM_KV_NAMES = { { LLM_KV_ATTENTION_INDEXER_HEAD_COUNT, "%s.attention.indexer.head_count" }, { LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, "%s.attention.indexer.key_length" }, { LLM_KV_ATTENTION_INDEXER_TOP_K, "%s.attention.indexer.top_k" }, + { LLM_KV_ATTENTION_OUTPUT_GROUP_COUNT, "%s.attention.output_group_count" }, + { LLM_KV_ATTENTION_OUTPUT_LORA_RANK, "%s.attention.output_lora_rank" }, + { LLM_KV_ATTENTION_COMPRESS_ROPE_FREQ_BASE,"%s.attention.compress_rope_freq_base"}, + { LLM_KV_ATTENTION_COMPRESS_RATIOS, "%s.attention.compress_ratios" }, { LLM_KV_FULL_ATTENTION_INTERVAL, "%s.full_attention_interval" }, { LLM_KV_ATTENTION_SHARED_KV_LAYERS, "%s.attention.shared_kv_layers" }, { LLM_KV_ATTENTION_KEY_LENGTH_SWA, "%s.attention.key_length_swa" }, { LLM_KV_ATTENTION_VALUE_LENGTH_SWA, "%s.attention.value_length_swa" }, { LLM_KV_ATTENTION_VALUE_SCALE, "%s.attention.value_scale" }, + { LLM_KV_HYPER_CONNECTION_COUNT, "%s.hyper_connection.count" }, + { LLM_KV_HYPER_CONNECTION_SINKHORN_ITERATIONS, "%s.hyper_connection.sinkhorn_iterations" }, + { LLM_KV_HYPER_CONNECTION_EPSILON, "%s.hyper_connection.epsilon" }, + + { LLM_KV_HASH_LAYER_COUNT, "%s.hash_layer_count" }, + { LLM_KV_ROPE_DIMENSION_COUNT, "%s.rope.dimension_count" }, { LLM_KV_ROPE_DIMENSION_COUNT_SWA, "%s.rope.dimension_count_swa" }, { LLM_KV_ROPE_DIMENSION_COUNT_PER_LAYER,"%s.rope.dimension_count_per_layer" }, diff --git a/src/llama-arch.h b/src/llama-arch.h index 45bcbbf9..211a3052 100644 --- a/src/llama-arch.h +++ b/src/llama-arch.h @@ -51,6 +51,7 @@ enum llm_arch { LLM_ARCH_OPENELM, LLM_ARCH_ARCTIC, LLM_ARCH_DEEPSEEK2, + LLM_ARCH_DEEPSEEK4, LLM_ARCH_CHATGLM, LLM_ARCH_GLM4, LLM_ARCH_GLM4_MOE, @@ -174,12 +175,22 @@ enum llm_kv { LLM_KV_ATTENTION_INDEXER_HEAD_COUNT, LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, LLM_KV_ATTENTION_INDEXER_TOP_K, + LLM_KV_ATTENTION_OUTPUT_GROUP_COUNT, + LLM_KV_ATTENTION_OUTPUT_LORA_RANK, + LLM_KV_ATTENTION_COMPRESS_ROPE_FREQ_BASE, + LLM_KV_ATTENTION_COMPRESS_RATIOS, LLM_KV_FULL_ATTENTION_INTERVAL, LLM_KV_ATTENTION_SHARED_KV_LAYERS, LLM_KV_ATTENTION_KEY_LENGTH_SWA, LLM_KV_ATTENTION_VALUE_LENGTH_SWA, LLM_KV_ATTENTION_VALUE_SCALE, + LLM_KV_HYPER_CONNECTION_COUNT, + LLM_KV_HYPER_CONNECTION_SINKHORN_ITERATIONS, + LLM_KV_HYPER_CONNECTION_EPSILON, + + LLM_KV_HASH_LAYER_COUNT, + LLM_KV_ROPE_DIMENSION_COUNT, LLM_KV_ROPE_DIMENSION_COUNT_SWA, LLM_KV_ROPE_DIMENSION_COUNT_PER_LAYER, @@ -369,6 +380,27 @@ enum llm_tensor { LLM_TENSOR_INDEXER_PROJ, LLM_TENSOR_INDEXER_ATTN_K, LLM_TENSOR_INDEXER_ATTN_Q_B, // 97 + LLM_TENSOR_ATTN_KV_LATENT, + LLM_TENSOR_ATTN_OUT_A, + LLM_TENSOR_ATTN_OUT_B, + LLM_TENSOR_ATTN_COMP_KV, + LLM_TENSOR_ATTN_COMP_GATE, + LLM_TENSOR_ATTN_COMP_APE, + LLM_TENSOR_ATTN_COMP_NORM, + LLM_TENSOR_INDEXER_COMP_KV, + LLM_TENSOR_INDEXER_COMP_GATE, + LLM_TENSOR_INDEXER_COMP_APE, + LLM_TENSOR_INDEXER_COMP_NORM, + LLM_TENSOR_FFN_GATE_TID2EID, + LLM_TENSOR_HC_HEAD_BASE, + LLM_TENSOR_HC_HEAD_FN, + LLM_TENSOR_HC_HEAD_SCALE, + LLM_TENSOR_HC_ATTN_BASE, + LLM_TENSOR_HC_ATTN_FN, + LLM_TENSOR_HC_ATTN_SCALE, + LLM_TENSOR_HC_FFN_BASE, + LLM_TENSOR_HC_FFN_FN, + LLM_TENSOR_HC_FFN_SCALE, LLM_TENSOR_PER_LAYER_TOKEN_EMBD, LLM_TENSOR_PER_LAYER_MODEL_PROJ, diff --git a/src/llama-build-context.cpp b/src/llama-build-context.cpp index 117ea246..62d07c92 100644 --- a/src/llama-build-context.cpp +++ b/src/llama-build-context.cpp @@ -123,6 +123,12 @@ void llm_build_context::init() { lctx.dflash.inputs.target_features = nullptr; lctx.dflash.inputs.pos_ctx = nullptr; lctx.dflash.inputs.kq_mask = nullptr; + lctx.dsv4.inputs.raw_k_write_src_idxs = nullptr; + lctx.dsv4.inputs.raw_k_write_idxs = nullptr; + lctx.dsv4.inputs.raw_k_read_idxs = nullptr; + lctx.dsv4.inputs.csa = {}; + lctx.dsv4.inputs.hca = {}; + lctx.dsv4.inputs.lid = {}; } } @@ -146,7 +152,7 @@ ggml_cgraph * llm_build_context::build_k_shift() { ? LLAMA_ROPE_TYPE_NEOX : hparams.rope_type; - const float yarn_attn_factor_shift = model.arch == LLM_ARCH_DEEPSEEK2 || model.arch == LLM_ARCH_MISTRAL4 + const float yarn_attn_factor_shift = model.arch == LLM_ARCH_DEEPSEEK2 || model.arch == LLM_ARCH_DEEPSEEK4 || model.arch == LLM_ARCH_MISTRAL4 ? 1.0f / (1.0f + 0.1f * logf(1.0f / freq_scale)) : cparams.yarn_attn_factor; @@ -543,6 +549,121 @@ ggml_tensor * llm_build_context::build_inp_KQ_mask_swa_win(int64_t n_kv_win, boo return flash_attn ? ggml_cast(ctx0, lctx.inp_KQ_mask_swa_win, GGML_TYPE_F16) : lctx.inp_KQ_mask_swa_win; } +//build_mhc_post: x = 4096 x 4096 x 1 x 1, post = 4 x 4096 x 1 x 1, residual = 4096 x 4 x 4096 x 1, comb = 4 x 4 x 4096 x 1 +//build_mhc_post: x = 4096 x 1 x 1 x 1, post = 4 x 1 x 1 x 1, residual = 4096 x 4 x 1 x 1, comb = 4 x 4 x 1 x 1 +// x = n_embd x n_tokens <--- y in Pangu +// post = 4 x n_tokens <--- h_post in Pangu +// residual = n_embd x 4 x n_tokens <--- Rin in Pangu +// comb = 4 x 4 x n_tokens + +ggml_tensor * llm_build_context::build_mhc_post( + ggml_tensor * x, + ggml_tensor * post, + ggml_tensor * residual, + ggml_tensor * comb, + int64_t n_embd, + int64_t n_stream, + bool comb_output_dim0) { + const int64_t n_tokens = x->ne[1]; + ggml_tensor * out = nullptr; + + //ggml_tensor * x3 = ggml_reshape_3d(ctx0, x, n_embd, 1, n_tokens); + //ggml_tensor * post3 = ggml_reshape_3d(ctx0, post, 1, n_stream, n_tokens); + //ggml_tensor repeater; + //repeater.ne[0] = n_embd; repeater.ne[1] = n_stream; repeater.ne[2] = n_tokens; repeater.ne[3] = 1; + //ggml_tensor * term1 = ggml_mul(ctx0, ggml_repeat(ctx0, x3, &repeater), post3); + + //ggml_tensor * term2 = nullptr; + //for (int s = 0; s < n_stream; ++s) { + // ggml_tensor * m_s = ggml_cont(ctx0, ggml_view_2d(ctx0, comb, n_stream, n_tokens, comb->nb[2], s*comb->nb[1])); + // ggml_tensor * m_s3 = ggml_reshape_3d(ctx0, m_s, 1, n_stream, n_tokens); + // ggml_tensor * acc = ggml_mul(ctx0, residual, m_s3); + // ggml_tensor * summed = ggml_sum_rows_ext(ctx0, acc, 1); + // term2 = term2 ? ggml_concat(ctx0, term2, summed, 1) : summed; + //} + //return ggml_add(ctx0, term1, term2); + //printf("%s: x = %ld x %ld x %ld x %ld, post = %ld x %ld x %ld x %ld, residual = %ld x %ld x %ld x %ld, comb = %ld x %ld x %ld x %ld\n", + // __func__, x->ne[0], x->ne[1], x->ne[2], x->ne[3], post->ne[0], post->ne[1], post->ne[2], post->ne[3], + // residual->ne[0], residual->ne[1], residual->ne[2], residual->ne[3], + // comb->ne[0], comb->ne[1], comb->ne[2], comb->ne[3]); + + return ggml_hc_post(ctx0, x, post, residual, comb); + + for (int64_t dst = 0; dst < n_stream; ++dst) { + ggml_tensor * post_dst = ggml_cont(ctx0, + ggml_view_2d(ctx0, post, 1, n_tokens, post->nb[1], dst*post->nb[0])); + ggml_tensor * cur = ggml_mul(ctx0, x, post_dst); + + for (int64_t src = 0; src < n_stream; ++src) { + ggml_tensor * res_src = ggml_cont(ctx0, + ggml_view_2d(ctx0, residual, n_embd, n_tokens, + residual->nb[2], src*residual->nb[1])); + const size_t comb_offset = comb_output_dim0 + ? dst*comb->nb[0] + src*comb->nb[1] + : src*comb->nb[0] + dst*comb->nb[1]; + ggml_tensor * comb_dst_src = ggml_cont(ctx0, + ggml_view_2d(ctx0, comb, 1, n_tokens, comb->nb[2], comb_offset)); + cur = ggml_add(ctx0, cur, ggml_mul(ctx0, res_src, comb_dst_src)); + } + + cur = ggml_reshape_3d(ctx0, cur, n_embd, 1, n_tokens); + if (out) { + out = ggml_concat(ctx0, out, cur, 1); + cb(out, "mhc_post", dst); + } else { + out = cur; + } + //out = out ? ggml_concat(ctx0, out, cur, 1) : cur; + } + + return out; +} + +ggml_tensor * llm_build_context::build_mhc_weighted_sum( + ggml_tensor * x, + ggml_tensor * weights, + int64_t n_embd, + int64_t n_stream) { + const int64_t n_tokens = x->ne[2]; + ggml_tensor * out = nullptr; + + if (weights->ne[3] == 1 && x->ne[3] == 1) { + auto w = ggml_reshape_4d(ctx0, weights, 1, weights->ne[0], weights->ne[1], weights->ne[2]); + return ggml_mul_multi_add(ctx0, x, w); + } + for (int64_t stream = 0; stream < n_stream; ++stream) { + ggml_tensor * x_stream = ggml_cont(ctx0, + ggml_view_2d(ctx0, x, n_embd, n_tokens, + x->nb[2], stream*x->nb[1])); + ggml_tensor * weight = ggml_cont(ctx0, + ggml_view_2d(ctx0, weights, 1, n_tokens, + weights->nb[1], stream*weights->nb[0])); + ggml_tensor * cur = ggml_mul(ctx0, x_stream, weight); + out = out ? ggml_add(ctx0, out, cur) : cur; + } + + return out; +} + +ggml_tensor * llm_build_context::build_mhc_pre_projection( + ggml_tensor * x, + ggml_tensor * fn, + ggml_tensor * gamma, + int64_t n_embd, + int64_t n_stream, + float norm_rms_eps, + bool force_contiguous) { + if (force_contiguous && !ggml_is_contiguous(x)) { + x = ggml_cont(ctx0, x); + } + + ggml_tensor * flat = ggml_reshape_2d(ctx0, x, n_embd*n_stream, x->ne[2]); + ggml_tensor * normed = gamma ? ggml_fused_rms_norm(ctx0, flat, gamma, hparams.f_norm_rms_eps) : + ggml_rms_norm(ctx0, flat, norm_rms_eps); + + return ggml_mul_mat(ctx0, fn, normed); +} + ggml_tensor * llm_build_context::build_inp_mean() { lctx.inp_mean = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_tokens, n_tokens); cb(lctx.inp_mean, "inp_mean", -1); @@ -925,7 +1046,7 @@ ggml_tensor * llm_build_context::llm_build_ffn( } cur = ggml_fused_up_gate(ctx, split_u, split_g, cur, unary_op); cb(cur, "ffn_up_gate", il_cb); - if (lctx.model.arch == LLM_ARCH_STEP35) { + if (lctx.model.arch == LLM_ARCH_STEP35 || lctx.model.arch == LLM_ARCH_DEEPSEEK4) { *(float *)(cur->op_params + 1) = lctx.model.hparams.swiglu_limits[il]; } cur = llm_build_lora_mm(lctx, ctx, split_d, cur); @@ -986,7 +1107,7 @@ ggml_tensor * llm_build_context::llm_build_ffn( type_op == LLM_FFN_GELU ? GGML_UNARY_OP_GELU : GGML_UNARY_OP_SWIGLU_OAI; cur = ggml_fused_up_gate(ctx, up, gate, cur, unary_op); cb(cur, "ffn_up_gate", il); - if (lctx.model.arch == LLM_ARCH_STEP35) { + if (lctx.model.arch == LLM_ARCH_STEP35 || lctx.model.arch == LLM_ARCH_DEEPSEEK4) { *(float *)(cur->op_params + 1) = lctx.model.hparams.swiglu_limits_shared[il]; } if (down) { @@ -1065,7 +1186,7 @@ ggml_tensor * llm_build_context::llm_build_ffn( (type_op == LLM_FFN_SILU || type_op == LLM_FFN_RELU || (type_op == LLM_FFN_GELU && !act_scales))) { cur = ggml_fused_mul_unary(ctx, cur, tmp, type_op == LLM_FFN_SILU ? GGML_UNARY_OP_SILU : type_op == LLM_FFN_RELU ? GGML_UNARY_OP_RELU : GGML_UNARY_OP_GELU); - if (lctx.model.arch == LLM_ARCH_STEP35) { + if (lctx.model.arch == LLM_ARCH_STEP35 || lctx.model.arch == LLM_ARCH_DEEPSEEK4) { *((float *)(cur->op_params + 1)) = lctx.model.hparams.swiglu_limits_shared[il]; } } @@ -1184,7 +1305,8 @@ ggml_tensor * llm_build_context::llm_build_moe_ffn( llm_expert_gating_func_type gating_op, const llm_build_cb & cb, int il, ggml_cgraph * graph, bool add_input, ggml_tensor * up_gate_exps, ggml_tensor * up_gate_exps_b, - ggml_tensor * input_logits, ggml_tensor * down_exps_s) { + ggml_tensor * input_logits, ggml_tensor * down_exps_s, + ggml_tensor * selected_experts) { GGML_ASSERT(gate_inp || input_logits); @@ -1202,6 +1324,9 @@ llm_expert_gating_func_type gating_op, cb(logits, "ffn_moe_logits_biased", il); } + if (gating_op == LLM_EXPERT_GATING_FUNC_TYPE_SQRT_SOFTPLUS) { + ggml_mul_mat_set_prec(logits, GGML_PREC_F32); + } //ggml_tensor * probs = ggml_soft_max(ctx, logits); // [n_expert, n_tokens] ggml_tensor * probs = nullptr; @@ -1218,6 +1343,11 @@ llm_expert_gating_func_type gating_op, { probs = logits; // [n_expert, n_tokens] } break; + case LLM_EXPERT_GATING_FUNC_TYPE_SQRT_SOFTPLUS: + { + //probs = ggml_sqrt(ctx, ggml_softplus(ctx, logits)); // [n_expert, n_tokens] + probs = ggml_sqrt_softplus(ctx, logits); // [n_expert, n_tokens] + } break; default: GGML_ABORT("fatal error"); } @@ -1238,14 +1368,15 @@ llm_expert_gating_func_type gating_op, } // select experts - ggml_tensor * selected_experts; - if (lctx.cparams.grouped_expert_routing && lctx.model.arch == LLM_ARCH_BAILINGMOE2 && n_tokens > 0) { - auto& hparams = lctx.model.hparams; - selected_experts = ggml_grouped_topk(ctx, selection_probs, hparams.n_expert_groups, hparams.n_group_used, 2, n_expert_used); - } else { - //selected_experts = ggml_top_k_thresh(ctx, selection_probs, n_expert_used, - // lctx.cparams.min_experts, lctx.cparams.thresh_experts); // [n_expert_used, n_tokens] - selected_experts = ggml_top_k(ctx, selection_probs, n_expert_used); // [n_expert_used, n_tokens] + if (selected_experts == nullptr) { + if (lctx.cparams.grouped_expert_routing && lctx.model.arch == LLM_ARCH_BAILINGMOE2 && n_tokens > 0) { + auto& hparams = lctx.model.hparams; + selected_experts = ggml_grouped_topk(ctx, selection_probs, hparams.n_expert_groups, hparams.n_group_used, 2, n_expert_used); + } else { + //selected_experts = ggml_top_k_thresh(ctx, selection_probs, n_expert_used, + // lctx.cparams.min_experts, lctx.cparams.thresh_experts); // [n_expert_used, n_tokens] + selected_experts = ggml_top_k(ctx, selection_probs, n_expert_used); // [n_expert_used, n_tokens] + } } cb(selected_experts, "ffn_moe_topk", il); ggml_tensor * weights = ggml_get_rows(ctx, @@ -1314,7 +1445,7 @@ llm_expert_gating_func_type gating_op, par = ggml_moe_up_gate(ctx, up_gate_exps, nullptr, cur, selected_experts, type_op == LLM_FFN_SILU ? GGML_UNARY_OP_SILU : GGML_UNARY_OP_GELU); } - if (lctx.model.arch == LLM_ARCH_STEP35) { + if (lctx.model.arch == LLM_ARCH_STEP35 || lctx.model.arch == LLM_ARCH_DEEPSEEK4) { *((float *)(par->op_params + 1)) = lctx.model.hparams.swiglu_limits[il]; } } else { @@ -1330,7 +1461,7 @@ llm_expert_gating_func_type gating_op, par = ggml_moe_up_gate(ctx, up_exps, gate_exps, cur, selected_experts, type_op == LLM_FFN_SILU ? GGML_UNARY_OP_SILU : GGML_UNARY_OP_GELU); } - if (lctx.model.arch == LLM_ARCH_STEP35) { + if (lctx.model.arch == LLM_ARCH_STEP35 || lctx.model.arch == LLM_ARCH_DEEPSEEK4) { *(float *)(par->op_params + 1) = lctx.model.hparams.swiglu_limits[il]; } } else { @@ -1358,7 +1489,7 @@ llm_expert_gating_func_type gating_op, if (type_op == LLM_FFN_SILU || type_op == LLM_FFN_GELU) { par = ggml_fused_mul_unary(ctx, gate, up, type_op == LLM_FFN_SILU ? GGML_UNARY_OP_SILU : GGML_UNARY_OP_GELU); - if (lctx.model.arch == LLM_ARCH_STEP35) { + if (lctx.model.arch == LLM_ARCH_STEP35 || lctx.model.arch == LLM_ARCH_DEEPSEEK4) { *((float *)(par->op_params + 1)) = lctx.model.hparams.swiglu_limits[il]; } } else if (type_op == LLM_FFN_SWIGLU_OAI) { @@ -2645,6 +2776,10 @@ ggml_cgraph * llm_build_context::llama_build_graph( { result = llm.build_deepseek2(); } break; + case LLM_ARCH_DEEPSEEK4: + { + result = llm.build_deepseek4(); + } break; case LLM_ARCH_OPENPANGU: { result = llm.build_openpangu(); diff --git a/src/llama-build-context.h b/src/llama-build-context.h index 854b7b5a..93bcb363 100644 --- a/src/llama-build-context.h +++ b/src/llama-build-context.h @@ -275,6 +275,7 @@ struct llm_build_context { ggml_cgraph * build_arctic(); ggml_cgraph * build_deepseek2(); + ggml_cgraph * build_deepseek4(); ggml_cgraph * build_openpangu(); // openPangu attention sublayer body (shared by base layers and the NextN/MTP head): @@ -309,6 +310,30 @@ struct llm_build_context { bool cache_writes_only = false, bool KQ_mask_swa_windowed = false); + ggml_tensor * build_mhc_post( + ggml_tensor * x, + ggml_tensor * post, + ggml_tensor * residual, + ggml_tensor * comb, + int64_t n_embd, + int64_t n_stream, + bool comb_output_dim0); + + ggml_tensor * build_mhc_weighted_sum( + ggml_tensor * x, + ggml_tensor * weights, + int64_t n_embd, + int64_t n_stream); + + ggml_tensor * build_mhc_pre_projection( + ggml_tensor * x, + ggml_tensor * fn, + ggml_tensor * gamma, + int64_t n_embd, + int64_t n_stream, + float norm_rms_eps, + bool force_contiguous); + ggml_tensor * build_deepseek2_tp_attention( ggml_cgraph * gf, int il, ggml_tensor * inpL, @@ -473,7 +498,8 @@ struct llm_build_context { llm_expert_gating_func_type gating_op, const llm_build_cb & cb, int il, ggml_cgraph * graph = nullptr, bool add_input = false, ggml_tensor * up_gate_exps = nullptr, ggml_tensor * up_gate_exps_b = nullptr, - ggml_tensor * input_logits = nullptr, ggml_tensor * down_exps_s = nullptr); + ggml_tensor * input_logits = nullptr, ggml_tensor * down_exps_s = nullptr, + ggml_tensor * selected_experts = nullptr); static ggml_tensor * llm_build_moe_ffn(ggml_context * ctx, llama_context & lctx, ggml_tensor * cur, @@ -501,7 +527,7 @@ llm_expert_gating_func_type gating_op, n_expert, n_expert_used, type_op, norm_w, scale_w, w_scale, gating_op, cb, il, graph, add_input, up_gate_exps, up_gate_exps_b, - input_logits, down_exps_s); + input_logits, down_exps_s, nullptr); } static ggml_tensor * llm_build_std_moe_ffn(ggml_context * ctx, llama_context & lctx, diff --git a/src/llama-context.h b/src/llama-context.h index c5cf1a85..676e09d4 100644 --- a/src/llama-context.h +++ b/src/llama-context.h @@ -392,6 +392,121 @@ struct llama_context { dflash_runtime dflash; using dflash_capture_state = dflash_runtime::capture_state; + struct dsv4_runtime { + static constexpr uint32_t CSA_RATIO = 4; + static constexpr uint32_t HCA_RATIO = 128; + + struct slot_info { + int32_t s0 = 0; + int32_t s1 = 0; + std::vector strm; + std::vector> idxs; + + void resize(size_t n) { + strm.resize(n); + idxs.resize(n); + } + + size_t size() const { + GGML_ASSERT(strm.size() == idxs.size()); + if (idxs.empty()) { + return 0; + } + + return idxs[0].size(); + } + + size_t n_stream() const { + GGML_ASSERT(strm.size() == idxs.size()); + return strm.size(); + } + + bool empty() const { + return idxs.empty(); + } + }; + + struct raw_context { + std::vector write_src_idxs; + std::vector write_dst_idxs; + std::vector read_dst_idxs; + std::vector write_counts; + std::vector read_counts; + slot_info sinfo_write; + slot_info sinfo_read; + int64_t graph_n_stream = 1; + int64_t n_kv = 0; + }; + + struct comp_context { + slot_info sinfo; + int64_t graph_n_stream = 1; + int64_t n_kv = 0; + }; + + struct comp_plan { + std::vector state_pos; + std::vector state_persist_src_idxs; + std::vector state_persist_dst_idxs; + std::vector state_read_idxs; + std::vector state_write_idxs; + std::vector state_write_pos; + std::vector n_visible; + int64_t n_stream = 1; + int64_t n_kv = 0; + }; + + struct comp_inputs { + struct ggml_tensor * state_pos = nullptr; + struct ggml_tensor * state_persist_src_idxs = nullptr; + struct ggml_tensor * state_persist_dst_idxs = nullptr; + struct ggml_tensor * state_read_idxs = nullptr; + struct ggml_tensor * state_write_idxs = nullptr; + struct ggml_tensor * state_write_pos = nullptr; + struct ggml_tensor * kq_mask = nullptr; + }; + + struct storage { + std::vector csa_k; + std::vector hca_k; + std::vector lid_k; + + std::vector csa_state_kv; + std::vector csa_state_score; + std::vector hca_state_kv; + std::vector hca_state_score; + std::vector lid_state_kv; + std::vector lid_state_score; + + struct ggml_context * cache_ctx = nullptr; + std::vector cache_bufs; + uint32_t n_stream = 1; + }; + + struct input_state { + struct ggml_tensor * raw_k_write_src_idxs = nullptr; + struct ggml_tensor * raw_k_write_idxs = nullptr; + struct ggml_tensor * raw_k_read_idxs = nullptr; + comp_inputs csa; + comp_inputs hca; + comp_inputs lid; + }; + + storage cache; + input_state inputs; + raw_context raw; + comp_context csa_ctx; + comp_context hca_ctx; + comp_context lid_ctx; + comp_plan csa_plan; + comp_plan hca_plan; + comp_plan lid_plan; + + std::vector csa_mask_data; + std::vector hca_mask_data; + }; + dsv4_runtime dsv4; + // input tensors struct ggml_tensor * inp_tokens; // I32 [n_batch] struct ggml_tensor * inp_embd; // F32 [n_embd, n_batch] @@ -464,6 +579,8 @@ struct llama_context { bool ensure_dflash_kv_cache_tensors(int32_t cross_ctx); void free_dflash_kv_cache_tensors(); + bool ensure_dsv4_cache_tensors(); + void free_dsv4_cache_tensors(); bool prepare_mtp_graph_inputs( struct llama_context & lctx); diff --git a/src/llama-cparams.h b/src/llama-cparams.h index b4fe0613..758ad69b 100644 --- a/src/llama-cparams.h +++ b/src/llama-cparams.h @@ -56,6 +56,7 @@ struct llama_cparams { enum ggml_type reduce_type; enum ggml_type graph_attn_precision; + enum ggml_type idx_type_k = GGML_TYPE_F16; enum llama_pooling_type pooling_type; enum llama_mtp_op_type mtp_op_type; diff --git a/src/llama-dsv4.cpp b/src/llama-dsv4.cpp new file mode 100644 index 00000000..7dad2ea8 --- /dev/null +++ b/src/llama-dsv4.cpp @@ -0,0 +1,1095 @@ +#include "llama-dsv4.h" + +#include "llama-context.h" +#include "llama-model.h" +#include "llama-impl.h" + +#include "ggml.h" +#include "ggml-backend.h" + +#include +#include +#include +#include +#include +#include + +static bool dsv4_cache_type_supported(ggml_type type) { + return type == GGML_TYPE_F16 || type == GGML_TYPE_BF16 || type == GGML_TYPE_Q8_0; +} + +static bool dsv4_validate_cache_type(ggml_type type, int64_t width, const char * name) { + if (!dsv4_cache_type_supported(type)) { + LLAMA_LOG_ERROR("%s: unsupported DSV4 %s cache type %s\n", __func__, name, ggml_type_name(type)); + return false; + } + if (ggml_is_quantized(type) && width % ggml_blck_size(type) != 0) { + LLAMA_LOG_ERROR("%s: DSV4 %s cache width %d is not aligned to %d elements for %s\n", + __func__, name, (int)width, (int)ggml_blck_size(type), ggml_type_name(type)); + return false; + } + return true; +} + +static ggml_backend_buffer_type_t llama_dsv4_layer_buft(const llama_context & lctx, int32_t il) { + if (il >= 0 && il < (int32_t) lctx.model.buft_layer.size() && lctx.model.buft_layer[il].buft != nullptr) { + return lctx.model.buft_layer[il].buft; + } + + if (il >= 0 && il < (int32_t) lctx.model.layers.size()) { + const ggml_tensor * ref = lctx.model.layers[il].attn_comp_wkv; + if (ref == nullptr) { + ref = lctx.model.layers[il].wq_a; + } + if (ref != nullptr && ref->buffer != nullptr) { + return ggml_backend_buffer_get_type(ref->buffer); + } + } + + return llama_default_buffer_type_cpu(true); +} + +static uint32_t dsv4_comp_size(uint32_t kv_size, uint32_t ratio) { + return std::max(1, (kv_size + ratio - 1)/ratio); +} + +static bool dsv4_validate_csa_lid_visibility( + const llama_context & lctx, + uint32_t csa_kv_size, + uint32_t lid_kv_size) { + const auto & csa_plan = lctx.dsv4.csa_plan; + const auto & lid_plan = lctx.dsv4.lid_plan; + const auto & csa_ctx = lctx.dsv4.csa_ctx; + const auto & lid_ctx = lctx.dsv4.lid_ctx; + + if (csa_kv_size != lid_kv_size || + csa_plan.n_stream != lid_plan.n_stream || + csa_plan.n_kv != lid_plan.n_kv || + csa_plan.n_visible != lid_plan.n_visible || + csa_ctx.graph_n_stream != lid_ctx.graph_n_stream || + csa_ctx.n_kv != lid_ctx.n_kv || + csa_ctx.sinfo.strm != lid_ctx.sinfo.strm || + csa_ctx.sinfo.idxs != lid_ctx.sinfo.idxs || + csa_ctx.sinfo.s0 != lid_ctx.sinfo.s0 || + csa_ctx.sinfo.s1 != lid_ctx.sinfo.s1) { + LLAMA_LOG_ERROR("%s: DSV4 CSA/LID visibility contracts differ\n", __func__); + return false; + } + + return true; +} + +static void dsv4_batch_shape( + const llama_batch & batch, + uint32_t & n_seqs, + uint32_t & n_seq_tokens) { + n_seqs = 1; + n_seq_tokens = (uint32_t) std::max(1, batch.n_tokens); + + if (batch.n_tokens <= 0 || batch.n_seq_id == nullptr || batch.seq_id == nullptr) { + return; + } + + std::map counts; + for (int32_t i = 0; i < batch.n_tokens; ++i) { + if (batch.n_seq_id[i] != 1 || batch.seq_id[i] == nullptr) { + return; + } + + counts[batch.seq_id[i][0]]++; + } + + if (counts.empty()) { + return; + } + + const uint32_t seq_tokens = counts.begin()->second; + for (const auto & [_, count] : counts) { + if (count != seq_tokens) { + return; + } + } + + n_seqs = (uint32_t) counts.size(); + n_seq_tokens = std::max(1, seq_tokens); +} + +static bool dsv4_batch_has_coupled(const llama_batch & batch) { + if (batch.n_tokens <= 0 || batch.n_seq_id == nullptr) { + return false; + } + + for (int32_t i = 0; i < batch.n_tokens; ++i) { + if (batch.n_seq_id[i] > 1) { + return true; + } + } + + return false; +} + +static bool dsv4_token_has_seq(const llama_batch & batch, int32_t i, llama_seq_id seq_id) { + if (batch.n_seq_id == nullptr || batch.seq_id == nullptr || batch.seq_id[i] == nullptr) { + return seq_id == 0; + } + + for (int32_t s = 0; s < batch.n_seq_id[i]; ++s) { + if (batch.seq_id[i][s] == seq_id) { + return true; + } + } + + return false; +} + +static std::vector dsv4_batch_unique_seq_ids(const llama_batch & batch) { + std::vector seq_ids; + std::unordered_set seen; + + if (batch.n_tokens <= 0 || batch.n_seq_id == nullptr || batch.seq_id == nullptr) { + seq_ids.push_back(0); + return seq_ids; + } + + for (int32_t i = 0; i < batch.n_tokens; ++i) { + if (batch.n_seq_id[i] <= 0 || batch.seq_id[i] == nullptr) { + continue; + } + + for (int32_t s = 0; s < batch.n_seq_id[i]; ++s) { + const llama_seq_id seq_id = batch.seq_id[i][s]; + if (seen.insert(seq_id).second) { + seq_ids.push_back(seq_id); + } + } + } + + if (seq_ids.empty()) { + seq_ids.push_back(0); + } + + return seq_ids; +} + +static int64_t dsv4_stream_offset(uint32_t n_stream, llama_seq_id seq_id, uint32_t size) { + if (n_stream <= 1) { + return 0; + } + + if (seq_id < 0 || (uint32_t) seq_id >= n_stream) { + LLAMA_LOG_ERROR("%s: DSV4 seq_id %d is outside stream range %u\n", __func__, seq_id, n_stream); + return -1; + } + + return (int64_t) seq_id*size; +} + +static int64_t dsv4_comp_graph_n_stream(const llama_batch & batch, uint32_t n_stream) { + if (n_stream <= 1) { + return 1; + } + + const std::vector seq_ids = dsv4_batch_unique_seq_ids(batch); + if (seq_ids.size() <= 1 || dsv4_batch_has_coupled(batch)) { + return 1; + } + + return (int64_t) seq_ids.size(); +} + +static std::vector dsv4_build_stream_seq_ids( + const llama_batch & batch, + uint32_t n_stream) { + if (n_stream <= 1) { + return { 0 }; + } + + const std::vector seq_ids = dsv4_batch_unique_seq_ids(batch); + if (seq_ids.size() <= 1 || dsv4_batch_has_coupled(batch)) { + return { seq_ids.empty() ? 0 : seq_ids.front() }; + } + + return seq_ids; +} + +static llama_context::dsv4_runtime::slot_info dsv4_build_comp_sinfo( + const llama_batch & batch, + uint32_t n_stream) { + llama_context::dsv4_runtime::slot_info sinfo; + + const std::vector seq_ids = dsv4_build_stream_seq_ids(batch, n_stream); + const int64_t graph_n_stream = (int64_t) seq_ids.size(); + bool have_stream = false; + + sinfo.s0 = INT_MAX; + sinfo.s1 = 0; + sinfo.resize((size_t) std::max(1, graph_n_stream)); + for (int64_t s = 0; s < graph_n_stream; ++s) { + const llama_seq_id seq_id = seq_ids[(size_t) s]; + const int64_t strm = dsv4_stream_offset(n_stream, seq_id, 1); + if (strm < 0) { + continue; + } + sinfo.strm[(size_t) s] = (llama_seq_id) strm; + sinfo.idxs[(size_t) s].assign(1, 0); + sinfo.s0 = std::min(sinfo.s0, (int32_t) strm); + sinfo.s1 = std::max(sinfo.s1, (int32_t) strm); + have_stream = true; + } + + if (!have_stream) { + sinfo.resize(1); + sinfo.strm[0] = 0; + sinfo.idxs[0].assign(1, 0); + sinfo.s0 = 0; + sinfo.s1 = 0; + } + + if (n_stream > 1 && sinfo.s1 - sinfo.s0 + 1 != (int32_t) sinfo.n_stream()) { + LLAMA_LOG_ERROR("%s: DSV4 compressed streams are not contiguous in batch\n", __func__); + } + + return sinfo; +} + +static llama_context::dsv4_runtime::slot_info dsv4_build_raw_read_sinfo( + const llama_context::dsv4_runtime::slot_info & sinfo_write, + const llama_batch & batch, + uint32_t n_stream) { + if (!dsv4_batch_has_coupled(batch)) { + return sinfo_write; + } + + const llama_seq_id seq_id = + (batch.n_seq_id != nullptr && batch.seq_id != nullptr && batch.n_tokens > 0 && batch.n_seq_id[0] > 0 && batch.seq_id[0] != nullptr) + ? batch.seq_id[0][0] + : 0; + const int64_t strm = dsv4_stream_offset(n_stream, seq_id, 1); + if (strm < 0) { + return {}; + } + + size_t i_stream = 0; + for (; i_stream < sinfo_write.n_stream(); ++i_stream) { + if ((int64_t) sinfo_write.strm[i_stream] == strm) { + break; + } + } + if (i_stream == sinfo_write.n_stream()) { + LLAMA_LOG_ERROR("%s: DSV4 raw write stream not found for coupled read\n", __func__); + return {}; + } + + llama_context::dsv4_runtime::slot_info sinfo; + sinfo.resize(1); + sinfo.strm[0] = sinfo_write.strm[i_stream]; + sinfo.idxs[0] = sinfo_write.idxs[i_stream]; + sinfo.s0 = (int32_t) strm; + sinfo.s1 = sinfo.s0; + + return sinfo; +} + +static bool dsv4_validate_batch_seq_ids( + const llama_context & lctx, + const llama_batch & batch) { + if (batch.n_tokens <= 0 || batch.n_seq_id == nullptr || batch.seq_id == nullptr) { + return true; + } + + const uint32_t n_stream = std::max(1, lctx.cparams.n_seq_max); + for (int32_t i = 0; i < batch.n_tokens; ++i) { + if (batch.n_seq_id[i] <= 0 || batch.seq_id[i] == nullptr) { + LLAMA_LOG_ERROR("%s: DSV4 token %d is missing seq_id ownership\n", __func__, i); + return false; + } + + for (int32_t s = 0; s < batch.n_seq_id[i]; ++s) { + const llama_seq_id seq_id = batch.seq_id[i][s]; + if (seq_id < 0 || (uint32_t) seq_id >= n_stream) { + LLAMA_LOG_ERROR("%s: DSV4 token %d seq_id %d is outside n_seq_max=%u\n", + __func__, i, seq_id, n_stream); + return false; + } + } + } + + return true; +} + +static bool dsv4_build_raw_context( + const llama_context & lctx, + const llama_batch & batch, + llama_context::dsv4_runtime::raw_context & raw) { + raw = {}; + const uint32_t n_stream = std::max(1, lctx.cparams.n_seq_max); + const std::vector write_seq_ids = dsv4_build_stream_seq_ids(batch, n_stream); + raw.sinfo_write = dsv4_build_comp_sinfo(batch, n_stream); + raw.sinfo_read = dsv4_build_raw_read_sinfo(raw.sinfo_write, batch, n_stream); + raw.graph_n_stream = (int64_t) raw.sinfo_write.n_stream(); + std::vector read_seq_ids = write_seq_ids; + + if (dsv4_batch_has_coupled(batch)) { + const llama_seq_id coupled_seq_id = + (batch.n_seq_id != nullptr && batch.seq_id != nullptr && batch.n_tokens > 0 && batch.n_seq_id[0] > 0 && batch.seq_id[0] != nullptr) + ? batch.seq_id[0][0] + : 0; + read_seq_ids.assign(1, coupled_seq_id); + } + + if (batch.n_tokens <= 0) { + return true; + } + + const llama_kv_cache & kv = lctx.kv_self; + if (kv.head + batch.n_tokens > (int32_t) kv.size) { + LLAMA_LOG_ERROR("%s: DSV4 raw write slots [%d, %d) are outside kv cache size %u\n", + __func__, kv.head, kv.head + batch.n_tokens, kv.size); + return false; + } + + raw.write_counts.push_back(batch.n_tokens); + for (int32_t i = 0; i < batch.n_tokens; ++i) { + const int32_t slot = kv.head + i; + const llama_kv_cell & cell = kv.cells[(size_t) slot]; + + if (batch.pos != nullptr && cell.pos != batch.pos[i]) { + LLAMA_LOG_ERROR("%s: DSV4 raw write slot %d pos mismatch: cell=%d batch=%d\n", + __func__, slot, cell.pos, batch.pos[i]); + return false; + } + + raw.write_src_idxs.push_back(i); + raw.write_dst_idxs.push_back(slot); + } + + raw.n_kv = 0; + + for (size_t s = 0; s < raw.sinfo_read.n_stream(); ++s) { + const llama_seq_id seq_id = read_seq_ids[s]; + raw.sinfo_read.idxs[s].clear(); + int32_t count = 0; + for (uint32_t slot = 0; slot < kv.size; ++slot) { + const llama_kv_cell & cell = kv.cells[slot]; + if (cell.is_empty() || cell.pos < 0) { + continue; + } + if (!cell.has_seq_id(seq_id)) { + continue; + } + raw.sinfo_read.idxs[s].push_back(slot); + raw.read_dst_idxs.push_back((int32_t) slot); + ++count; + } + raw.read_counts.push_back(count); + raw.n_kv = std::max(raw.n_kv, count); + } + + if (raw.read_counts.empty()) { + raw.read_counts.push_back(0); + } + + for (size_t s = 0; s < raw.sinfo_write.n_stream(); ++s) { + const llama_seq_id seq_id = write_seq_ids[s]; + raw.sinfo_write.idxs[s].clear(); + for (int32_t i = 0; i < batch.n_tokens; ++i) { + if (!dsv4_token_has_seq(batch, i, seq_id)) { + continue; + } + raw.sinfo_write.idxs[s].push_back((uint32_t) (kv.head + i)); + } + } + + if (raw.sinfo_write.n_stream() > 1) { + std::vector write_src_idxs; + std::vector write_dst_idxs; + const size_t rows_per_stream = raw.sinfo_write.size(); + for (size_t s = 0; s < raw.sinfo_write.n_stream(); ++s) { + if (raw.sinfo_write.idxs[s].size() != rows_per_stream) { + LLAMA_LOG_ERROR("%s: DSV4 packed batch has unequal raw-write rows per stream\n", __func__); + return false; + } + + for (int32_t i = 0; i < batch.n_tokens; ++i) { + if (dsv4_token_has_seq(batch, i, write_seq_ids[s])) { + write_src_idxs.push_back(i); + } + } + + for (uint32_t slot : raw.sinfo_write.idxs[s]) { + write_dst_idxs.push_back((int32_t) slot); + } + } + + raw.write_src_idxs = std::move(write_src_idxs); + raw.write_dst_idxs = std::move(write_dst_idxs); + } + + // The graph exposes a rectangular raw-key view. Repeat the last valid row + // for shorter streams; the corresponding mask entries remain -INFINITY. + // This preserves the logical visibility while allowing one get_rows op to + // serve all streams. + if (raw.n_kv > 0) { + raw.read_dst_idxs.clear(); + const size_t read_rows = GGML_PAD((size_t) raw.n_kv, 256u); + for (size_t s = 0; s < raw.sinfo_read.n_stream(); ++s) { + const auto & rows = raw.sinfo_read.idxs[s]; + for (uint32_t slot : rows) { + raw.read_dst_idxs.push_back((int32_t) slot); + } + + const int32_t pad = rows.empty() ? 0 : (int32_t) rows.back(); + for (size_t i = rows.size(); i < read_rows; ++i) { + raw.read_dst_idxs.push_back(pad); + } + } + } + + return true; +} + +static llama_context::dsv4_runtime::comp_context dsv4_build_comp_context( + const llama_batch & batch, + uint32_t n_stream, + int64_t n_kv) { + llama_context::dsv4_runtime::comp_context ctx; + ctx.sinfo = dsv4_build_comp_sinfo(batch, n_stream); + ctx.graph_n_stream = dsv4_comp_graph_n_stream(batch, n_stream); + ctx.n_kv = n_kv; + return ctx; +} + +static llama_context::dsv4_runtime::comp_plan dsv4_build_reserve_comp_plan( + const llama_batch & batch, + uint32_t ratio, + bool overlap, + uint32_t state_size, + uint32_t kv_size, + uint32_t n_stream) { + llama_context::dsv4_runtime::comp_plan plan; + plan.n_visible.resize((size_t) batch.n_tokens, (int32_t) kv_size); + plan.n_stream = dsv4_comp_graph_n_stream(batch, n_stream); + plan.n_kv = kv_size; + + if (batch.n_tokens == 0) { + return plan; + } + + uint32_t n_seqs = 1; + uint32_t n_seq_tokens = 1; + dsv4_batch_shape(batch, n_seqs, n_seq_tokens); + + plan.n_visible.assign((size_t) batch.n_tokens, 0); + + const uint64_t n_blocks_u64 = (uint64_t) n_seqs*((n_seq_tokens + ratio - 1)/ratio); + const size_t n_blocks = (size_t) std::max(1, n_blocks_u64); + GGML_ASSERT((uint64_t) n_blocks == std::max(1, n_blocks_u64)); + const uint64_t state_rows = (uint64_t) state_size*(uint64_t) n_stream; + const size_t n_persist = (size_t) std::min((uint64_t) batch.n_tokens, state_rows); + + plan.state_pos.resize((size_t) batch.n_tokens); + plan.state_persist_src_idxs.resize(n_persist); + plan.state_persist_dst_idxs.resize(n_persist); + plan.state_read_idxs.resize((overlap ? 2u : 1u)*ratio*n_blocks); + plan.state_write_idxs.resize(n_blocks); + plan.state_write_pos.resize(n_blocks); + + return plan; +} + +static uint32_t dsv4_cache_kv_size(const std::vector & tensors) { + for (ggml_tensor * tensor : tensors) { + if (tensor != nullptr) { + return (uint32_t) tensor->ne[1]; + } + } + + return 0; +} + +static uint32_t dsv4_cache_state_size(const std::vector & tensors) { + for (ggml_tensor * tensor : tensors) { + if (tensor != nullptr) { + return (uint32_t) tensor->ne[1]; + } + } + + return 0; +} + +static bool dsv4_validate_comp_plan( + const char * tag, + const llama_batch & batch, + const llama_context::dsv4_runtime::comp_plan & plan, + uint32_t ratio, + bool overlap, + uint32_t state_size, + uint32_t kv_size, + uint32_t n_stream) { + const int64_t max_state_read_idx = (int64_t) state_size*n_stream + batch.n_tokens + (overlap ? 0 : -1); + + if (plan.n_visible.size() != (size_t) std::max(0, batch.n_tokens)) { + LLAMA_LOG_ERROR("%s: DSV4 %s plan n_visible size mismatch: got=%zu expected=%d\n", + __func__, tag, plan.n_visible.size(), std::max(0, batch.n_tokens)); + return false; + } + + if (plan.state_pos.size() > (size_t) std::max(0, batch.n_tokens)) { + LLAMA_LOG_ERROR("%s: DSV4 %s plan has too many state_pos rows: %zu > %d\n", + __func__, tag, plan.state_pos.size(), std::max(0, batch.n_tokens)); + return false; + } + + if (plan.state_persist_src_idxs.size() != plan.state_persist_dst_idxs.size()) { + LLAMA_LOG_ERROR("%s: DSV4 %s persist idx size mismatch: src=%zu dst=%zu\n", + __func__, tag, plan.state_persist_src_idxs.size(), plan.state_persist_dst_idxs.size()); + return false; + } + + if (plan.state_write_idxs.size() != plan.state_write_pos.size()) { + LLAMA_LOG_ERROR("%s: DSV4 %s write idx size mismatch: idxs=%zu pos=%zu\n", + __func__, tag, plan.state_write_idxs.size(), plan.state_write_pos.size()); + return false; + } + + for (size_t i = 0; i < plan.n_visible.size(); ++i) { + const int32_t n_visible = plan.n_visible[i]; + if (n_visible < 0 || (uint32_t) n_visible > kv_size) { + LLAMA_LOG_ERROR("%s: DSV4 %s n_visible[%zu]=%d exceeds kv_size=%u\n", + __func__, tag, i, n_visible, kv_size); + return false; + } + } + + for (size_t i = 0; i < plan.state_pos.size(); ++i) { + const int64_t pos = plan.state_pos[i]; + if (pos < 0 || pos >= (int64_t) ratio) { + LLAMA_LOG_ERROR("%s: DSV4 %s state_pos[%zu]=%lld outside ratio=%u\n", + __func__, tag, i, (long long) pos, ratio); + return false; + } + } + + for (size_t i = 0; i < plan.state_persist_src_idxs.size(); ++i) { + const int64_t src = plan.state_persist_src_idxs[i]; + const int64_t dst = plan.state_persist_dst_idxs[i]; + if (src < 0 || src >= batch.n_tokens) { + LLAMA_LOG_ERROR("%s: DSV4 %s persist src[%zu]=%lld outside current batch rows=%d\n", + __func__, tag, i, (long long) src, batch.n_tokens); + return false; + } + if (dst < 0 || (uint32_t) dst >= state_size*n_stream) { + LLAMA_LOG_ERROR("%s: DSV4 %s persist dst[%zu]=%lld outside state_size*n_stream=%u\n", + __func__, tag, i, (long long) dst, state_size*n_stream); + return false; + } + } + + for (size_t i = 0; i < plan.state_read_idxs.size(); ++i) { + const int64_t idx = plan.state_read_idxs[i]; + if (idx < 0 || idx > max_state_read_idx) { + LLAMA_LOG_ERROR("%s: DSV4 %s read idx[%zu]=%lld outside max source row=%lld\n", + __func__, tag, i, (long long) idx, (long long) max_state_read_idx); + return false; + } + } + + for (size_t i = 0; i < plan.state_write_idxs.size(); ++i) { + const int64_t idx = plan.state_write_idxs[i]; + if (idx < 0 || (uint32_t) idx >= kv_size*n_stream) { + LLAMA_LOG_ERROR("%s: DSV4 %s write idx[%zu]=%lld outside kv_size*n_stream=%u\n", + __func__, tag, i, (long long) idx, kv_size*n_stream); + return false; + } + } + + if (plan.n_kv == 0 || (uint32_t) plan.n_kv > kv_size) { + LLAMA_LOG_ERROR("%s: DSV4 %s plan n_kv=%lld outside kv_size=%u\n", + __func__, tag, (long long) plan.n_kv, kv_size); + return false; + } + + return true; +} + +static llama_context::dsv4_runtime::comp_plan dsv4_build_comp_plan( + const llama_batch & batch, + uint32_t ratio, + bool overlap, + uint32_t state_size, + uint32_t kv_size, + uint32_t n_stream) { + llama_context::dsv4_runtime::comp_plan plan; + plan.n_visible.resize((size_t) batch.n_tokens); + plan.n_stream = dsv4_comp_graph_n_stream(batch, n_stream); + + if (n_stream <= 1 && dsv4_batch_unique_seq_ids(batch).size() > 1) { + LLAMA_LOG_ERROR("%s: DSV4 single compressed stream cannot serve multiple sequences\n", __func__); + return plan; + } + + const int64_t state_rows = (int64_t) state_size*n_stream; + + struct persist_row { + int32_t dst; + int32_t src; + llama_pos pos; + }; + + std::vector persist_rows; + std::vector overlap_prev_reads; + std::vector overlap_cur_reads; + std::map, int32_t> curr_token_idx_map; + + for (int32_t i = 0; i < batch.n_tokens; ++i) { + const llama_seq_id seq_id = + batch.n_seq_id != nullptr && batch.seq_id != nullptr && batch.n_seq_id[i] > 0 && batch.seq_id[i] != nullptr + ? batch.seq_id[i][0] + : 0; + curr_token_idx_map[std::make_pair(seq_id, batch.pos[i])] = i; + } + + const auto state_source_idx = [&](llama_seq_id seq_id, llama_pos pos) -> int32_t { + if (pos < 0) { + return (int32_t) (state_rows + batch.n_tokens); + } + + const auto it = curr_token_idx_map.find(std::make_pair(seq_id, pos)); + if (it != curr_token_idx_map.end()) { + return (int32_t) (state_rows + it->second); + } + + const int64_t stream_off = dsv4_stream_offset(n_stream, seq_id, state_size); + GGML_ASSERT(stream_off >= 0); + return (int32_t) (stream_off + pos%state_size); + }; + + for (int32_t i = 0; i < batch.n_tokens; ++i) { + const llama_pos pos = batch.pos[i]; + if (pos < 0) { + continue; + } + + plan.state_pos.push_back((int32_t) (pos%ratio)); + + const int64_t n_visible = (int64_t) (pos + 1)/ratio; + plan.n_visible[(size_t) i] = (int32_t) n_visible; + plan.n_kv = std::max(plan.n_kv, n_visible); + + const int32_t n_token_seqs = + batch.n_seq_id != nullptr && batch.seq_id != nullptr && batch.seq_id[i] != nullptr + ? batch.n_seq_id[i] + : 1; + for (int32_t s = 0; s < n_token_seqs; ++s) { + const llama_seq_id seq_id = + batch.n_seq_id != nullptr && batch.seq_id != nullptr && batch.seq_id[i] != nullptr + ? batch.seq_id[i][s] + : 0; + const int64_t stream_off = dsv4_stream_offset(n_stream, seq_id, state_size); + const int32_t state_idx = (int32_t) (stream_off + pos%state_size); + const auto it = std::find_if(persist_rows.begin(), persist_rows.end(), [state_idx](const persist_row & row) { + return row.dst == state_idx; + }); + if (it == persist_rows.end()) { + persist_rows.push_back({ state_idx, i, pos }); + } else if (pos > it->pos) { + it->src = i; + it->pos = pos; + } + + if ((pos + 1) % ratio != 0) { + continue; + } + + const llama_pos source_start = pos + 1 - ratio; + const int64_t cache_off = dsv4_stream_offset(n_stream, seq_id, kv_size); + plan.state_write_idxs.push_back(cache_off + pos/ratio); + plan.state_write_pos.push_back((int32_t) source_start); + + if (overlap) { + const llama_pos prev_start = source_start - ratio; + for (uint32_t j = 0; j < ratio; ++j) { + overlap_prev_reads.push_back(state_source_idx(seq_id, prev_start + j)); + } + for (uint32_t j = 0; j < ratio; ++j) { + overlap_cur_reads.push_back(state_source_idx(seq_id, source_start + j)); + } + } else { + for (uint32_t j = 0; j < ratio; ++j) { + plan.state_read_idxs.push_back(state_source_idx(seq_id, source_start + j)); + } + } + } + } + + if (ratio == llama_context::dsv4_runtime::CSA_RATIO && plan.state_write_idxs.empty() && !plan.state_pos.empty()) { + const llama_seq_id seq_id0 = + batch.n_seq_id != nullptr && batch.seq_id != nullptr && batch.n_seq_id[0] > 0 && batch.seq_id[0] != nullptr + ? batch.seq_id[0][0] + : 0; + const uint32_t source_idx = (uint32_t) state_source_idx(seq_id0, batch.pos[0]); + const int64_t cache_off = std::max(0, dsv4_stream_offset(n_stream, seq_id0, kv_size)); + plan.state_write_idxs.push_back(cache_off + (int64_t) kv_size - 1); + plan.state_write_pos.push_back(0); + + if (overlap) { + for (uint32_t j = 0; j < ratio; ++j) { + overlap_prev_reads.push_back(source_idx); + overlap_cur_reads.push_back(source_idx); + } + } else { + for (uint32_t j = 0; j < ratio; ++j) { + plan.state_read_idxs.push_back(source_idx); + } + } + } + + if (overlap) { + plan.state_read_idxs.reserve(overlap_prev_reads.size() + overlap_cur_reads.size()); + plan.state_read_idxs.insert(plan.state_read_idxs.end(), overlap_prev_reads.begin(), overlap_prev_reads.end()); + plan.state_read_idxs.insert(plan.state_read_idxs.end(), overlap_cur_reads.begin(), overlap_cur_reads.end()); + } + + plan.n_kv = GGML_PAD(plan.n_kv, 256u); + + std::sort(persist_rows.begin(), persist_rows.end(), [](const persist_row & a, const persist_row & b) { + return a.dst < b.dst; + }); + + for (const persist_row & row : persist_rows) { + plan.state_persist_src_idxs.push_back(row.src); + plan.state_persist_dst_idxs.push_back(row.dst); + } + + if (plan.n_kv == 0) { + plan.n_kv = GGML_PAD(1, 256u); + } + + return plan; +} + +template +static void dsv4_set_input_tensor(ggml_tensor * tensor, const std::vector & values) { + if (tensor == nullptr || tensor->buffer == nullptr || values.empty()) { + return; + } + ggml_backend_tensor_set(tensor, values.data(), 0, values.size()*sizeof(T)); +} + +static void dsv4_set_mask_tensor( + ggml_tensor * tensor, + const llama_context::dsv4_runtime::comp_plan & plan, + int32_t n_tokens) { + if (tensor == nullptr) { + return; + } + + if (tensor->buffer == nullptr) { + return; + } + + const int64_t width = tensor->ne[0]; + const int64_t height = tensor->ne[1]; + auto type = tensor->type; + GGML_ASSERT(type == GGML_TYPE_F16 || type == GGML_TYPE_F32); + + //printf("%s: preparing mask %s of type %s with %ld x %ld entries\n", __func__, tensor->name, ggml_type_name(type), tensor->ne[0], tensor->ne[1]); + if (type == GGML_TYPE_F16) { + auto h_inf = ggml_fp32_to_fp16(-INFINITY); + auto h_zero = ggml_fp32_to_fp16(0.0f); + std::vector storage((size_t) width*height, h_inf); + for (int32_t i = 0; i < n_tokens; ++i) { + const int32_t n_visible = i < (int32_t) plan.n_visible.size() ? plan.n_visible[(size_t) i] : 0; + //if (i == 0) printf(" n_visible = %d\n", n_visible); + for (int32_t j = 0; j < n_visible && j < width; ++j) { + storage[(size_t) i*width + j] = h_zero; + } + } + ggml_backend_tensor_set(tensor, storage.data(), 0, storage.size()*sizeof(ggml_fp16_t)); + } else { + std::vector storage((size_t) width*height, -INFINITY); + for (int32_t i = 0; i < n_tokens; ++i) { + const int32_t n_visible = i < (int32_t) plan.n_visible.size() ? plan.n_visible[(size_t) i] : 0; + for (int32_t j = 0; j < n_visible && j < width; ++j) { + storage[(size_t) i*width + j] = 0.0f; + } + } + ggml_backend_tensor_set(tensor, storage.data(), 0, storage.size()*sizeof(float)); + } +} + +bool llama_context::ensure_dsv4_cache_tensors() { + const int32_t n_layer = model.hparams.n_layer; + const int64_t n_embd_head = model.hparams.n_embd_head_k(0); + const int64_t n_indexer_head = model.hparams.indexer_head_size; + const uint32_t n_stream = std::max(1, cparams.n_seq_max); + const uint32_t csa_kv = GGML_PAD(dsv4_comp_size(cparams.n_ctx, dsv4_runtime::CSA_RATIO), 256u); + const uint32_t hca_kv = GGML_PAD(dsv4_comp_size(cparams.n_ctx, dsv4_runtime::HCA_RATIO), 256u); + + if (!dsv4_validate_cache_type(kv_self.type_k, n_embd_head, "raw/CSA/HCA") || + !dsv4_validate_cache_type(cparams.idx_type_k, n_indexer_head, "LID")) { + return false; + } + + if (dsv4.cache.cache_ctx != nullptr && + (int32_t) dsv4.cache.csa_k.size() == n_layer && + dsv4.cache.n_stream == n_stream) { + return true; + } + + free_dsv4_cache_tensors(); + + ggml_init_params params = { + /*.mem_size =*/ (size_t) (16 * std::max(1, n_layer)) * ggml_tensor_overhead(), + /*.mem_buffer =*/ nullptr, + /*.no_alloc =*/ true, + }; + + dsv4.cache.cache_ctx = ggml_init(params); + if (dsv4.cache.cache_ctx == nullptr) { + LLAMA_LOG_ERROR("%s: failed to allocate DSV4 cache context\n", __func__); + return false; + } + + auto & cache = dsv4.cache; + cache.n_stream = n_stream; + cache.csa_k.resize((size_t) n_layer, nullptr); + cache.hca_k.resize((size_t) n_layer, nullptr); + cache.lid_k.resize((size_t) n_layer, nullptr); + cache.csa_state_kv.resize((size_t) n_layer, nullptr); + cache.csa_state_score.resize((size_t) n_layer, nullptr); + cache.hca_state_kv.resize((size_t) n_layer, nullptr); + cache.hca_state_score.resize((size_t) n_layer, nullptr); + cache.lid_state_kv.resize((size_t) n_layer, nullptr); + cache.lid_state_score.resize((size_t) n_layer, nullptr); + + auto alloc_tensor = [&](ggml_tensor * tensor, ggml_backend_buffer_type_t buft) -> bool { + const size_t tensor_bytes = ggml_backend_buft_get_alloc_size(buft, tensor); + ggml_backend_buffer_t buf = ggml_backend_buft_alloc_buffer(buft, tensor_bytes); + if (buf == nullptr) { + return false; + } + ggml_backend_buffer_set_usage(buf, GGML_BACKEND_BUFFER_USAGE_COMPUTE); + ggml_backend_tensor_alloc(buf, tensor, ggml_backend_buffer_get_base(buf)); + ggml_backend_buffer_clear(buf, 0); + cache.cache_bufs.push_back(buf); + return true; + }; + + for (int32_t il = 0; il < n_layer; ++il) { + const uint32_t ratio = model.hparams.dsv4_compress_ratios[(size_t) il]; + ggml_backend_buffer_type_t buft = llama_dsv4_layer_buft(*this, il); + + if (ratio == dsv4_runtime::CSA_RATIO) { + cache.csa_k[(size_t) il] = ggml_new_tensor_3d(cache.cache_ctx, kv_self.type_k, n_embd_head, csa_kv*n_stream, 1); + cache.lid_k[(size_t) il] = ggml_new_tensor_3d(cache.cache_ctx, cparams.idx_type_k, n_indexer_head, csa_kv*n_stream, 1); + cache.csa_state_kv[(size_t) il] = ggml_new_tensor_2d(cache.cache_ctx, GGML_TYPE_F32, 2*n_embd_head, 2*dsv4_runtime::CSA_RATIO*n_stream); + cache.csa_state_score[(size_t) il] = ggml_new_tensor_2d(cache.cache_ctx, GGML_TYPE_F32, 2*n_embd_head, 2*dsv4_runtime::CSA_RATIO*n_stream); + cache.lid_state_kv[(size_t) il] = ggml_new_tensor_2d(cache.cache_ctx, GGML_TYPE_F32, 2*n_indexer_head, 2*dsv4_runtime::CSA_RATIO*n_stream); + cache.lid_state_score[(size_t) il] = ggml_new_tensor_2d(cache.cache_ctx, GGML_TYPE_F32, 2*n_indexer_head, 2*dsv4_runtime::CSA_RATIO*n_stream); + + if (!alloc_tensor(cache.csa_k[(size_t) il], buft) || + !alloc_tensor(cache.lid_k[(size_t) il], buft) || + !alloc_tensor(cache.csa_state_kv[(size_t) il], buft) || + !alloc_tensor(cache.csa_state_score[(size_t) il], buft) || + !alloc_tensor(cache.lid_state_kv[(size_t) il], buft) || + !alloc_tensor(cache.lid_state_score[(size_t) il], buft)) { + LLAMA_LOG_ERROR("%s: failed to allocate DSV4 CSA/LID buffers for layer %d\n", __func__, il); + free_dsv4_cache_tensors(); + return false; + } + } else if (ratio == dsv4_runtime::HCA_RATIO) { + cache.hca_k[(size_t) il] = ggml_new_tensor_3d(cache.cache_ctx, kv_self.type_k, n_embd_head, hca_kv*n_stream, 1); + cache.hca_state_kv[(size_t) il] = ggml_new_tensor_2d(cache.cache_ctx, GGML_TYPE_F32, n_embd_head, dsv4_runtime::HCA_RATIO*n_stream); + cache.hca_state_score[(size_t) il] = ggml_new_tensor_2d(cache.cache_ctx, GGML_TYPE_F32, n_embd_head, dsv4_runtime::HCA_RATIO*n_stream); + + if (!alloc_tensor(cache.hca_k[(size_t) il], buft) || + !alloc_tensor(cache.hca_state_kv[(size_t) il], buft) || + !alloc_tensor(cache.hca_state_score[(size_t) il], buft)) { + LLAMA_LOG_ERROR("%s: failed to allocate DSV4 HCA buffers for layer %d\n", __func__, il); + free_dsv4_cache_tensors(); + return false; + } + } + } + + auto bytes = [](const auto & tensors) { + size_t total = 0; + for (const ggml_tensor * tensor : tensors) { + if (tensor != nullptr) { + total += ggml_nbytes(tensor); + } + } + return total; + }; + + const size_t csa_k_bytes = bytes(cache.csa_k); + const size_t hca_k_bytes = bytes(cache.hca_k); + const size_t lid_k_bytes = bytes(cache.lid_k); + const size_t csa_state_bytes = bytes(cache.csa_state_kv) + bytes(cache.csa_state_score); + const size_t hca_state_bytes = bytes(cache.hca_state_kv) + bytes(cache.hca_state_score); + const size_t lid_state_bytes = bytes(cache.lid_state_kv) + bytes(cache.lid_state_score); + + LLAMA_LOG_INFO("%s: DSV4 cache: CSA K=%7.2f MiB (%s), HCA K=%7.2f MiB (%s), LID K=%7.2f MiB (%s), states=%7.2f MiB, total=%7.2f MiB, streams=%u\n", + __func__, + (float) csa_k_bytes / (1024.0f * 1024.0f), ggml_type_name(kv_self.type_k), + (float) hca_k_bytes / (1024.0f * 1024.0f), ggml_type_name(kv_self.type_k), + (float) lid_k_bytes / (1024.0f * 1024.0f), ggml_type_name(cparams.idx_type_k), + (float) (csa_state_bytes + hca_state_bytes + lid_state_bytes) / (1024.0f * 1024.0f), + (float) (csa_k_bytes + hca_k_bytes + lid_k_bytes + csa_state_bytes + hca_state_bytes + lid_state_bytes) / (1024.0f * 1024.0f), + n_stream); + + return true; +} + +void llama_context::free_dsv4_cache_tensors() { + auto release_vector = [](auto & v) { + using vec_type = std::decay_t; + vec_type().swap(v); + }; + + for (ggml_backend_buffer_t buf : dsv4.cache.cache_bufs) { + if (buf != nullptr) { + ggml_backend_buffer_free(buf); + } + } + release_vector(dsv4.cache.cache_bufs); + release_vector(dsv4.cache.csa_k); + release_vector(dsv4.cache.hca_k); + release_vector(dsv4.cache.lid_k); + release_vector(dsv4.cache.csa_state_kv); + release_vector(dsv4.cache.csa_state_score); + release_vector(dsv4.cache.hca_state_kv); + release_vector(dsv4.cache.hca_state_score); + release_vector(dsv4.cache.lid_state_kv); + release_vector(dsv4.cache.lid_state_score); + dsv4.cache.n_stream = 1; + if (dsv4.cache.cache_ctx != nullptr) { + ggml_free(dsv4.cache.cache_ctx); + dsv4.cache.cache_ctx = nullptr; + } +} + +void llama_reset_dsv4_state(llama_context * ctx, int32_t seq_id) { + if (ctx == nullptr) { + return; + } + + const uint32_t n_stream = std::max(1, ctx->dsv4.cache.n_stream); + if (seq_id >= (llama_seq_id) n_stream) { + LLAMA_LOG_ERROR("%s: DSV4 seq_id %d is outside stream range %u\n", __func__, seq_id, n_stream); + return; + } + + if (seq_id < 0) { + for (ggml_backend_buffer_t buf : ctx->dsv4.cache.cache_bufs) { + ggml_backend_buffer_clear(buf, 0); + } + return; + } + + auto clear_tensor = [seq_id, n_stream](ggml_tensor * tensor) { + if (tensor == nullptr) { + return; + } + + GGML_ASSERT(tensor->ne[1] % n_stream == 0); + const size_t row_bytes = tensor->nb[1]; + const size_t rows_per_stream = (size_t) tensor->ne[1] / n_stream; + const size_t offset = (size_t) seq_id * rows_per_stream * row_bytes; + const size_t bytes = rows_per_stream * row_bytes; + std::vector zeros(bytes, 0); + ggml_backend_tensor_set(tensor, zeros.data(), offset, bytes); + }; + + for (ggml_tensor * tensor : ctx->dsv4.cache.csa_k) clear_tensor(tensor); + for (ggml_tensor * tensor : ctx->dsv4.cache.hca_k) clear_tensor(tensor); + for (ggml_tensor * tensor : ctx->dsv4.cache.lid_k) clear_tensor(tensor); + for (ggml_tensor * tensor : ctx->dsv4.cache.csa_state_kv) clear_tensor(tensor); + for (ggml_tensor * tensor : ctx->dsv4.cache.csa_state_score) clear_tensor(tensor); + for (ggml_tensor * tensor : ctx->dsv4.cache.hca_state_kv) clear_tensor(tensor); + for (ggml_tensor * tensor : ctx->dsv4.cache.hca_state_score) clear_tensor(tensor); + for (ggml_tensor * tensor : ctx->dsv4.cache.lid_state_kv) clear_tensor(tensor); + for (ggml_tensor * tensor : ctx->dsv4.cache.lid_state_score) clear_tensor(tensor); +} + +bool llama_prepare_dsv4_graph_inputs(llama_context & lctx, const llama_batch & batch, bool set_tensors, bool reserve_plan) { + if (lctx.model.arch != LLM_ARCH_DEEPSEEK4) { + return true; + } + + if (!dsv4_validate_batch_seq_ids(lctx, batch)) { + return false; + } + + if (!lctx.ensure_dsv4_cache_tensors()) { + return false; + } + + const uint32_t cache_n_stream = std::max(1, lctx.dsv4.cache.n_stream); + const uint32_t csa_kv_size = dsv4_cache_kv_size(lctx.dsv4.cache.csa_k)/cache_n_stream; + const uint32_t hca_kv_size = dsv4_cache_kv_size(lctx.dsv4.cache.hca_k)/cache_n_stream; + const uint32_t lid_kv_size = dsv4_cache_kv_size(lctx.dsv4.cache.lid_k)/cache_n_stream; + const uint32_t csa_state_size = dsv4_cache_state_size(lctx.dsv4.cache.csa_state_kv)/cache_n_stream; + const uint32_t hca_state_size = dsv4_cache_state_size(lctx.dsv4.cache.hca_state_kv)/cache_n_stream; + const uint32_t lid_state_size = dsv4_cache_state_size(lctx.dsv4.cache.lid_state_kv)/cache_n_stream; + + const auto build_plan = [&](uint32_t ratio, bool overlap, uint32_t state_size, uint32_t kv_size, uint32_t n_stream) { + return reserve_plan + ? dsv4_build_reserve_comp_plan(batch, ratio, overlap, state_size, kv_size, n_stream) + : dsv4_build_comp_plan(batch, ratio, overlap, state_size, kv_size, n_stream); + }; + + lctx.dsv4.raw = {}; + if (!reserve_plan && !dsv4_build_raw_context(lctx, batch, lctx.dsv4.raw)) { + return false; + } + + //auto tim1 = ggml_time_us(); + lctx.dsv4.csa_plan = build_plan(llama_context::dsv4_runtime::CSA_RATIO, true, csa_state_size, csa_kv_size, cache_n_stream); + lctx.dsv4.hca_plan = build_plan(llama_context::dsv4_runtime::HCA_RATIO, false, hca_state_size, hca_kv_size, cache_n_stream); + lctx.dsv4.lid_plan = build_plan(llama_context::dsv4_runtime::CSA_RATIO, true, lid_state_size, lid_kv_size, cache_n_stream); + lctx.dsv4.csa_ctx = dsv4_build_comp_context(batch, cache_n_stream, lctx.dsv4.csa_plan.n_kv); + lctx.dsv4.hca_ctx = dsv4_build_comp_context(batch, cache_n_stream, lctx.dsv4.hca_plan.n_kv); + lctx.dsv4.lid_ctx = dsv4_build_comp_context(batch, cache_n_stream, lctx.dsv4.lid_plan.n_kv); + //auto tim2 = ggml_time_us(); + //fprintf(stderr, "%s: %ld us to buils plans\n", __func__, tim2-tim1); + + if (!dsv4_validate_comp_plan("csa", batch, lctx.dsv4.csa_plan, llama_context::dsv4_runtime::CSA_RATIO, true, csa_state_size, csa_kv_size, cache_n_stream) || + !dsv4_validate_comp_plan("hca", batch, lctx.dsv4.hca_plan, llama_context::dsv4_runtime::HCA_RATIO, false, hca_state_size, hca_kv_size, cache_n_stream) || + !dsv4_validate_comp_plan("lid", batch, lctx.dsv4.lid_plan, llama_context::dsv4_runtime::CSA_RATIO, true, lid_state_size, lid_kv_size, cache_n_stream) || + !dsv4_validate_csa_lid_visibility(lctx, csa_kv_size, lid_kv_size)) { + return false; + } + + if (!set_tensors) { + return true; + } + + //tim1 = ggml_time_us(); + + dsv4_set_input_tensor(lctx.dsv4.inputs.raw_k_write_src_idxs, lctx.dsv4.raw.write_src_idxs); + dsv4_set_input_tensor(lctx.dsv4.inputs.raw_k_write_idxs, lctx.dsv4.raw.write_dst_idxs); + dsv4_set_input_tensor(lctx.dsv4.inputs.raw_k_read_idxs, lctx.dsv4.raw.read_dst_idxs); + + auto set_comp = [&](llama_context::dsv4_runtime::comp_inputs & inputs, llama_context::dsv4_runtime::comp_plan & plan, bool set_mask) { + dsv4_set_input_tensor(inputs.state_pos, plan.state_pos); + dsv4_set_input_tensor(inputs.state_persist_src_idxs, plan.state_persist_src_idxs); + dsv4_set_input_tensor(inputs.state_persist_dst_idxs, plan.state_persist_dst_idxs); + dsv4_set_input_tensor(inputs.state_read_idxs, plan.state_read_idxs); + dsv4_set_input_tensor(inputs.state_write_idxs, plan.state_write_idxs); + dsv4_set_input_tensor(inputs.state_write_pos, plan.state_write_pos); + if (set_mask) { + dsv4_set_mask_tensor(inputs.kq_mask, plan, batch.n_tokens); + } + }; + + set_comp(lctx.dsv4.inputs.csa, lctx.dsv4.csa_plan, true); + set_comp(lctx.dsv4.inputs.hca, lctx.dsv4.hca_plan, true); + set_comp(lctx.dsv4.inputs.lid, lctx.dsv4.lid_plan, false); + + //tim2 = ggml_time_us(); + //fprintf(stderr, "%s: setting input tensors took %ld us\n", __func__, tim2 - tim1); + + return true; +} diff --git a/src/llama-dsv4.h b/src/llama-dsv4.h new file mode 100644 index 00000000..43cd130c --- /dev/null +++ b/src/llama-dsv4.h @@ -0,0 +1,8 @@ +#pragma once + +#include +struct llama_batch; +struct llama_context; + +bool llama_prepare_dsv4_graph_inputs(llama_context & lctx, const llama_batch & batch, bool set_tensors, bool reserve_plan); +void llama_reset_dsv4_state(llama_context * ctx, int32_t seq_id = -1); diff --git a/src/llama-hparams.cpp b/src/llama-hparams.cpp index cc32e796..d5bad468 100644 --- a/src/llama-hparams.cpp +++ b/src/llama-hparams.cpp @@ -3,6 +3,7 @@ #include "llama-model-loader.h" #include "llama-model.h" +#include #include #include @@ -33,6 +34,7 @@ static inline const char * llm_expert_gating_func_name(llm_expert_gating_func_ty case LLM_EXPERT_GATING_FUNC_SOFTMAX: return "softmax"; case LLM_EXPERT_GATING_FUNC_SIGMOID: return "sigmoid"; case LLM_EXPERT_GATING_FUNC_TYPE_SOFTMAX_WEIGHT: return "weight"; + case LLM_EXPERT_GATING_FUNC_TYPE_SQRT_SOFTPLUS: return "sqrtsoftplus"; default: return "none"; } } @@ -1650,6 +1652,7 @@ void llm_load_hparams( default: model.type = e_model::MODEL_UNKNOWN; } } break; + case LLM_ARCH_DEEPSEEK4: case LLM_ARCH_GLM_DSA: { ml.get_key(LLM_KV_EXPERT_FEED_FORWARD_LENGTH, hparams.n_ff_exp); @@ -1672,9 +1675,48 @@ void llm_load_hparams( ml.get_key(LLM_KV_EXPERT_WEIGHTS_SCALE, hparams.expert_weights_scale); ml.get_key(LLM_KV_EXPERT_WEIGHTS_NORM, hparams.expert_weights_norm, false); - // deepseek MLA parameters + if (model.arch == LLM_ARCH_DEEPSEEK4) { + ml.get_key_or_arr(LLM_KV_SWIGLU_CLAMP_EXP, hparams.swiglu_limits, hparams.n_layer); + if (!ml.get_key_or_arr(LLM_KV_SWIGLU_CLAMP_SHEXP, hparams.swiglu_limits_shared, hparams.n_layer, 0)) { + hparams.swiglu_limits_shared = hparams.swiglu_limits; + } + } + + // Shared latent-attention ranks used by common hparams and + // cache helpers. DSV4 has its own CSA/HCA execution graph; + // this does not select the DeepSeek V3 MLA path. ml.get_key(LLM_KV_ATTENTION_Q_LORA_RANK, hparams.n_lora_q); - ml.get_key(LLM_KV_ATTENTION_KV_LORA_RANK, hparams.n_lora_kv); + ml.get_key(LLM_KV_ATTENTION_KV_LORA_RANK, hparams.n_lora_kv, false); + if (model.arch == LLM_ARCH_DEEPSEEK4 && hparams.n_lora_kv == 0) { + if (auto * kv_norm = ml.get_tensor_meta("blk.0.attn_kv_a_norm.weight")) { + hparams.n_lora_kv = (uint32_t) kv_norm->ne[0]; + } else if (auto * kv = ml.get_tensor_meta("blk.0.attn_kv.weight")) { + const int64_t kv_inner = kv->ne[0] == hparams.n_embd ? kv->ne[1] : kv->ne[0]; + hparams.n_lora_kv = (uint32_t) kv_inner; + } else { + auto * kv_a = ml.get_tensor_meta("blk.0.attn_kv_latent.weight"); + bool subtract_rope = false; + if (kv_a == nullptr) { + kv_a = ml.require_tensor_meta("blk.0.attn_kv_a_mqa.weight"); + subtract_rope = true; + } + + const int64_t kv_a_inner = kv_a->ne[0] == hparams.n_embd ? kv_a->ne[1] : kv_a->ne[0]; + if (subtract_rope) { + if (kv_a_inner <= hparams.n_rot) { + throw std::runtime_error(format( + "%s: unable to infer %s from blk.0.attn_kv_a_mqa.weight shape [%lld, %lld]", + __func__, + ml.llm_kv(LLM_KV_ATTENTION_KV_LORA_RANK).c_str(), + (long long) kv_a->ne[0], + (long long) kv_a->ne[1])); + } + hparams.n_lora_kv = (uint32_t) (kv_a_inner - hparams.n_rot); + } else { + hparams.n_lora_kv = (uint32_t) kv_a_inner; + } + } + } //ml.get_key(LLM_KV_ATTENTION_KEY_LENGTH_MLA, hparams.n_embd_head_k_mla_impl, false); //ml.get_key(LLM_KV_ATTENTION_VALUE_LENGTH_MLA, hparams.n_embd_head_v_mla_impl, false); ml.get_key(LLM_KV_EXPERT_FEED_FORWARD_LENGTH, hparams.n_ff_exp); @@ -1685,13 +1727,86 @@ void llm_load_hparams( ml.get_key(LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, hparams.indexer_head_size); ml.get_key(LLM_KV_ATTENTION_INDEXER_TOP_K, hparams.indexer_top_k); - // GLM-5.2 IndexShare: per-layer full/shared indexer map. "full" layers compute their own - // top-k; "shared" layers reuse the previous full layer's selection (transformers - // modeling_glm_moe_dsa.py: shared layer indexer=None, topk_indices=prev_topk_indices). - // Derived from GLM-5.2's config indexer_types rule (full iff il<=1 or il%4==2), verified - // to reproduce the config's full set {0,1,2,6,10,...} exactly. Existing GGUFs carry no - // per-layer metadata, so the derivation is the source of truth; a future metadata key can - // override this for GLM-DSA variants with a different pattern. + if (model.arch == LLM_ARCH_DEEPSEEK4) { + ml.get_key(LLM_KV_ATTENTION_SLIDING_WINDOW, hparams.n_swa, false); + ml.get_key(LLM_KV_ROPE_FREQ_BASE_SWA, hparams.rope_freq_base_train_swa, false); + hparams.rope_freq_scale_train_swa = 1.0f; + if (!ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, hparams.swa_layers, hparams.n_layer, false) && hparams.n_swa > 0) { + std::fill(hparams.swa_layers.begin(), hparams.swa_layers.end(), true); + } + + const auto * hc_head_base = ml.get_tensor_meta("hc_head_base"); + const auto * wo_a_0 = ml.get_tensor_meta("blk.0.attn_output_a.weight"); + const auto * wo_b_0 = ml.get_tensor_meta("blk.0.attn_output_b.weight"); + + if (!ml.get_key(LLM_KV_ATTENTION_OUTPUT_GROUP_COUNT, hparams.dsv4_o_group_count, false) && wo_a_0 != nullptr) { + GGML_ASSERT(wo_a_0->ne[0] > 0); + hparams.dsv4_o_group_count = (uint32_t) ((hparams.n_head() * hparams.n_embd_head_k(0)) / wo_a_0->ne[0]); + } + if (!ml.get_key(LLM_KV_ATTENTION_OUTPUT_LORA_RANK, hparams.dsv4_o_lora_rank, false) && wo_b_0 != nullptr) { + GGML_ASSERT(hparams.dsv4_o_group_count > 0); + hparams.dsv4_o_lora_rank = (uint32_t) (wo_b_0->ne[0] / hparams.dsv4_o_group_count); + } + if (!ml.get_key(LLM_KV_ATTENTION_COMPRESS_ROPE_FREQ_BASE, hparams.dsv4_compress_rope_base, false)) { + hparams.dsv4_compress_rope_base = hparams.rope_freq_base_train_swa != 0.0f + ? hparams.rope_freq_base_train_swa + : hparams.rope_freq_base_train; + } + if (!ml.get_key(LLM_KV_HYPER_CONNECTION_COUNT, hparams.dsv4_hc_mult, false)) { + if (hc_head_base != nullptr) { + hparams.dsv4_hc_mult = (uint32_t) hc_head_base->ne[0]; + } else if (wo_a_0 != nullptr) { + hparams.dsv4_hc_mult = (uint32_t) (wo_a_0->ne[1] / hparams.n_embd); + } + } + if (!ml.get_key(LLM_KV_HYPER_CONNECTION_SINKHORN_ITERATIONS, hparams.dsv4_hc_sinkhorn_iters, false)) { + hparams.dsv4_hc_sinkhorn_iters = 3; + } + if (!ml.get_key(LLM_KV_HYPER_CONNECTION_EPSILON, hparams.dsv4_hc_eps, false)) { + hparams.dsv4_hc_eps = hparams.f_norm_rms_eps; + } + ml.get_key(LLM_KV_HASH_LAYER_COUNT, hparams.dsv4_hash_layer_count, false); + + uint32_t n_compress_ratios = 0; + if (ml.get_arr_n(LLM_KV_ATTENTION_COMPRESS_RATIOS, n_compress_ratios, false)) { + if (n_compress_ratios < hparams.n_layer) { + throw std::runtime_error("DeepSeek-V4 compress_ratios is shorter than block_count"); + } + std::vector compress_ratios; + ml.get_arr(ml.llm_kv(LLM_KV_ATTENTION_COMPRESS_RATIOS), compress_ratios); + std::copy_n(compress_ratios.begin(), hparams.n_layer, hparams.dsv4_compress_ratios.begin()); + } else { + for (uint32_t il = 0; il < hparams.n_layer; ++il) { + const bool has_attn_compress = + ml.get_tensor_meta(format("blk.%u.attn_compress_kv.weight", il).c_str()) != nullptr || + ml.get_tensor_meta(format("blk.%u.attn_compressor_kv.weight", il).c_str()) != nullptr; + const bool has_indexer = + ml.get_tensor_meta(format("blk.%u.indexer.attn_q_b.weight", il).c_str()) != nullptr || + ml.get_tensor_meta(format("blk.%u.indexer.compress_kv.weight", il).c_str()) != nullptr || + ml.get_tensor_meta(format("blk.%u.indexer_compressor_kv.weight", il).c_str()) != nullptr; + + if (has_indexer && !has_attn_compress) { + throw std::runtime_error(format("DeepSeek-V4 layer %u has indexer tensors without attention compressor tensors", il)); + } + + hparams.dsv4_compress_ratios[il] = has_indexer ? 4 : (has_attn_compress ? 128 : 0); + } + } + + if (hparams.dsv4_hc_mult == 0) { + throw std::runtime_error("DeepSeek-V4 hyper_connection.count is missing and could not be inferred"); + } + if (hparams.dsv4_o_group_count == 0 || hparams.dsv4_o_lora_rank == 0) { + throw std::runtime_error("DeepSeek-V4 output projection metadata is missing and could not be inferred"); + } + if (wo_b_0 != nullptr && (int64_t) hparams.dsv4_o_group_count * hparams.dsv4_o_lora_rank != wo_b_0->ne[0]) { + throw std::runtime_error("DeepSeek-V4 inferred output_group_count/output_lora_rank does not match attn_output_b shape"); + } + if (wo_a_0 != nullptr && (int64_t) hparams.dsv4_o_group_count * wo_a_0->ne[0] != (int64_t) hparams.n_head() * hparams.n_embd_head_k(0)) { + throw std::runtime_error("DeepSeek-V4 inferred output_group_count does not match attn_output_a shape"); + } + } + for (uint32_t il = 0; il < hparams.n_layer; ++il) { hparams.indexer_is_full[il] = (il <= 1) || (il % 4 == 2); } @@ -1702,21 +1817,26 @@ void llm_load_hparams( hparams.expert_gating_func = LLM_EXPERT_GATING_FUNC_SIGMOID; } - // NextN/MTP parameters - ml.get_key(LLM_KV_NEXTN_PREDICT_LAYERS, hparams.nextn_predict_layers, false); - - if (model.mtp) { - hparams.n_layer_kv_from_start = hparams.n_layer; + if (model.arch == LLM_ARCH_DEEPSEEK4 && + hparams.expert_gating_func != LLM_EXPERT_GATING_FUNC_TYPE_SQRT_SOFTPLUS) { + throw std::runtime_error("DeepSeek-V4 loader currently expects sqrtsoftplus MoE scoring"); } - else { - hparams.n_layer_kv_from_start = hparams.n_layer - hparams.nextn_predict_layers; + + if (model.arch == LLM_ARCH_DEEPSEEK4) { + hparams.n_layer_kv_from_start = hparams.n_layer; + } else { + ml.get_key(LLM_KV_NEXTN_PREDICT_LAYERS, hparams.nextn_predict_layers, false); + hparams.n_layer_kv_from_start = model.mtp + ? hparams.n_layer + : hparams.n_layer - hparams.nextn_predict_layers; } switch (hparams.n_layer) { + case 61: model.type = MODEL_290B; break; case 79: model.type = MODEL_744B_A40B; break; default: model.type = MODEL_UNKNOWN; } - if (hparams.n_head_kv() == 1) { + if (model.arch != LLM_ARCH_DEEPSEEK4 && hparams.n_head_kv() == 1) { int n_nead_kv = hparams.n_gqa(); if (n_nead_kv%4 != 0 || hparams.n_embd_head_k_full != 576 || hparams.n_embd_head_v_full != 512 || hparams.n_rot != 64) { diff --git a/src/llama-hparams.h b/src/llama-hparams.h index c52c842f..f21272dd 100644 --- a/src/llama-hparams.h +++ b/src/llama-hparams.h @@ -13,6 +13,7 @@ enum llm_expert_gating_func_type { LLM_EXPERT_GATING_FUNC_SOFTMAX = 1, LLM_EXPERT_GATING_FUNC_SIGMOID = 2, LLM_EXPERT_GATING_FUNC_TYPE_SOFTMAX_WEIGHT = 3, + LLM_EXPERT_GATING_FUNC_TYPE_SQRT_SOFTPLUS = 4, }; struct llama_hparams { @@ -141,6 +142,16 @@ struct llama_hparams { uint32_t n_swa_mtp = 0; std::array openpangu_window = {}; + // DeepSeek-V4 + uint32_t dsv4_o_group_count = 0; + uint32_t dsv4_o_lora_rank = 0; + uint32_t dsv4_hc_mult = 0; + uint32_t dsv4_hc_sinkhorn_iters = 0; + uint32_t dsv4_hash_layer_count = 0; + float dsv4_compress_rope_base = 0.0f; + float dsv4_hc_eps = 0.0f; + std::array dsv4_compress_ratios = {}; + // qwen3vl deepstack uint32_t n_deepstack_layers = 0; diff --git a/src/llama-load-tensors.cpp b/src/llama-load-tensors.cpp index 8aa5a02c..3ece0f96 100644 --- a/src/llama-load-tensors.cpp +++ b/src/llama-load-tensors.cpp @@ -119,6 +119,7 @@ struct create_tensors_helper : public create_tensors_helper_interface { bool create_arctix_tensors(const LLM_TN & tn); bool create_deepseek2_tensors(const LLM_TN & tn); + bool create_deepseek4_tensors(const LLM_TN & tn); bool create_openpangu_tensors(const LLM_TN & tn); bool create_glm_dsa_tensors(const LLM_TN & tn); @@ -2733,6 +2734,137 @@ bool create_tensors_helper::create_deepseek2_tensors(const LLM_TN & tn) { return use_mmap_buffer; } +bool create_tensors_helper::create_deepseek4_tensors(const LLM_TN &) { + LOADING_PRELUDE + + auto create_tensor_from_meta = [&](ggml_context * ctx, const std::string & name, int flags = 0) -> ggml_tensor * { + ggml_tensor * meta = (flags & (llama_model_loader::TENSOR_NOT_REQUIRED | llama_model_loader::TENSOR_SKIP)) + ? ml.get_tensor_meta(name.c_str()) + : ml.require_tensor_meta(name.c_str()); + if (meta == nullptr) { + return nullptr; + } + + std::vector ne; + const int n_dims = ggml_n_dims(meta); + ne.reserve(n_dims); + for (int d = 0; d < n_dims; ++d) { + ne.push_back(meta->ne[d]); + } + + return create_tensor(ctx, name, ne, flags); + }; + + auto pick_tensor_name = [&](std::initializer_list candidates) -> std::string { + for (const auto & candidate : candidates) { + if (ml.get_tensor_meta(candidate.c_str()) != nullptr) { + return candidate; + } + } + return *candidates.begin(); + }; + + auto layer_weight_name = [](int i, const char * stem) { + return format("blk.%d.%s.weight", i, stem); + }; + + model.tok_embd = create_tensor_from_meta(ctx_input, "token_embd.weight"); + model.output_norm = create_tensor_from_meta(ctx_output, "output_norm.weight"); + model.output = create_tensor_from_meta(ctx_output, "output.weight"); + + model.hc_head_base = create_tensor_from_meta(ctx_output, pick_tensor_name({"hc_head_base.weight", "output_hc_base.weight"}), llama_model_loader::TENSOR_NOT_REQUIRED); + model.hc_head_fn = create_tensor_from_meta(ctx_output, pick_tensor_name({"hc_head_fn.weight", "output_hc_fn.weight"}), llama_model_loader::TENSOR_NOT_REQUIRED); + model.hc_head_scale = create_tensor_from_meta(ctx_output, pick_tensor_name({"hc_head_scale.weight", "output_hc_scale.weight"}), llama_model_loader::TENSOR_NOT_REQUIRED); + + for (int i = 0; i < n_layer; ++i) { + ggml_context * ctx_split = ctx_for_layer_split(i); + auto & layer = model.layers[i]; + + layer.attn_norm = create_tensor_from_meta(ctx_split, layer_weight_name(i, "attn_norm")); + layer.attn_sinks = create_tensor_from_meta(ctx_split, format("blk.%d.attn_sinks.weight", i), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.wq_a = create_tensor_from_meta(ctx_split, layer_weight_name(i, "attn_q_a")); + layer.attn_q_a_norm = create_tensor_from_meta(ctx_split, layer_weight_name(i, "attn_q_a_norm")); + layer.wq_b = create_tensor_from_meta(ctx_split, layer_weight_name(i, "attn_q_b")); + layer.wkv_latent = create_tensor_from_meta(ctx_split, pick_tensor_name({ + layer_weight_name(i, "attn_kv_latent"), + layer_weight_name(i, "attn_kv"), + layer_weight_name(i, "attn_kv_a_mqa"), + })); + layer.wkv_b = layer.wkv_latent; + layer.wkv_a_mqa = layer.wkv_latent; + layer.attn_kv_a_norm = create_tensor_from_meta(ctx_split, layer_weight_name(i, "attn_kv_a_norm")); + layer.attn_kv_norm = layer.attn_kv_a_norm; + layer.wo_a = create_tensor_from_meta(ctx_split, layer_weight_name(i, "attn_output_a")); + layer.wo_b = create_tensor_from_meta(ctx_split, layer_weight_name(i, "attn_output_b")); + layer.wo = layer.wo_b; + + layer.hc_attn_base = create_tensor_from_meta(ctx_split, format("blk.%d.hc_attn_base.weight", i), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.hc_attn_fn = create_tensor_from_meta(ctx_split, format("blk.%d.hc_attn_fn.weight", i), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.hc_attn_scale = create_tensor_from_meta(ctx_split, format("blk.%d.hc_attn_scale.weight", i), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.hc_ffn_base = create_tensor_from_meta(ctx_split, format("blk.%d.hc_ffn_base.weight", i), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.hc_ffn_fn = create_tensor_from_meta(ctx_split, format("blk.%d.hc_ffn_fn.weight", i), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.hc_ffn_scale = create_tensor_from_meta(ctx_split, format("blk.%d.hc_ffn_scale.weight", i), llama_model_loader::TENSOR_NOT_REQUIRED); + + layer.attn_comp_wkv = create_tensor_from_meta(ctx_split, pick_tensor_name({ + layer_weight_name(i, "attn_compress_kv"), + layer_weight_name(i, "attn_compressor_kv"), + }), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.attn_comp_wgate = create_tensor_from_meta(ctx_split, pick_tensor_name({ + layer_weight_name(i, "attn_compress_gate"), + layer_weight_name(i, "attn_compressor_gate"), + }), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.attn_comp_ape = create_tensor_from_meta(ctx_split, pick_tensor_name({ + format("blk.%d.attn_compress_ape.weight", i), + format("blk.%d.attn_compressor_ape.weight", i), + }), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.attn_comp_norm = create_tensor_from_meta(ctx_split, pick_tensor_name({ + layer_weight_name(i, "attn_compress_norm"), + layer_weight_name(i, "attn_compressor_norm"), + }), llama_model_loader::TENSOR_NOT_REQUIRED); + + layer.indexer_k_norm = create_tensor_from_meta(ctx_split, layer_weight_name(i, "indexer.k_norm"), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.indexer_attn_k = create_tensor_from_meta(ctx_split, layer_weight_name(i, "indexer.attn_k"), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.indexer_proj = create_tensor_from_meta(ctx_split, layer_weight_name(i, "indexer.proj"), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.indexer_attn_q_b = create_tensor_from_meta(ctx_split, layer_weight_name(i, "indexer.attn_q_b"), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.indexer_comp_wkv = create_tensor_from_meta(ctx_split, pick_tensor_name({ + layer_weight_name(i, "indexer.compress_kv"), + layer_weight_name(i, "indexer_compressor_kv"), + }), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.indexer_comp_wgate = create_tensor_from_meta(ctx_split, pick_tensor_name({ + layer_weight_name(i, "indexer.compress_gate"), + layer_weight_name(i, "indexer_compressor_gate"), + }), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.indexer_comp_ape = create_tensor_from_meta(ctx_split, pick_tensor_name({ + format("blk.%d.indexer.compress_ape.weight", i), + format("blk.%d.indexer_compressor_ape.weight", i), + }), llama_model_loader::TENSOR_NOT_REQUIRED); + layer.indexer_comp_norm = create_tensor_from_meta(ctx_split, pick_tensor_name({ + layer_weight_name(i, "indexer.compress_norm"), + layer_weight_name(i, "indexer_compressor_norm"), + }), llama_model_loader::TENSOR_NOT_REQUIRED); + + layer.ffn_gate_inp = create_tensor_from_meta(ctx_split, layer_weight_name(i, "ffn_gate_inp")); + layer.ffn_norm = create_tensor_from_meta(ctx_split, layer_weight_name(i, "ffn_norm")); + layer.ffn_gate_exps = create_tensor_from_meta(ctx_split, layer_weight_name(i, "ffn_gate_exps")); + layer.ffn_down_exps = create_tensor_from_meta(ctx_split, layer_weight_name(i, "ffn_down_exps")); + layer.ffn_up_exps = create_tensor_from_meta(ctx_split, layer_weight_name(i, "ffn_up_exps")); + layer.ffn_gate_shexp = create_tensor_from_meta(ctx_split, layer_weight_name(i, "ffn_gate_shexp")); + layer.ffn_down_shexp = create_tensor_from_meta(ctx_split, layer_weight_name(i, "ffn_down_shexp")); + layer.ffn_up_shexp = create_tensor_from_meta(ctx_split, layer_weight_name(i, "ffn_up_shexp")); + + layer.ffn_gate_tid2eid = create_tensor_from_meta(ctx_split, format("blk.%d.ffn_gate_tid2eid.weight", i), llama_model_loader::TENSOR_NOT_REQUIRED); + if (layer.ffn_gate_tid2eid == nullptr) { + layer.ffn_exp_probs_b = create_tensor_from_meta(ctx_split, pick_tensor_name({ + format("blk.%d.exp_probs_b.bias", i), + format("blk.%d.exp_probs_b.weight", i), + }), llama_model_loader::TENSOR_NOT_REQUIRED); + } + + } + + return use_mmap_buffer; +} + bool create_tensors_helper::create_glm_dsa_tensors(const LLM_TN & tn) { LOADING_PRELUDE @@ -4679,6 +4811,8 @@ bool create_tensors_helper::create_tensors() { case LLM_ARCH_DEEPSEEK2: case LLM_ARCH_MISTRAL4: use_mmap_buffer = create_deepseek2_tensors(tn); break; + case LLM_ARCH_DEEPSEEK4: + use_mmap_buffer = create_deepseek4_tensors(tn); break; case LLM_ARCH_GLM_DSA: use_mmap_buffer = create_glm_dsa_tensors(tn); break; case LLM_ARCH_OPENPANGU: diff --git a/src/llama-model-loader.cpp b/src/llama-model-loader.cpp index dc544c8c..ff981634 100644 --- a/src/llama-model-loader.cpp +++ b/src/llama-model-loader.cpp @@ -1391,6 +1391,6 @@ template bool llama_model_loader::get_key_or_arr>(enum ll template std::enable_if::value, bool>::type llama_model_loader::get_arr_n(const std::string &, unsigned int &, bool); template std::enable_if::value, bool>::type llama_model_loader::get_arr_n(enum llm_kv, unsigned int&, bool); +template bool llama_model_loader::get_arr(const std::string &, std::vector &, bool); template bool llama_model_loader::get_arr(const std::string &, std::array &, bool); template bool llama_model_loader::get_arr(const std::string &, std::array &, bool); -template bool llama_model_loader::get_arr(const std::string &, std::vector &, bool); diff --git a/src/llama-model.cpp b/src/llama-model.cpp index f9048b60..eab0e9b9 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -1064,6 +1064,67 @@ static const std::map> LLM_TENSOR_NA { LLM_TENSOR_FFN_EXP_PROBS_B, "blk.%d.exp_probs_b" }, }, }, + { + LLM_ARCH_DEEPSEEK4, + { + { LLM_TENSOR_TOKEN_EMBD, "token_embd" }, + { LLM_TENSOR_OUTPUT_NORM, "output_norm" }, + { LLM_TENSOR_OUTPUT, "output" }, + { LLM_TENSOR_ATTN_NORM, "blk.%d.attn_norm" }, + { LLM_TENSOR_ATTN_SINKS, "blk.%d.attn_sinks" }, + { LLM_TENSOR_ATTN_Q_A_NORM, "blk.%d.attn_q_a_norm" }, + { LLM_TENSOR_ATTN_KV_A_NORM, "blk.%d.attn_kv_a_norm" }, + { LLM_TENSOR_ATTN_Q, "blk.%d.attn_q" }, + { LLM_TENSOR_ATTN_Q_A, "blk.%d.attn_q_a" }, + { LLM_TENSOR_ATTN_Q_B, "blk.%d.attn_q_b" }, + { LLM_TENSOR_ATTN_KV_LATENT, "blk.%d.attn_kv" }, + { LLM_TENSOR_ATTN_OUT_A, "blk.%d.attn_output_a" }, + { LLM_TENSOR_ATTN_OUT_B, "blk.%d.attn_output_b" }, + { LLM_TENSOR_ATTN_COMP_KV, "blk.%d.attn_compressor_kv" }, + { LLM_TENSOR_ATTN_COMP_GATE, "blk.%d.attn_compressor_gate" }, + { LLM_TENSOR_ATTN_COMP_APE, "blk.%d.attn_compressor_ape" }, + { LLM_TENSOR_ATTN_COMP_NORM, "blk.%d.attn_compressor_norm" }, + { LLM_TENSOR_ATTN_OUT, "blk.%d.attn_output" }, + { LLM_TENSOR_FFN_NORM, "blk.%d.ffn_norm" }, + { LLM_TENSOR_FFN_GATE, "blk.%d.ffn_gate" }, + { LLM_TENSOR_FFN_UP, "blk.%d.ffn_up" }, + { LLM_TENSOR_FFN_DOWN, "blk.%d.ffn_down" }, + { LLM_TENSOR_FFN_GATE_INP, "blk.%d.ffn_gate_inp" }, + { LLM_TENSOR_FFN_GATE_EXPS, "blk.%d.ffn_gate_exps" }, + { LLM_TENSOR_FFN_DOWN_EXPS, "blk.%d.ffn_down_exps" }, + { LLM_TENSOR_FFN_UP_EXPS, "blk.%d.ffn_up_exps" }, + { LLM_TENSOR_FFN_GATE_UP_EXPS, "blk.%d.ffn_gate_up_exps" }, + { LLM_TENSOR_FFN_GATE_INP_SHEXP, "blk.%d.ffn_gate_inp_shexp" }, + { LLM_TENSOR_FFN_GATE_SHEXP, "blk.%d.ffn_gate_shexp" }, + { LLM_TENSOR_FFN_DOWN_SHEXP, "blk.%d.ffn_down_shexp" }, + { LLM_TENSOR_FFN_UP_SHEXP, "blk.%d.ffn_up_shexp" }, + { LLM_TENSOR_FFN_EXP_PROBS_B, "blk.%d.exp_probs_b" }, + { LLM_TENSOR_FFN_GATE_TID2EID, "blk.%d.ffn_gate_tid2eid" }, + { LLM_TENSOR_INDEXER_K_NORM, "blk.%d.indexer.k_norm" }, + { LLM_TENSOR_INDEXER_PROJ, "blk.%d.indexer.proj" }, + { LLM_TENSOR_INDEXER_ATTN_K, "blk.%d.indexer.attn_k" }, + { LLM_TENSOR_INDEXER_ATTN_Q_B, "blk.%d.indexer.attn_q_b" }, + { LLM_TENSOR_INDEXER_COMP_KV, "blk.%d.indexer_compressor_kv" }, + { LLM_TENSOR_INDEXER_COMP_GATE, "blk.%d.indexer_compressor_gate" }, + { LLM_TENSOR_INDEXER_COMP_APE, "blk.%d.indexer_compressor_ape" }, + { LLM_TENSOR_INDEXER_COMP_NORM, "blk.%d.indexer_compressor_norm" }, + { LLM_TENSOR_HC_HEAD_BASE, "output_hc_base" }, + { LLM_TENSOR_HC_HEAD_FN, "output_hc_fn" }, + { LLM_TENSOR_HC_HEAD_SCALE, "output_hc_scale" }, + { LLM_TENSOR_HC_ATTN_BASE, "blk.%d.hc_attn_base" }, + { LLM_TENSOR_HC_ATTN_FN, "blk.%d.hc_attn_fn" }, + { LLM_TENSOR_HC_ATTN_SCALE, "blk.%d.hc_attn_scale" }, + { LLM_TENSOR_HC_FFN_BASE, "blk.%d.hc_ffn_base" }, + { LLM_TENSOR_HC_FFN_FN, "blk.%d.hc_ffn_fn" }, + { LLM_TENSOR_HC_FFN_SCALE, "blk.%d.hc_ffn_scale" }, + { LLM_TENSOR_NEXTN_EH_PROJ, "blk.%d.nextn.eh_proj" }, + { LLM_TENSOR_NEXTN_EMBED_TOKENS, "blk.%d.nextn.embed_tokens" }, + { LLM_TENSOR_NEXTN_ENORM, "blk.%d.nextn.enorm" }, + { LLM_TENSOR_NEXTN_HNORM, "blk.%d.nextn.hnorm" }, + { LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, "blk.%d.nextn.shared_head_head" }, + { LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, "blk.%d.nextn.shared_head_norm" }, + }, + }, { LLM_ARCH_MISTRAL4, { @@ -2145,15 +2206,13 @@ bool llama_model_is_split_mode_graph(const struct llama_model * model) { } bool llama_model_supports_ctx_shift(const struct llama_model * model) { - // openPangu's latent K rows carry baked-in rope (k_pe) and the DSA indexer cache is - // keyed by absolute position; neither survives K-shift/defrag-style repositioning. - return model && model->arch != LLM_ARCH_OPENPANGU; + // openPangu and DeepSeek4 keep position-dependent private state outside the generic KV cache. + return model && model->arch != LLM_ARCH_OPENPANGU && model->arch != LLM_ARCH_DEEPSEEK4; } bool llama_model_supports_partial_kv_reuse(const struct llama_model * model) { - // openPangu keeps only the current MoME conv state, so a sequence can be extended or - // reset, but rewinding into its decoded middle cannot reconstruct the state at that point. - return model && model->arch != LLM_ARCH_OPENPANGU; + // These architectures cannot reconstruct their private per-position state after a mid-sequence rewind. + return model && model->arch != LLM_ARCH_OPENPANGU && model->arch != LLM_ARCH_DEEPSEEK4; } llm_tensor llm_tensor_type(llm_arch arch, const std::string & tensor_name, int il) { @@ -2219,6 +2278,31 @@ size_t llama_model::cache_size(int il, ggml_type type_k, ggml_type type_v, ggml_ } return size; } + if (arch == LLM_ARCH_DEEPSEEK4) { + constexpr uint32_t csa_ratio = 4; + constexpr uint32_t hca_ratio = 128; + constexpr uint32_t cache_pad = 256; + + const uint32_t n_stream = std::max(1, n_seq_max); + const uint32_t csa_kv = GGML_PAD(std::max(1, (kv_size + csa_ratio - 1)/csa_ratio), cache_pad); + const uint32_t hca_kv = GGML_PAD(std::max(1, (kv_size + hca_ratio - 1)/hca_ratio), cache_pad); + const uint32_t ratio = hparams.dsv4_compress_ratios[(size_t) il]; + const int64_t n_embd_head = hparams.n_embd_head_k(il); + const int64_t n_indexer_head = hparams.indexer_head_size; + + size_t size = ggml_row_size(type_k, n_embd_head) * hparams.n_head_kv(il) * kv_size; + if (ratio == csa_ratio) { + size += ggml_row_size(type_k, n_embd_head) * csa_kv * n_stream; + size += ggml_row_size(idx_type_k, n_indexer_head) * csa_kv * n_stream; + size += (size_t) 2 * n_embd_head * 2 * csa_ratio * n_stream * sizeof(float) * 2; + size += (size_t) 2 * n_indexer_head * 2 * csa_ratio * n_stream * sizeof(float) * 2; + } else if (ratio == hca_ratio) { + size += ggml_row_size(type_k, n_embd_head) * hca_kv * n_stream; + size += (size_t) n_embd_head * hca_ratio * n_stream * sizeof(float) * 2; + } + return size; + } + auto n_head_kv = hparams.n_head_kv(il); auto k_size = ggml_row_size(type_k, hparams.n_embd_head_k(il)) * n_head_kv*kv_size; auto v_size = ggml_row_size(type_v, hparams.n_embd_v_gqa(il)) * kv_size; diff --git a/src/llama-model.h b/src/llama-model.h index b940d66e..8de5b220 100644 --- a/src/llama-model.h +++ b/src/llama-model.h @@ -185,6 +185,9 @@ struct llama_layer { // as "attn_kv_b.weight". Materialized under -sm graph + mla>1; mla=1 skips. struct ggml_tensor * wk_b_pp = nullptr; struct ggml_tensor * wv_b = nullptr; + struct ggml_tensor * wkv_latent = nullptr; + struct ggml_tensor * wo_a = nullptr; + struct ggml_tensor * wo_b = nullptr; struct ggml_tensor * wq_cross = nullptr; struct ggml_tensor * wk_cross = nullptr; struct ggml_tensor * wv_cross = nullptr; @@ -226,7 +229,7 @@ struct llama_layer { llama_split_tensor split_sinks; llama_split_tensor split_wqkv_gate; - // MLA per-device shards (-sm graph for DEEPSEEK2/GLM_DSA/MISTRAL4). + // MLA per-device shards (-sm graph for DEEPSEEK2/DEEPSEEK4/GLM_DSA/MISTRAL4). llama_split_tensor split_wq_a; llama_split_tensor split_wq_b; llama_split_tensor split_wkv_a_mqa; @@ -329,6 +332,7 @@ struct llama_layer { struct ggml_tensor * ffn_up_b = nullptr; // b3 struct ggml_tensor * ffn_act = nullptr; struct ggml_tensor * ffn_exp_probs_b = nullptr; + struct ggml_tensor * ffn_gate_tid2eid = nullptr; llama_split_tensor split_ffn_gate_b; llama_split_tensor split_ffn_down_b; @@ -366,6 +370,21 @@ struct llama_layer { struct ggml_tensor * indexer_proj = nullptr; struct ggml_tensor * indexer_attn_k = nullptr; struct ggml_tensor * indexer_attn_q_b = nullptr; // note: for lora a/b, not bias + struct ggml_tensor * indexer_comp_wkv = nullptr; + struct ggml_tensor * indexer_comp_wgate = nullptr; + struct ggml_tensor * indexer_comp_ape = nullptr; + struct ggml_tensor * indexer_comp_norm = nullptr; + struct ggml_tensor * attn_kv_norm = nullptr; + struct ggml_tensor * hc_attn_base = nullptr; + struct ggml_tensor * hc_attn_fn = nullptr; + struct ggml_tensor * hc_attn_scale = nullptr; + struct ggml_tensor * hc_ffn_base = nullptr; + struct ggml_tensor * hc_ffn_fn = nullptr; + struct ggml_tensor * hc_ffn_scale = nullptr; + struct ggml_tensor * attn_comp_wkv = nullptr; + struct ggml_tensor * attn_comp_wgate = nullptr; + struct ggml_tensor * attn_comp_ape = nullptr; + struct ggml_tensor * attn_comp_norm = nullptr; // long rope factors struct ggml_tensor * rope_long = nullptr; @@ -461,6 +480,9 @@ struct llama_model { struct ggml_tensor * output_b; struct ggml_tensor * output_norm_enc; struct ggml_tensor * output_mtp = nullptr; + struct ggml_tensor * hc_head_base = nullptr; + struct ggml_tensor * hc_head_fn = nullptr; + struct ggml_tensor * hc_head_scale = nullptr; // openPangu-2.0: global mHC stream-merge module (non-block) struct ggml_tensor * mhc_merge_phi = nullptr; diff --git a/src/llama.cpp b/src/llama.cpp index 15323fb9..ac0843a7 100644 --- a/src/llama.cpp +++ b/src/llama.cpp @@ -19,6 +19,7 @@ #include "llama-context.h" #include "llama-spec-features.h" #include "llama-dflash.h" +#include "llama-dsv4.h" #include "llama-quantize.h" #include "unicode.h" @@ -27,6 +28,8 @@ #include "ggml-alloc.h" #include "ggml-backend.h" +#include + void llama_set_mtp_target_context(struct llama_context * ctx, struct llama_context * target_ctx); void llama_set_mtp_step_idx(struct llama_context * ctx, int32_t mtp_step_idx); void llama_set_mtp_n_heads(struct llama_context * ctx, int32_t mtp_n_heads); @@ -87,7 +90,6 @@ void llama_set_mtp_n_heads(struct llama_context * ctx, int32_t mtp_n_heads); //#endif #define LU8(x) (const char*)(u8##x) #include -#include #include #include #include @@ -265,6 +267,7 @@ static const std::map LLM_CHAT_TEMPLATES = { { "deepseek", LLM_CHAT_TEMPLATE_DEEPSEEK }, { "deepseek2", LLM_CHAT_TEMPLATE_DEEPSEEK_2 }, { "deepseek3", LLM_CHAT_TEMPLATE_DEEPSEEK_3 }, + { "deepseek4", LLM_CHAT_TEMPLATE_DEEPSEEK_3 }, { "command-r", LLM_CHAT_TEMPLATE_COMMAND_R }, { "llama3", LLM_CHAT_TEMPLATE_LLAMA_3 }, { "chatglm3", LLM_CHAT_TEMPLATE_CHATGLM_3 }, @@ -409,6 +412,7 @@ static const char * llama_expert_gating_func_name(llm_expert_gating_func_type ty case LLM_EXPERT_GATING_FUNC_SOFTMAX: return "softmax"; case LLM_EXPERT_GATING_FUNC_SIGMOID: return "sigmoid"; case LLM_EXPERT_GATING_FUNC_TYPE_SOFTMAX_WEIGHT: return "softmax_weight"; + case LLM_EXPERT_GATING_FUNC_TYPE_SQRT_SOFTPLUS: return "sqrtsoftplus"; default: return "unknown"; } } @@ -590,6 +594,7 @@ static void why_not_reuse_previous(const llama_batch & u_batch, const llama_cont bool llama_context::can_reuse_graph(const llama_batch & u_batch) { if (!cparams.graph_reuse) return false; + if (model.arch == LLM_ARCH_DEEPSEEK4) return false; auto the_prev = cparams.mtp_op_type == MTP_OP_NONE ? prev.get() : prev_mtp.get(); if (!the_prev || !the_prev->graph) return false; if (u_batch.embd) return false; @@ -770,6 +775,7 @@ llama_context::~llama_context() { ggml_backend_sched_free(dflash.kv.cache_sched); } free_dflash_kv_cache_tensors(); + free_dsv4_cache_tensors(); ggml_backend_sched_free(sched); for (ggml_backend_t backend : backends) { @@ -1086,11 +1092,12 @@ static bool llama_kv_cache_init( } } + const bool is_dsv4_k_only = model.arch == LLM_ARCH_DEEPSEEK4; bool is_mla_attn = model.is_mla_model(); bool split_cache = false; bool replicate_mla = false; - if ((model.split_mode == LLAMA_SPLIT_MODE_GRAPH || model.split_mode == LLAMA_SPLIT_MODE_ATTN) && !is_mla_attn && offload) { + if ((model.split_mode == LLAMA_SPLIT_MODE_GRAPH || model.split_mode == LLAMA_SPLIT_MODE_ATTN) && !is_mla_attn && offload && !is_dsv4_k_only) { cache.split_k_l.reserve(n_layer); cache.split_v_l.reserve(n_layer); if (llama_model_has_recurrent(&model)) { @@ -1174,7 +1181,7 @@ static bool llama_kv_cache_init( if (is_mla_attn && cparams.mla_attn) { needs_v_cache = cparams.mla_attn == 1 && !cparams.flash_attn; } - if (needs_v_cache) cache.v_l.reserve(n_layer); + if (needs_v_cache && !is_dsv4_k_only) cache.v_l.reserve(n_layer); cache.s_l.resize(n_layer, nullptr); // DSA indexer-key cache: one [indexer_head_size, kv_size] tensor per indexer layer. @@ -1201,7 +1208,7 @@ static bool llama_kv_cache_init( // For MTP-only context, skip KV allocation for non-MTP layers if (cparams.mtp_op_type != MTP_OP_NONE && i < n_mtp_first_layer) { cache.k_l.push_back(nullptr); - if (model.arch != LLM_ARCH_OPENPANGU && + if (!is_dsv4_k_only && model.arch != LLM_ARCH_OPENPANGU && (!is_mla_attn || !cparams.mla_attn || (cparams.mla_attn == 1 && !cparams.flash_attn))) { cache.v_l.push_back(nullptr); } @@ -1272,7 +1279,9 @@ static bool llama_kv_cache_init( const bool is_mtp_layer = (cparams.mtp_op_type != MTP_OP_NONE && i >= (int)n_mtp_first_layer); if (!hparams.has_kv(i) && !is_mtp_layer) { cache.k_l.push_back(nullptr); - cache.v_l.push_back(nullptr); + if (!is_dsv4_k_only) { + cache.v_l.push_back(nullptr); + } continue; } if (qnext_recurrent) { @@ -1347,6 +1356,8 @@ static bool llama_kv_cache_init( // The value-side latent is rederived from k_l per graph; no persistent V store. const int64_t n_lat = (int64_t) hparams.n_lora_kv + hparams.n_rot; // 576 k = ggml_new_tensor_2d(ctx, this_type_k, n_lat, kv_size); + } else if (is_dsv4_k_only) { + k = ggml_new_tensor_2d(ctx, this_type_k, n_embd_head_k, n_head_kv*kv_size); } else { k = ggml_new_tensor_2d(ctx, this_type_k, n_embd_head_k, n_head_kv*kv_size); v = ggml_new_tensor_1d(ctx, this_type_v, v_ne); @@ -1422,7 +1433,7 @@ static bool llama_kv_cache_init( v->extra = (void *)&split_v_l.ggml; } cache.k_l.push_back(k); - if (model.arch != LLM_ARCH_OPENPANGU) { + if (!is_dsv4_k_only && model.arch != LLM_ARCH_OPENPANGU) { cache.v_l.push_back(v); } } @@ -6036,6 +6047,10 @@ static int llama_decode_internal( #if IK_PRINT_TIMING tim1 = ggml_time_us(); #endif + if (lctx.model.arch == LLM_ARCH_DEEPSEEK4 && !llama_prepare_dsv4_graph_inputs(lctx, u_batch, false, false)) { + return GGML_STATUS_FAILED; + } + gf = llm_build_context::llama_build_graph(lctx, u_batch, false); #if IK_PRINT_TIMING tim2 = ggml_time_us(); @@ -6074,6 +6089,10 @@ static int llama_decode_internal( return GGML_STATUS_FAILED; } + if (lctx.model.arch == LLM_ARCH_DEEPSEEK4 && !llama_prepare_dsv4_graph_inputs(lctx, u_batch, true, false)) { + return GGML_STATUS_FAILED; + } + // the output is always the last tensor in the graph struct ggml_tensor * res = gf->nodes[gf->n_nodes - 1]; struct ggml_tensor * embd = nullptr; @@ -6130,6 +6149,7 @@ static int llama_decode_internal( #if IK_PRINT_TIMING == 1 tim1 = ggml_time_us(); #endif + //fprintf(stderr, "%s: setting inputs\n", __func__); llama_set_inputs(lctx, u_batch); #if IK_PRINT_TIMING == 1 tim2 = ggml_time_us(); @@ -6138,7 +6158,9 @@ static int llama_decode_internal( #if IK_PRINT_TIMING tim1 = ggml_time_us(); #endif + //fprintf(stderr, "%s: invoking llama_graph_compute\n", __func__); llama_graph_compute(lctx, gf, n_threads); + #if IK_PRINT_TIMING llama_synchronize(&lctx); tim2 = ggml_time_us(); @@ -6762,6 +6784,7 @@ static void llama_kv_cache_defrag_internal(struct llama_context & lctx) { static bool get_can_shift(struct llama_context & lctx) { bool no_shift = lctx.model.is_mla_model(); + no_shift = no_shift || lctx.model.arch == LLM_ARCH_DEEPSEEK4; no_shift = no_shift || lctx.model.hparams.rope_type == LLAMA_ROPE_TYPE_IMROPE; return !no_shift; } @@ -6772,6 +6795,10 @@ static int32_t llama_kv_cache_update_internal(struct llama_context & lctx) { // apply K-shift if needed if (lctx.model.hparams.rope_type != LLAMA_ROPE_TYPE_NONE && lctx.kv_self.has_shift) { if (!get_can_shift(lctx)) { + if (lctx.model.arch == LLM_ARCH_DEEPSEEK4) { + LLAMA_LOG_WARN("%s: DeepSeek4 does not support context shifting; use --no-context-shift or increase context size\n", + __func__); + } return 1; } @@ -6842,7 +6869,11 @@ static int32_t llama_kv_cache_update_internal(struct llama_context & lctx) { int n_tokens = (int)std::min(lctx.cparams.n_ctx, lctx.cparams.n_ubatch); int n_past = lctx.cparams.n_ctx - n_tokens; llama_token token = llama_token_bos(&lctx.model); // not actually used by llama_build_graph, but required to choose between token and embedding inputs graph - ggml_cgraph * gf = llm_build_context::llama_build_graph(lctx, llama_batch_get_one(&token, n_tokens, n_past, 0), true, lctx.cparams.worst_graph_tokens); + llama_batch reserve_batch = llama_batch_get_one(&token, n_tokens, n_past, 0); + if (lctx.model.arch == LLM_ARCH_DEEPSEEK4 && !llama_prepare_dsv4_graph_inputs(lctx, reserve_batch, false, true)) { + return GGML_STATUS_FAILED; + } + ggml_cgraph * gf = llm_build_context::llama_build_graph(lctx, reserve_batch, true, lctx.cparams.worst_graph_tokens); // initialize scheduler with the worst-case graph lctx.reset_scheduler(); @@ -7174,7 +7205,7 @@ struct llama_context_params llama_context_default_params() { /*.rope_cache =*/ false, /*.graph_reuse =*/ true, /*.dsa =*/ false, - /*.fused_idx_topk =*/ false, + /*.fused_idx_topk =*/ true, /*.dsa_top_k =*/ -1, /*.min_experts =*/ -1, /*.thtesh_experts =*/ 0.0f, @@ -7555,12 +7586,36 @@ struct llama_context * llama_init_from_model( // params.flash_attn = false; //} + if (model->arch == LLM_ARCH_DEEPSEEK4 && params.type_v != GGML_TYPE_F16) { + LLAMA_LOG_WARN("%s: DeepSeek4 has no independent V-cache; ignoring requested V-cache type %s\n", + __func__, ggml_type_name(params.type_v)); + params.type_v = GGML_TYPE_F16; + } + if (model->arch != LLM_ARCH_OPENPANGU && params.type_v != GGML_TYPE_F16 && params.type_v != GGML_TYPE_BF16 && !params.flash_attn) { LLAMA_LOG_ERROR("%s: V cache quantization requires flash_attn\n", __func__); return nullptr; } + if (model->arch == LLM_ARCH_DEEPSEEK4 && params.k_cache_hadamard) { + LLAMA_LOG_ERROR("%s: DeepSeek4 K-cache Hadamard is not supported; use an untransformed K-cache\n", + __func__); + return nullptr; + } + + if (model->arch == LLM_ARCH_DEEPSEEK4 && + params.type_k != GGML_TYPE_F16 && params.type_k != GGML_TYPE_BF16 && params.type_k != GGML_TYPE_Q8_0) { + LLAMA_LOG_ERROR("%s: DeepSeek4 K-cache supports only F16, BF16, and Q8_0 (requested %s)\n", + __func__, ggml_type_name(params.type_k)); + return nullptr; + } + + if (model->arch == LLM_ARCH_DEEPSEEK4 && params.v_cache_hadamard) { + LLAMA_LOG_WARN("%s: DeepSeek4 has no independent V-cache; ignoring -vhad\n", __func__); + params.v_cache_hadamard = false; + } + if (params.k_cache_hadamard && !ggml_is_quantized(params.type_k)) { LLAMA_LOG_WARN("%s: there is no point in Hadamard transforms with not quantized K-cache. Turning K-cache Hadamard off\n", __func__); params.k_cache_hadamard = false; @@ -7658,6 +7713,7 @@ struct llama_context * llama_init_from_model( cparams.reduce_type = params.type_reduce; cparams.graph_attn_precision = params.type_graph_attn; + cparams.idx_type_k = params.idx_type_k; if (cparams.graph_attn_precision != GGML_TYPE_F16 && cparams.graph_attn_precision != GGML_TYPE_F32) { throw std::runtime_error(format("--graph-attn-precision must be f16 or f32, got %s", ggml_type_name(cparams.graph_attn_precision))); @@ -7757,6 +7813,7 @@ struct llama_context * llama_init_from_model( if (model->arch != LLM_ARCH_GLM4_MOE && model->arch != LLM_ARCH_QWEN35 && model->arch != LLM_ARCH_QWEN35MOE && model->arch != LLM_ARCH_GEMMA4 && model->arch != LLM_ARCH_GEMMA4_MTP && model->arch != LLM_ARCH_GLM_DSA && + model->arch != LLM_ARCH_DEEPSEEK4 && model->arch != LLM_ARCH_GEMMA4_ASSISTANT && model->arch != LLM_ARCH_OPENPANGU && cparams.mtp != 0) { @@ -8024,6 +8081,10 @@ struct llama_context * llama_init_from_model( LLAMA_LOG_INFO("%s: KV self size = %7.2f MiB, c^KV (%s): %7.2f MiB, kv^T: not used\n", __func__, (float)(memory_size_k + memory_size_v) / (1024.0f * 1024.0f), ggml_type_name(type_k), (float)memory_size_k / (1024.0f * 1024.0f)); + } else if (model->arch == LLM_ARCH_DEEPSEEK4) { + LLAMA_LOG_INFO("%s: KV self size = %7.2f MiB, K-only (%s): %7.2f MiB; independent V-cache: not used\n", __func__, + (float) memory_size_k / (1024.0f * 1024.0f), + ggml_type_name(type_k), (float) memory_size_k / (1024.0f * 1024.0f)); } else { LLAMA_LOG_INFO("%s: KV self size = %7.2f MiB, K (%s): %7.2f MiB, V (%s): %7.2f MiB\n", __func__, (float)(memory_size_k + memory_size_v) / (1024.0f * 1024.0f), @@ -8092,7 +8153,12 @@ struct llama_context * llama_init_from_model( // build worst-case graph int n_past = cparams.n_ctx - n_tokens; llama_token token = llama_token_bos(&ctx->model); // not actually used by llama_build_graph, but required to choose between token and embedding inputs graph - ggml_cgraph * gf = llm_build_context::llama_build_graph(*ctx, llama_batch_get_one(&token, n_tokens, n_past, 0), true, cparams.worst_graph_tokens); + llama_batch reserve_batch = llama_batch_get_one(&token, n_tokens, n_past, 0); + if (ctx->model.arch == LLM_ARCH_DEEPSEEK4 && !llama_prepare_dsv4_graph_inputs(*ctx, reserve_batch, false, true)) { + llama_free(ctx); + return nullptr; + } + ggml_cgraph * gf = llm_build_context::llama_build_graph(*ctx, reserve_batch, true, cparams.worst_graph_tokens); // initialize scheduler with the worst-case graph bool gf_success = ggml_backend_sched_reserve(ctx->sched, gf); @@ -8266,6 +8332,7 @@ enum llama_rope_type llama_rope_type(const struct llama_model * model) { case LLM_ARCH_OLMO: case LLM_ARCH_ARCTIC: case LLM_ARCH_DEEPSEEK2: + case LLM_ARCH_DEEPSEEK4: case LLM_ARCH_CHATGLM: case LLM_ARCH_GLM4: case LLM_ARCH_GRANITE: @@ -8688,6 +8755,7 @@ int32_t llama_get_kv_cache_used_cells(const struct llama_context * ctx) { void llama_kv_cache_clear(struct llama_context * ctx) { llama_kv_cache_clear(ctx->kv_self); + llama_reset_dsv4_state(ctx); } // Unified speculative-checkpoint @@ -8968,7 +9036,11 @@ void llama_spec_ckpt_discard(struct llama_context * ctx) { } bool llama_kv_cache_seq_rm(struct llama_context * ctx, llama_seq_id seq_id, llama_pos p0, llama_pos p1) { - return llama_kv_cache_seq_rm(ctx->kv_self, seq_id, p0, p1); + const bool result = llama_kv_cache_seq_rm(ctx->kv_self, seq_id, p0, p1); + if (result && ctx->model.arch == LLM_ARCH_DEEPSEEK4 && p0 <= 0 && p1 < 0) { + llama_reset_dsv4_state(ctx, seq_id); + } + return result; } void llama_kv_cache_seq_cp(struct llama_context * ctx, llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) { @@ -10154,13 +10226,11 @@ struct llama_data_read_file : llama_data_read { } }; -// openPangu: refuse state I/O outright until the side state is part of the format. K/V -// rows alone are not a complete snapshot; the MoME conv slot (s_l) and the DSA -// indexer cache are needed to resume a sequence, and restoring without them diverges -// silently instead of failing. +// Refuse state I/O when private per-position state is not part of the format. static bool llama_state_io_supported(const struct llama_context * ctx, const char * func) { - if (ctx->model.arch == LLM_ARCH_OPENPANGU) { - LLAMA_LOG_ERROR("%s: state save/restore is not supported for openPangu (conv slot and indexer cache are not serialized)\n", func); + if (ctx->model.arch == LLM_ARCH_OPENPANGU || ctx->model.arch == LLM_ARCH_DEEPSEEK4) { + const char * arch = ctx->model.arch == LLM_ARCH_OPENPANGU ? "openPangu" : "DeepSeek4"; + LLAMA_LOG_ERROR("%s: state save/restore is not supported for %s (private cache and side state are not serialized)\n", func, arch); return false; } return true; @@ -10322,7 +10392,9 @@ static bool llama_state_save_file_internal(struct llama_context * ctx, const cha // save the context state using stream saving llama_data_write_file data_ctx(&file, ctx->model); - llama_state_get_data_internal(ctx, data_ctx); + if (llama_state_get_data_internal(ctx, data_ctx) == 0) { + return false; + } return true; } @@ -10398,7 +10470,9 @@ static size_t llama_state_seq_save_file_internal(struct llama_context * ctx, con // save the context state using stream saving llama_data_write_file data_ctx(&file, ctx->model); - llama_state_seq_get_data_internal(ctx, data_ctx, seq_id, 0); + if (llama_state_seq_get_data_internal(ctx, data_ctx, seq_id, 0) == 0) { + return 0; + } const size_t res = file.tell(); GGML_ASSERT(res == sizeof(uint32_t) * 3 + sizeof(llama_token) * n_token_count + data_ctx.get_size_written());