DS4: slowly approaching a meaningful performance (#2165)

* initial map to load deepseek 4 arch

* wip

* wip: match graph build and attn logic for dpv4

* wip: Enhance DeepSeek-V4 architecture with new tensor types and sqrtsoftplus gating function

* Update DeepSeek-V4 to support raw key indexing with read/write indices

* fix mismatch in attn_raw

* Enable FA with CSA/HCA

* Fix logit mismatch with FA path

* Clean traces and logs for debug

* Refactor DSV4 tensor handling for MTP execution and improve raw context management

* Refactor DeepSeek4 tensor operations: replace manual weighted sum and post-processing with new helper functions

* Share mHC pre-projection and fix packed DSV4 writes

* DSV4: add shared top-k selection and improve mask handling

* Fix DSV4 c2048 view stride and duplicate loader instantiation

* Reuse shared RMS normalization in DSV4 graph

* Replace DSV4 indexer rotation with shared Hadamard

* Share CSA visibility mask with DSV4 LID

* dsv4: document dependency ordering and reset state

* Remove DSV4 zero-dependency graph shim

* Fix DSV4 packed stream execution

* Remove DSV4 l_out backend override

* Enable DSV4 quantized K-only cache

* Revert "Enable DSV4 quantized K-only cache"

This reverts commit 04f9b425321f62ba60e16d1bea2f8de714cfe855.

* Fix DSV4 quantized cache accounting

* Fail closed on unsupported DSV4 cache lifecycle operations

* Various optimizations

* llama: fix GGML_METAL=ON build - missing ggml-metal.h include in llama-dflash.cpp (#2134)

llama-dflash.cpp calls ggml_backend_is_metal() and
ggml_backend_metal_set_n_cb() inside an #ifdef GGML_USE_METAL block but
never includes ggml-metal.h, so any Metal-enabled build fails to
compile. Add the same guarded include llama.cpp already uses.

* New op: ggml_sum_rows_ext (#2132)

* Add ggml_sum_rows_ext

* openPangu: use ggml_sum_rows_ext also in mhc_post

* openPangu: use ggml_sum_rows_ext also in mhc_tail

* Minor

* Reuse shared inverse RoPE operation for DSV4

* Reuse maintainer CUDA concat implementation

* WIP

* hc_pre

* hc_post

* Remove unnecessary mask manipulations

* WIP

* Take into account swiglu limits

* Turn on fused indexer by default

* Give names to mat mul results

* More named ops

* dsv4: do not uselessly copy the KV cache

+20% TG at 32k tokens

* mask_to_index and make CPU FA work with that

* Much better CPU-only, CUDA still not functional

* Better CPU TG

I'm now at 9.7 t/s for zero context and 6.5 t/s for context of 32k.
PP is 120 t/s for short context and 101 t/s at 32k.

* Even better CPU TG

I'm now at 8.1 t/s for context of 32k tokens.

* Turn off DSA on CUDA for now

* Fix CUDA DSA

* Remove again the unnecessary softmax result buffer

* Experiments

* Various

* More named ops

* Forgot to uncomment

---------

Co-authored-by: samuel <samueloliveira32df@gmail.com>
Co-authored-by: hchengit <95317477+hchengit@users.noreply.github.com>
This commit is contained in:
Kawrakow 2026-07-22 17:18:57 +03:00 committed by GitHub
parent 9d07d8681e
commit 7945404458
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
52 changed files with 5175 additions and 303 deletions

View File

@ -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;

View File

@ -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)]

View File

@ -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<server_prompt_cache>(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;

View File

@ -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();

View File

@ -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 *);

View File

@ -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];

View File

