echodict/llama.cpp
version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786
0479
1#include "models.h"2 3#include "llama-impl.h"4 5// utility to get one slice from the third dimension6// input dim: [x, y, c, b]7// output dim: [x, y, 1, b]8static ggml_tensor * get_slice_2d(ggml_context * ctx0, ggml_tensor * t, int64_t c) {9 return ggml_view_4d(ctx0, t, t->ne[0], t->ne[1], 1, t->ne[3],10 t->nb[1], t->nb[2], t->nb[3], t->nb[2] * c);11}12 13llm_build_delta_net_base::llm_build_delta_net_base(const llm_graph_params & params) : llm_graph_context(params) {}14 15std::pair<ggml_tensor *, ggml_tensor *> llm_build_delta_net_base::build_delta_net_chunking(16 ggml_tensor * q,17 ggml_tensor * k,18 ggml_tensor * v,19 ggml_tensor * g,20 ggml_tensor * b,21 ggml_tensor * s,22 int il) {23 const int64_t S_k = q->ne[0];24 const int64_t H_k = q->ne[1];25 const int64_t n_tokens = q->ne[2];26 const int64_t n_seqs = q->ne[3];27 28 const int64_t S_v = v->ne[0];29 const int64_t H_v = v->ne[1];30 const bool kda = (g->ne[0] == S_k && g->ne[1] == H_k);31 32 GGML_ASSERT(S_k == S_v);33 GGML_ASSERT(H_v % H_k == 0);34 35 GGML_ASSERT(q->ne[0] == S_k && q->ne[1] == H_k && q->ne[2] == n_tokens && q->ne[3] == n_seqs);36 GGML_ASSERT(k->ne[0] == S_k && k->ne[1] == H_k && k->ne[2] == n_tokens && k->ne[3] == n_seqs);37 GGML_ASSERT(v->ne[0] == S_v && v->ne[1] == H_v && v->ne[2] == n_tokens && v->ne[3] == n_seqs);38 39 GGML_ASSERT(g->ne[0] == 1 || g->ne[0] == S_v);40 GGML_ASSERT( g->ne[1] == H_v && g->ne[2] == n_tokens && g->ne[3] == n_seqs);41 GGML_ASSERT(b->ne[0] == 1 && b->ne[1] == H_v && b->ne[2] == n_tokens && b->ne[3] == n_seqs);42 GGML_ASSERT(s->ne[0] == S_v && s->ne[1] == S_v && s->ne[2] == H_v && s->ne[3] == n_seqs);43 44 const float scale = 1.0f / sqrtf(S_k);45 46 q = ggml_scale(ctx0, q, scale);47 48 cb(q, "q_in", il);49 cb(k, "k_in", il);50 cb(v, "v_in", il);51 cb(b, "b_in", il);52 cb(g, "g_in", il);53 54 q = ggml_permute(ctx0, q, 0, 2, 1, 3); // [S_k, n_tokens, H_k, n_seqs]55 k = ggml_permute(ctx0, k, 0, 2, 1, 3); // [S_k, n_tokens, H_k, n_seqs]56 v = ggml_permute(ctx0, v, 0, 2, 1, 3); // [S_v, n_tokens, H_v, n_seqs]57 g = ggml_permute(ctx0, g, 0, 2, 1, 3); // [g_0, n_tokens, H_v, n_seqs]58 b = ggml_permute(ctx0, b, 0, 2, 1, 3); // [ 1, n_tokens, H_v, n_seqs]59 60 const int CS = kda ? 16 : 64; // chunk size61 62 const int pad = (CS - n_tokens % CS) % CS;63 const int n_chunks = (n_tokens + pad) / CS;64 65 q = ggml_pad(ctx0, q, 0, pad, 0, 0);66 k = ggml_pad(ctx0, k, 0, pad, 0, 0);67 v = ggml_pad(ctx0, v, 0, pad, 0, 0);68 g = ggml_pad(ctx0, g, 0, pad, 0, 0);69 b = ggml_pad(ctx0, b, 0, pad, 0, 0);70 71 ggml_tensor * v_b = ggml_mul(ctx0, v, b);72 ggml_tensor * k_b = ggml_mul(ctx0, k, b);73 74 cb(v_b, "v_b", il);75 cb(k_b, "k_b", il);76 77 q = ggml_reshape_4d(ctx0, q, S_k, CS, n_chunks, H_k * n_seqs);78 k = ggml_reshape_4d(ctx0, k, S_k, CS, n_chunks, H_k * n_seqs);79 k_b = ggml_reshape_4d(ctx0, k_b, S_k, CS, n_chunks, H_v * n_seqs);80 v = ggml_reshape_4d(ctx0, v, S_v, CS, n_chunks, H_v * n_seqs);81 v_b = ggml_reshape_4d(ctx0, v_b, S_v, CS, n_chunks, H_v * n_seqs);82 83 g = ggml_reshape_4d(ctx0, g, g->ne[0], CS, n_chunks, H_v * n_seqs);84 b = ggml_reshape_4d(ctx0, b, 1, CS, n_chunks, H_v * n_seqs);85 86 // [CS, g_0, n_chunks, H_v * n_seqs]87 // TODO: extend ggml_cumsum with axis parameter to avoid transpose88 ggml_tensor * g_cs = ggml_cumsum(ctx0, ggml_cont(ctx0, ggml_transpose(ctx0, g)));89 cb(g_cs, "g_cs", il);90 91 ggml_tensor * kb = nullptr;92 ggml_tensor * kq = nullptr;93 if (kda) {94 const int64_t CHB = n_chunks * H_k * n_seqs;95 96 ggml_tensor * g_cs_i = ggml_reshape_4d(ctx0, g_cs, CS, 1, S_k, CHB); // [chunk_size, 1, S_k, CHB]97 ggml_tensor * g_cs_j = ggml_reshape_4d(ctx0, g_cs, 1, CS, S_k, CHB); // [1, chunk_size, S_k, CHB]98 99 g_cs_j = ggml_repeat_4d(ctx0, g_cs_j, CS, CS, S_k, CHB); // [1, chunk_size, S_k, CHB] -> [chunk_size, chunk_size, S_k, CHB]100 101 // decay_mask [chunk_size,chunk_size,S_k,CHB]102 ggml_tensor * decay_mask;103 decay_mask = ggml_sub(ctx0, g_cs_j, g_cs_i);104 decay_mask = ggml_tri(ctx0, decay_mask, GGML_TRI_TYPE_LOWER_DIAG);105 decay_mask = ggml_exp(ctx0, decay_mask);106 cb(decay_mask, "decay_mask", il);107 108 // decay_mask [S_k,BT_j,BT_i,CHB] *Note* second and third chunk_sizes are switched109 decay_mask = ggml_cont_4d(ctx0, ggml_permute(ctx0, decay_mask, 2, 1, 0, 3), S_k, CS, CS, CHB);110 111 ggml_tensor * k_b_i = ggml_reshape_4d(ctx0, k_b, S_k, CS, 1, CHB);112 ggml_tensor * k_j = ggml_reshape_4d(ctx0, k, S_k, 1, CS, CHB);113 ggml_tensor * q_i = ggml_reshape_4d(ctx0, q, S_k, CS, 1, CHB);114 115 ggml_tensor * decay_k_b_i = ggml_mul(ctx0, decay_mask, k_b_i);116 ggml_tensor * decay_q_i = ggml_mul(ctx0, decay_mask, q_i);117 118 // decay_k_b_i [S,BT,BT,CHB] @ k_j [S,1,BT,CHB] = Akk [BT,1,BT,CHB]119 kb = ggml_mul_mat(ctx0, decay_k_b_i, k_j);120 kq = ggml_mul_mat(ctx0, decay_q_i, k_j);121 122 kb = ggml_cont(ctx0, ggml_transpose(ctx0, ggml_reshape_4d(ctx0, kb, CS, CS, n_chunks, H_v * n_seqs)));123 kq = ggml_cont(ctx0, ggml_transpose(ctx0, ggml_reshape_4d(ctx0, kq, CS, CS, n_chunks, H_v * n_seqs)));124 } else {125 ggml_tensor * g_cs_i = g_cs;126 ggml_tensor * g_cs_j = ggml_reshape_4d(ctx0, g_cs, 1, CS, n_chunks, H_v * n_seqs);127 128 g_cs_j = ggml_repeat_4d(ctx0, g_cs_j, CS, CS, n_chunks, H_v * n_seqs);129 130 // [CS, CS, n_chunks, H_v * n_seqs]131 ggml_tensor * decay_mask;132 decay_mask = ggml_sub(ctx0, g_cs_j, g_cs_i);133 decay_mask = ggml_tri(ctx0, decay_mask, GGML_TRI_TYPE_LOWER_DIAG);134 decay_mask = ggml_exp(ctx0, decay_mask);135 cb(decay_mask, "decay_mask", il);136 137 // [CS, CS, n_chunks, H_k * n_seqs]138 kb = ggml_mul_mat(ctx0, k, k_b);139 kb = ggml_mul (ctx0, kb, decay_mask);140 141 // [CS, CS, n_chunks, H_k * n_seqs]142 kq = ggml_mul_mat(ctx0, k, q);143 kq = ggml_mul(ctx0, kq, decay_mask);144 }145 146 kq = ggml_tri(ctx0, kq, GGML_TRI_TYPE_LOWER_DIAG);147 cb(kq, "kq", il);148 149 // [CS, CS, n_chunks, H_k * n_seqs]150 ggml_tensor * attn;151 attn = ggml_tri(ctx0, kb, GGML_TRI_TYPE_LOWER);152 cb(attn, "attn", il);153 154 ggml_tensor * identity;155 identity = ggml_view_1d(ctx0, attn, CS, 0);156 identity = ggml_fill (ctx0, identity, 1.0f);157 identity = ggml_diag (ctx0, identity);158 159 ggml_tensor * lhs = ggml_add(ctx0, attn, identity);160 cb(lhs, "dnet_add_ch_lhs", il);161 162 attn = ggml_neg(ctx0, attn);163 cb(attn, "attn_pre_solve", il);164 165 ggml_tensor * lin_solve = ggml_solve_tri(ctx0, lhs, attn, true, true, false);166 attn = ggml_add(ctx0, lin_solve, identity);167 cb(attn, "dnet_add_ch_attn_solved", il); // [CS, CS, n_chunks, H_k * n_seqs]168 169 // [S_v, CS, n_chunks, H_v * n_seqs]170 v = ggml_mul_mat(ctx0, ggml_cont(ctx0, ggml_transpose(ctx0, v_b)), attn);171 172 // [CS, 1, n_chunks, H_v * n_seqs] KDA: [CS, S_k, n_chunks, H_v * n_seqs]173 ggml_tensor * g_exp = ggml_exp(ctx0, g_cs);174 175 k_b = ggml_cont(ctx0, ggml_transpose(ctx0, k_b));176 177 // [CS, S_k, n_chunks, H_k * n_seqs]178 ggml_tensor * kbg = ggml_mul(ctx0, k_b, g_exp);179 cb(kbg, "k_beta_g_exp", il);180 181 // [S_k, CS, n_chunks, H_k * n_seqs]182 ggml_tensor * k_cd = ggml_mul_mat(ctx0, kbg, attn);183 cb(k_cd, "k_cumdecay", il);184 185 // [1, CS, n_chunks, H_k * n_seqs] KDA: [S_k, CS, n_chunks, H_k * n_seqs]186 ggml_tensor * g_exp_t = ggml_cont(ctx0, ggml_transpose(ctx0, g_exp));187 ggml_tensor * q_g_exp = ggml_mul(ctx0, q, g_exp_t);188 189 // vectorized calculation of key_gdiff190 // improved from the chunked version:191 // g_last = torch.clamp(g_cum[:, :, -1], max=50.0).exp().unsqueeze(-1).unsqueeze(-1)192 // g_diff = torch.clamp(g_cum[:, :, -1:] - g_cum, max=50.0).exp()193 // key_gdiff = key * g_diff.unsqueeze(-1)194 // kgdmulvnew = (key_gdiff).transpose(-1, -2) @ v_new195 // last_recurrent_state = last_recurrent_state * g_last + kgdmulvnew196 197 // get last element in g_cumsum along CS dimension (ne0)198 // example: [[x, y, z, ..., last], ...] -> [[last], ...]199 // [1, 1, n_chunks, H_v * n_seqs] KDA: [1, S_k, n_chunks, H_v * n_seqs]200 ggml_tensor * g_last = ggml_view_4d(ctx0, g_cs, 1, g_cs->ne[1], g_cs->ne[2], g_cs->ne[3],201 g_cs->nb[1],202 g_cs->nb[2],203 g_cs->nb[3],204 ggml_row_size(g_cs->type, g_cs->ne[0] - 1));205 cb(g_last, "g_last", il);206 207 // TODO: remove this cont when CUDA supports non-cont unary ops208 g_last = ggml_cont(ctx0, g_last);209 210 // [1, 1, n_chunks, H_v * n_seqs] KDA: [S_k, 1, n_chunks, H_v * n_seqs]211 ggml_tensor * g_last_exp_t = ggml_transpose(ctx0, ggml_exp(ctx0, g_last));212 cb(g_last_exp_t, "g_last_exp_t", il);213 214 // [CS, 1, n_chunks, H_v * n_seqs] KDA: [CS, S_k, n_chunks, H_v * n_seqs]215 ggml_tensor * g_diff = ggml_neg(ctx0, ggml_sub(ctx0, g_cs, g_last));216 cb(g_diff, "g_diff", il);217 218 ggml_tensor * g_diff_exp_t = ggml_cont(ctx0, ggml_transpose(ctx0, ggml_exp(ctx0, g_diff)));219 220 // [S_k, CS, n_chunks, H_v * n_seqs]221 ggml_tensor * kg = ggml_mul(ctx0, k, g_diff_exp_t);222 cb(kg, "key_gdiff", il);223 224 // [CS, S_k, n_chunks, H_v * n_seqs]225 ggml_tensor * kg_t = ggml_cont(ctx0, ggml_transpose(ctx0, kg));226 cb(kg_t, "key_gdiff_t", il);227 228 s = ggml_reshape_4d(ctx0, s, S_v, S_v, 1, H_v * n_seqs);229 cb(s, "dnet_add_ch_state", il);230 231 // [CS, S_v, n_chunks, H_v * n_seqs]232 ggml_tensor * v_t = ggml_cont(ctx0, ggml_transpose(ctx0, v));233 234 for (int64_t chunk = 0; chunk < n_chunks; chunk++) {235 ggml_tensor * ch_k_cd = get_slice_2d(ctx0, k_cd, chunk); // [S_k, CS, 1, H_k * n_seqs]236 ggml_tensor * ch_v_t = get_slice_2d(ctx0, v_t, chunk); // [ CS, S_v, 1, H_v * n_seqs]237 ggml_tensor * ch_kq = get_slice_2d(ctx0, kq, chunk); // [ CS, CS, 1, H_k * n_seqs]238 ggml_tensor * ch_q_g_exp = get_slice_2d(ctx0, q_g_exp, chunk); // [S_k, CS, 1, H_k * n_seqs]239 ggml_tensor * ch_kg_t = get_slice_2d(ctx0, kg_t, chunk); // [ CS, S_k, 1, H_v * n_seqs]240 241 // [CS, S_v, 1, H_v * n_seqs]242 ggml_tensor * v_t_p = ggml_mul_mat(ctx0, ch_k_cd, s);243 cb(v_t_p, "v_prime", il);244 245 // [CS, S_v, 1, H_v * n_seqs]246 ggml_tensor * v_t_new = ggml_sub(ctx0, ch_v_t, v_t_p);247 cb(v_t_new, "v_t_new", il);248 249 // [S_v, CS, 1, H_v * n_seqs]250 ggml_tensor * v_attn = ggml_mul_mat(ctx0, v_t_new, ch_kq);251 cb(v_attn, "v_attn", il);252 253 // [S_v, CS, 1, H_v * n_seqs]254 ggml_tensor * attn_inter = ggml_mul_mat(ctx0, s, ch_q_g_exp);255 cb(attn_inter, "attn_inter", il);256 257 // [S_v, CS, 1, H_v * n_seqs]258 ggml_tensor * o_ch = ggml_add(ctx0, attn_inter, v_attn);259 cb(o_ch, "dnet_add_ch_attn_out", il);260 261 v = ggml_set_inplace(ctx0, v, o_ch, v->nb[1], v->nb[2], v->nb[3], chunk * v->nb[2]);262 263 // kgdmulvnew = (key_gdiff).transpose(-1, -2) @ v_new264 // TODO: head broadcast might not work here - probably will need a transpose265 ggml_tensor * kgv = ggml_mul_mat(ctx0, ch_kg_t, v_t_new); // [S_k, S_v, 1, H_k * n_seqs]266 267 // last_recurrent_state = last_recurrent_state * g_last + kgdmulvnew268 ggml_tensor * ch_g_last_exp_t = get_slice_2d(ctx0, g_last_exp_t, chunk);269 270 s = ggml_mul(ctx0, s, ch_g_last_exp_t);271 s = ggml_add(ctx0, s, kgv);272 cb(s, "dnet_add_ch_state", il);273 }274 275 // truncate padded tokens276 ggml_tensor * o = ggml_view_4d(ctx0, v,277 S_v, n_tokens, H_v, n_seqs,278 ggml_row_size(v->type, S_v),279 ggml_row_size(v->type, S_v * CS * n_chunks),280 ggml_row_size(v->type, S_v * CS * n_chunks * H_v), 0);281 o = ggml_permute (ctx0, o, 0, 2, 1, 3); // [S_v, H_v, n_tokens, n_seqs]282 s = ggml_reshape_4d(ctx0, s, S_v, S_v, H_v, n_seqs);283 cb(s, "output_state", il);284 285 return {o, s};286}287 288std::pair<ggml_tensor *, ggml_tensor *> llm_build_delta_net_base::build_delta_net_autoregressive(289 ggml_tensor * q,290 ggml_tensor * k,291 ggml_tensor * v,292 ggml_tensor * g,293 ggml_tensor * b, // beta294 ggml_tensor * s, // state295 int il) {296 const int64_t S_k = q->ne[0];297 const int64_t H_k = q->ne[1];298 const int64_t n_tokens = q->ne[2];299 const int64_t n_seqs = q->ne[3];300 301 const int64_t S_v = v->ne[0];302 const int64_t H_v = v->ne[1];303 304 GGML_ASSERT(n_tokens == 1);305 306 GGML_ASSERT(S_k == S_v);307 GGML_ASSERT(H_v % H_k == 0);308 309 GGML_ASSERT(q->ne[0] == S_k && q->ne[1] == H_k && q->ne[2] == n_tokens && q->ne[3] == n_seqs);310 GGML_ASSERT(k->ne[0] == S_k && k->ne[1] == H_k && k->ne[2] == n_tokens && k->ne[3] == n_seqs);311 GGML_ASSERT(v->ne[0] == S_v && v->ne[1] == H_v && v->ne[2] == n_tokens && v->ne[3] == n_seqs);312 313 GGML_ASSERT(g->ne[0] == 1 || g->ne[0] == S_v);314 GGML_ASSERT( g->ne[1] == H_v && g->ne[2] == n_tokens && g->ne[3] == n_seqs);315 GGML_ASSERT(b->ne[0] == 1 && b->ne[1] == H_v && b->ne[2] == n_tokens && b->ne[3] == n_seqs);316 GGML_ASSERT(s->ne[0] == S_v && s->ne[1] == S_v && s->ne[2] == H_v && s->ne[3] == n_seqs);317 318 const float scale = 1.0f / sqrtf(S_k);319 320 q = ggml_scale(ctx0, q, scale);321 322 q = ggml_permute(ctx0, q, 0, 2, 1, 3); // [S_k, n_tokens, H_k, n_seqs]323 k = ggml_permute(ctx0, k, 0, 2, 1, 3); // [S_k, n_tokens, H_k, n_seqs]324 v = ggml_permute(ctx0, v, 0, 2, 1, 3); // [S_v, n_tokens, H_v, n_seqs]325 326 cb(q, "q_in", il);327 cb(k, "k_in", il);328 cb(v, "v_in", il);329 cb(b, "b_in", il);330 cb(g, "g_in", il);331 332 // GDA: [1, 1, H_v, n_seqs]333 // KDA: [1, S_k, H_v, n_seqs]334 g = ggml_reshape_4d(ctx0, g, 1, g->ne[0], H_v, n_seqs);335 b = ggml_reshape_4d(ctx0, b, 1, 1, H_v, n_seqs);336 337 // [S_v, S_v, H_v, n_seqs]338 g = ggml_exp(ctx0, g);339 s = ggml_mul(ctx0, s, g);340 341 // [1, S_v, H_v, n_seqs]342 ggml_tensor * sk;343 sk = ggml_mul (ctx0, s, k);344 sk = ggml_sum_rows(ctx0, sk);345 346 // [S_v, 1, H_v, n_seqs]347 ggml_tensor * d;348 d = ggml_sub(ctx0, v, ggml_transpose(ctx0, sk));349 d = ggml_mul(ctx0, d, b);350 351 // [1, S_v, H_v, n_seqs]352 ggml_tensor * d_t;353 d_t = ggml_transpose(ctx0, d);354 355 // [S_v, S_v, H_v, n_seqs]356 ggml_tensor * kd;357 k = ggml_repeat(ctx0, k, s);358 kd = ggml_mul (ctx0, k, d_t);359 360 s = ggml_add(ctx0, s, kd);361 362 cb(s, "dnet_add_ar_state", il);363 364 ggml_tensor * s_q = ggml_mul (ctx0, s, q);365 ggml_tensor * o = ggml_sum_rows(ctx0, s_q);366 367 o = ggml_permute (ctx0, o, 2, 0, 1, 3); // [S_v, H_v, n_tokens, n_seqs]368 369 return {o, s};370}371 372std::pair<ggml_tensor *, ggml_tensor *> llm_build_delta_net_base::build_delta_net_fused(373 ggml_tensor * q,374 ggml_tensor * k,375 ggml_tensor * v,376 ggml_tensor * g,377 ggml_tensor * b,378 ggml_tensor * s,379 int il) {380 const int64_t S_k = q->ne[0];381 const int64_t H_k = q->ne[1];382 const int64_t n_tokens = q->ne[2];383 const int64_t n_seqs = q->ne[3];384 385 const int64_t S_v = v->ne[0];386 const int64_t H_v = v->ne[1];387 388 GGML_ASSERT(S_k == S_v);389 GGML_ASSERT(H_v % H_k == 0);390 391 GGML_ASSERT(q->ne[0] == S_k && q->ne[1] == H_k && q->ne[2] == n_tokens && q->ne[3] == n_seqs);392 GGML_ASSERT(k->ne[0] == S_k && k->ne[1] == H_k && k->ne[2] == n_tokens && k->ne[3] == n_seqs);393 GGML_ASSERT(v->ne[0] == S_v && v->ne[1] == H_v && v->ne[2] == n_tokens && v->ne[3] == n_seqs);394 395 GGML_ASSERT(g->ne[0] == 1 || g->ne[0] == S_v);396 GGML_ASSERT( g->ne[1] == H_v && g->ne[2] == n_tokens && g->ne[3] == n_seqs);397 GGML_ASSERT(b->ne[0] == 1 && b->ne[1] == H_v && b->ne[2] == n_tokens && b->ne[3] == n_seqs);398 GGML_ASSERT(s->ne[0] == S_v && s->ne[1] == S_v && s->ne[2] == H_v && s->ne[3] == n_seqs);399 400 ggml_tensor * result = ggml_gated_delta_net(ctx0, q, k, v, g, b, s);401 if (n_tokens == 1) {402 cb(result, LLAMA_TENSOR_NAME_FGDN_AR, il);403 } else {404 cb(result, LLAMA_TENSOR_NAME_FGDN_CH, il);405 }406 407 ggml_tensor * output = ggml_view_4d(ctx0, result,408 S_v, H_v, n_tokens, n_seqs,409 ggml_row_size(result->type, S_v),410 ggml_row_size(result->type, S_v * H_v),411 ggml_row_size(result->type, S_v * H_v * n_tokens), 0);412 413 ggml_tensor * new_state = ggml_view_4d(ctx0, result,414 S_v, S_v, H_v, n_seqs,415 ggml_row_size(result->type, S_v),416 ggml_row_size(result->type, S_v * S_v),417 ggml_row_size(result->type, S_v * S_v * H_v),418 ggml_row_size(result->type, S_v * H_v * n_tokens * n_seqs));419 420 return {output, new_state};421}422 423std::pair<ggml_tensor *, ggml_tensor *> llm_build_delta_net_base::build_delta_net(424 ggml_tensor * q,425 ggml_tensor * k,426 ggml_tensor * v,427 ggml_tensor * g,428 ggml_tensor * b,429 ggml_tensor * s,430 int il) {431 const int64_t n_seq_tokens = q->ne[2];432 433 if (n_seq_tokens == 1) {434 if (cparams.fused_gdn_ar) {435 return build_delta_net_fused(q, k, v, g, b, s, il);436 }437 return build_delta_net_autoregressive(q, k, v, g, b, s, il);438 }439 440 if (cparams.fused_gdn_ch) {441 return build_delta_net_fused(q, k, v, g, b, s, il);442 }443 444 return build_delta_net_chunking(q, k, v, g, b, s, il);445}446 