@ -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 <int ncols_template, int block_size_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><<<block_nums, block_dims, shmem, stream>>>(x, mask, ncols_x, nrows_y, scale);
soft_max_f16_simple<32, 32><<<block_nums, block_dims, shmem, stream>>>(x, mask, sinks, ncols_x, nrows_y, scale);
break;
case 64:
soft_max_f16_simple<64, 64><<<block_nums, block_dims, shmem, stream>>>(x, mask, ncols_x, nrows_y, scale);
soft_max_f16_simple<64, 64><<<block_nums, block_dims, shmem, stream>>>(x, mask, sinks, ncols_x, nrows_y, scale);
break;
case 128:
soft_max_f16_simple<128, 128><<<block_nums, block_dims, shmem, stream>>>(x, mask, ncols_x, nrows_y, scale);
soft_max_f16_simple<128, 128><<<block_nums, block_dims, shmem, stream>>>(x, mask, sinks, ncols_x, nrows_y, scale);
break;
case 256:
soft_max_f16_simple<256, 256><<<block_nums, block_dims, shmem, stream>>>(x, mask, ncols_x, nrows_y, scale);
soft_max_f16_simple<256, 256><<<block_nums, block_dims, shmem, stream>>>(x, mask, sinks, ncols_x, nrows_y, scale);
break;
case 512:
soft_max_f16_simple<512, 512><<<block_nums, block_dims, shmem, stream>>>(x, mask, ncols_x, nrows_y, scale);
soft_max_f16_simple<512, 512><<<block_nums, block_dims, shmem, stream>>>(x, mask, sinks, ncols_x, nrows_y, scale);
break;
case 1024:
soft_max_f16_simple<1024, 1024><<<block_nums, block_dims, shmem, stream>>>(x, mask, ncols_x, nrows_y, scale);
soft_max_f16_simple<1024, 1024><<<block_nums, block_dims, shmem, stream>>>(x, mask, sinks, ncols_x, nrows_y, scale);
break;
case 2048:
soft_max_f16_simple<2048, 1024><<<block_nums, block_dims, shmem, stream>>>(x, mask, ncols_x, nrows_y, scale);
soft_max_f16_simple<2048, 1024><<<block_nums, block_dims, shmem, stream>>>(x, mask, sinks, ncols_x, nrows_y, scale);
break;
case 4096:
soft_max_f16_simple<4096, 1024><<<block_nums, block_dims, shmem, stream>>>(x, mask, ncols_x, nrows_y, scale);
soft_max_f16_simple<4096, 1024><<<block_nums, block_dims, shmem, stream>>>(x, mask, sinks, ncols_x, nrows_y, scale);
break;
default:
soft_max_f16_simple<0, 0><<<block_nums, block_dims, shmem, stream>>>(x, mask, ncols_x, nrows_y, scale);
soft_max_f16_simple<0, 0><<<block_nums, block_dims, shmem, stream>>>(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());

View File

@ -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:

View File

@ -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 <typename mask_t>
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<mask_t, half>) {
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<<<grid, WARP_SIZE, 0, ctx.stream()>>>(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<<<grid, WARP_SIZE, 0, ctx.stream()>>>(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);
}
}

View File

@ -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);

View File

@ -62,6 +62,113 @@ static __global__ void k_sinkhorn(const float * __restrict__ x, float * __restri
}
}
template <int S>
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 <int S>
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><<<grid, block, 0, stream>>>(x, s, b, y, T, iters, eps, src0->nb[1]); break;
case 2: k_hc_pre<2><<<grid, block, 0, stream>>>(x, s, b, y, T, iters, eps, src0->nb[1]); break;
case 3: k_hc_pre<3><<<grid, block, 0, stream>>>(x, s, b, y, T, iters, eps, src0->nb[1]); break;
case 4: k_hc_pre<4><<<grid, block, 0, stream>>>(x, s, b, y, T, iters, eps, src0->nb[1]); break;
case 5: k_hc_pre<5><<<grid, block, 0, stream>>>(x, s, b, y, T, iters, eps, src0->nb[1]); break;
case 6: k_hc_pre<6><<<grid, block, 0, stream>>>(x, s, b, y, T, iters, eps, src0->nb[1]); break;
case 7: k_hc_pre<7><<<grid, block, 0, stream>>>(x, s, b, y, T, iters, eps, src0->nb[1]); break;
case 8: k_hc_pre<8><<<grid, block, 0, stream>>>(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><<<nblock, kBlockSize, 0, ctx.stream()>>>(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><<<nblock, kBlockSize, 0, ctx.stream()>>>(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><<<nblock, kBlockSize, 0, ctx.stream()>>>(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><<<nblock, kBlockSize, 0, ctx.stream()>>>(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><<<nblock, kBlockSize, 0, ctx.stream()>>>(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><<<nblock, kBlockSize, 0, ctx.stream()>>>(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><<<nblock, kBlockSize, 0, ctx.stream()>>>(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><<<nblock, kBlockSize, 0, ctx.stream()>>>(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");
}
}

View File

@ -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);

View File

@ -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<op_softplus>(ctx, dst);
}
void ggml_cuda_op_sqrt_softplus(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
ggml_cuda_op_unary<op_sqrt_softplus>(ctx, dst);
}
// === gated ops
template <float (*op)(float), typename T>

View File

@ -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);

View File

@ -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;

View File

@ -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);
}

View File

@ -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);
}

View File

@ -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);
}

View File

@ -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);
}

View File

@ -10,43 +10,44 @@ template <int step_k, typename KHelper, typename VHelper>
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 <int step_k>
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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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);
}

View File

@ -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);
}

View File

@ -10,43 +10,44 @@ template <int step_k, typename KHelper, typename VHelper>
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 <int step_k>
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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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<step_k>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_deepseek_helper<step_k>(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);
}

View File

@ -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);
}

View File

@ -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);
}

View File

@ -1156,11 +1156,11 @@ struct FlashQKV {
}
template <typename FMS>
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 <typename FMS>
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 <typename FMS>
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<q_step, k_step>& fms,
FlashQKV<Dv, q_step, k_step>& 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<q_step, k_step>& fms,
FlashQKV<Dv, q_step, k_step>& 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 <typename KHelper, typename VHelper>
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<Dk> khr4(nk1, kh);
#endif
compute_helper_q<Dk, Dv, q_step, k_step, HelperQ80R8<Dk>, VHelper, FlashQKfp32<Dk, q_step, k_step>>(
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<Dk> khr4(nk1, kh);
#endif
compute_helper_q<Dk, Dv, q_step, k_step, HelperQ8KVR8<Dk>, VHelper, FlashQKfp32<Dk, q_step, k_step>>(
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<Dk, Dv, q_step, k_step, KHelper, VHelper, FlashQKfp32<Dk, q_step, k_step>>(
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<Dk, Dv, q_step, k_step, KHelper, VHelper, FlashQKfp32<Dk, q_step, k_step>>(
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<Dk, Dv, q_step, k_step, KHelper, VHelper, FlashQKfp32<Dk, q_step, k_step>>(
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<q_step, k_step> fms;
FlashQKV<Dv, q_step, k_step> 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 <typename KHelper, typename VHelper>
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<q_step, k_step> fms;
FlashQKV<Dv, q_step, k_step> fqkv;
const float * sinkf;
int sink_stride;
};
#endif
template <int Dk, int Dv, int k_step, typename KHelper, typename VHelper>
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<Dk, Dv, 64, k_step> fa(scale, softcap, sinkf);
FlashAttn<Dk, Dv, 64, k_step> 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<Dk, Dv, 64, k_step> fa(scale, softcap, sinkf);
FlashAttn<Dk, Dv, 64, k_step> 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<Dk, Dv, 32, k_step> fa(scale, softcap, sinkf);
FlashAttn<Dk, Dv, 32, k_step> 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<Dk, Dv, 16, k_step> fa(scale, softcap, sinkf);
FlashAttn<Dk, Dv, 16, k_step> 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<Dk, Dv, 12, k_step> fa(scale, softcap, sinkf);
FlashAttn<Dk, Dv, 12, k_step> 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<Dk, Dv, 8, k_step> fa(scale, softcap, sinkf);
FlashAttn<Dk, Dv, 8, k_step> 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<Dk, Dv, 4, k_step> fa(scale, softcap, sinkf);
FlashAttn<Dk, Dv, 4, k_step> 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<Dk, Dv, 2, k_step> fa(scale, softcap, sinkf);
FlashAttn<Dk, Dv, 2, k_step> 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<Dk, Dv, 1, k_step> fa(scale, softcap, sinkf);
FlashAttn<Dk, Dv, 1, k_step> 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 <int Dk, int Dv, int k_step>
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<Dk, k_step> kh(k, stride_k);
HelperBF16<Dv, k_step> vh(v, stride_v);
if (nk1 >= 4096) {
if (nq1 >= 64) {
FlashAttnBF16<Dk, Dv, 64, k_step> fa(scale, softcap, sinkf);
FlashAttnBF16<Dk, Dv, 64, k_step> 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<Dk, Dv, 16, k_step> fa(scale, softcap, sinkf);
FlashAttnBF16<Dk, Dv, 16, k_step> 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<Dk, Dv, 8, k_step> fa(scale, softcap, sinkf);
FlashAttnBF16<Dk, Dv, 8, k_step> 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<Dk, Dv, 1, k_step> fa(scale, softcap, sinkf);
FlashAttnBF16<Dk, Dv, 1, k_step> 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 <int Dk, int Dv, int k_step, typename KHelper>
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<Dk, Dv, k_step>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_flash_helper<Dk, Dv, k_step>(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<Dv, k_step> vh(v, stride_v);
iqk_flash_helper<Dk, Dv, k_step>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_flash_helper<Dk, Dv, k_step>(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<Dk, Dv, k_step>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_flash_helper<Dk, Dv, k_step>(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<Dv> vh(v, stride_v);
iqk_flash_helper<Dk, Dv, k_step>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_flash_helper<Dk, Dv, k_step>(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<Dk, Dv, k_step>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_flash_helper<Dk, Dv, k_step>(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<Dk, Dv, k_step>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_flash_helper<Dk, Dv, k_step>(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<Dk, Dv, k_step>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_flash_helper<Dk, Dv, k_step>(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<Dk, Dv, k_step>(kh, vh, nq1, nk1, stride_q, stride_m, stride_qkv, q, mask, scale, softcap, qkv, sinkf, M, S);
iqk_flash_helper<Dk, Dv, k_step>(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 <int Dk, int Dv, int k_step>
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<Dk, Dv, k_step>(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<Dk, Dv, k_step>(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<Dk, Dv, k_step>(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<Dk, Dv, k_step>(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<Dk> kh(k, stride_k);
result = iqk_flash_helper_T<Dk, Dv, k_step>(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<Dk, Dv, k_step>(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<Dk, Dv, k_step>(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<Dk, Dv, k_step>(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<Dk> kh(k, stride_k);
result = iqk_flash_helper_T<Dk, Dv, k_step>(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<Dk, Dv, k_step>(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<Dk, Dv, k_step>(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<Dk, Dv, k_step>(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<Dk, Dv, k_step>(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<Dk, Dv, k_step>(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<Dk, Dv, k_step>(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<Dk, Dv, k_step>(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);

View File

@ -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);

View File

@ -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;
}

View File

@ -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))

View File

@ -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;

View File

@ -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,
],

View File

@ -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)

View File

@ -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

View File

@ -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
```

View File

@ -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 = '<think>' -%}
{%- set thinking_end_token = '</think>' -%}
{%- 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</' + dsml_token + 'parameter>\n...\n</' + dsml_token + 'invoke>\n<' + dsml_token + 'invoke name="$TOOL_NAME2">\n...\n</' + dsml_token + 'invoke>\n</' + dsml_token + 'tool_calls>\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 -%}
{{- '<tool_result>' + (message['content'] or '') + '</tool_result>' -}}
{%- 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 + '</' + dsml_token + 'parameter>\n' -}}
{%- else -%}
{{- '<' + dsml_token + 'parameter name="' + key + '" string="false">' + (val | tojson) + '</' + dsml_token + 'parameter>\n' -}}
{%- endif -%}
{%- endfor -%}
{{- '</' + dsml_token + 'invoke>\n' -}}
{%- endfor -%}
{{- '</' + dsml_token + 'tool_calls>' -}}
{%- endif -%}
{{- '<end▁of▁sentence>' -}}
{%- endif -%}
{%- endfor -%}
{%- if add_generation_prompt -%}
{{- '<Assistant>' -}}
{%- if thinking -%}
{{- thinking_start_token -}}
{%- else -%}
{{- thinking_end_token -}}
{%- endif -%}
{%- endif -%}

View File

@ -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

File diff suppressed because it is too large Load Diff

View File

@ -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);

View File

@ -53,6 +53,7 @@ static const std::map<llm_arch, const char *> 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, const char *> 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" },

View File

@ -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,

View File

@ -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();

View File

@ -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,

View File

@ -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<llama_seq_id> strm;
std::vector<std::vector<uint32_t>> 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<int32_t> write_src_idxs;
std::vector<int32_t> write_dst_idxs;
std::vector<int32_t> read_dst_idxs;
std::vector<int32_t> write_counts;
std::vector<int32_t> 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<int32_t> state_pos;
std::vector<int32_t> state_persist_src_idxs;
std::vector<int32_t> state_persist_dst_idxs;
std::vector<int32_t> state_read_idxs;
std::vector<int64_t> state_write_idxs;
std::vector<int32_t> state_write_pos;
std::vector<int32_t> 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<struct ggml_tensor *> csa_k;
std::vector<struct ggml_tensor *> hca_k;
std::vector<struct ggml_tensor *> lid_k;
std::vector<struct ggml_tensor *> csa_state_kv;
std::vector<struct ggml_tensor *> csa_state_score;
std::vector<struct ggml_tensor *> hca_state_kv;
std::vector<struct ggml_tensor *> hca_state_score;
std::vector<struct ggml_tensor *> lid_state_kv;
std::vector<struct ggml_tensor *> lid_state_score;
struct ggml_context * cache_ctx = nullptr;
std::vector<ggml_backend_buffer_t> 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<float> csa_mask_data;
std::vector<float> 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);

View File

@ -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;

1095
src/llama-dsv4.cpp Normal file

File diff suppressed because it is too large Load Diff

8
src/llama-dsv4.h Normal file
View File

@ -0,0 +1,8 @@
#pragma once
#include <cstdint>
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);

View File

@ -3,6 +3,7 @@
#include "llama-model-loader.h"
#include "llama-model.h"
#include <algorithm>
#include <limits>
#include <map>
@ -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<uint32_t> 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) {

View File

@ -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<uint32_t, LLAMA_MAX_LAYERS> 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<uint32_t, LLAMA_MAX_LAYERS> dsv4_compress_ratios = {};
// qwen3vl deepstack
uint32_t n_deepstack_layers = 0;

View File

@ -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<int64_t> 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<std::string> 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:

View File

@ -1391,6 +1391,6 @@ template bool llama_model_loader::get_key_or_arr<std::array<float, 512>>(enum ll
template std::enable_if<std::is_integral<unsigned int>::value, bool>::type llama_model_loader::get_arr_n<unsigned int>(const std::string &, unsigned int &, bool);
template std::enable_if<std::is_integral<unsigned int>::value, bool>::type llama_model_loader::get_arr_n<unsigned int>(enum llm_kv, unsigned int&, bool);
template bool llama_model_loader::get_arr<uint32_t>(const std::string &, std::vector<uint32_t> &, bool);
template bool llama_model_loader::get_arr<int32_t, 8>(const std::string &, std::array<int32_t, 8> &, bool);
template bool llama_model_loader::get_arr<uint32_t, 8>(const std::string &, std::array<uint32_t, 8> &, bool);
template bool llama_model_loader::get_arr<uint32_t>(const std::string &, std::vector<uint32_t> &, bool);

View File

@ -1064,6 +1064,67 @@ static const std::map<llm_arch, std::map<llm_tensor, std::string>> 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<uint32_t>(1, n_seq_max);
const uint32_t csa_kv = GGML_PAD(std::max<uint32_t>(1, (kv_size + csa_ratio - 1)/csa_ratio), cache_pad);
const uint32_t hca_kv = GGML_PAD(std::max<uint32_t>(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;

View File

@ -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;

View File

@ -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 <cmath>
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 <algorithm>
#include <array>
#include <cassert>
#include <cctype>
#include <cfloat>
@ -265,6 +267,7 @@ static const std::map<std::string, llm_chat_template> 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());