Felipe97/llama-cpp-compiled
01.1k
1#include "models.h"2 3#include "llama-impl.h"4#include "llama-memory-recurrent.h"5 6// utility to get one slice from the third dimension7// input dim: [x, y, c, b]8// output dim: [x, y, 1, b]9static ggml_tensor * get_slice_2d(ggml_context * ctx0, ggml_tensor * t, int64_t c) {10 return ggml_view_4d(ctx0, t, t->ne[0], t->ne[1], 1, t->ne[3],11 t->nb[1], t->nb[2], t->nb[3], t->nb[2] * c);12}13 14llm_build_delta_net_base::llm_build_delta_net_base(const llm_graph_params & params) : llm_graph_context(params) {}15 16std::pair<ggml_tensor *, ggml_tensor *> llm_build_delta_net_base::build_delta_net_chunking(17 ggml_tensor * q,18 ggml_tensor * k,19 ggml_tensor * v,20 ggml_tensor * g,21 ggml_tensor * b,22 ggml_tensor * s,23 int il) {24 const int64_t S_k = q->ne[0];25 const int64_t H_k = q->ne[1];26 const int64_t n_tokens = q->ne[2];27 const int64_t n_seqs = q->ne[3];28 29 const int64_t S_v = v->ne[0];30 const int64_t H_v = v->ne[1];31 const bool kda = (g->ne[0] == S_k && g->ne[1] == H_k);32 33 GGML_ASSERT(S_k == S_v);34 GGML_ASSERT(H_v % H_k == 0);35 36 GGML_ASSERT(q->ne[0] == S_k && q->ne[1] == H_k && q->ne[2] == n_tokens && q->ne[3] == n_seqs);37 GGML_ASSERT(k->ne[0] == S_k && k->ne[1] == H_k && k->ne[2] == n_tokens && k->ne[3] == n_seqs);38 GGML_ASSERT(v->ne[0] == S_v && v->ne[1] == H_v && v->ne[2] == n_tokens && v->ne[3] == n_seqs);39 40 GGML_ASSERT(g->ne[0] == 1 || g->ne[0] == S_v);41 GGML_ASSERT( g->ne[1] == H_v && g->ne[2] == n_tokens && g->ne[3] == n_seqs);42 GGML_ASSERT(b->ne[0] == 1 && b->ne[1] == H_v && b->ne[2] == n_tokens && b->ne[3] == n_seqs);43 GGML_ASSERT(s->ne[0] == S_v && s->ne[1] == S_v && s->ne[2] == H_v && s->ne[3] == n_seqs);44 45 const float scale = 1.0f / sqrtf(S_k);46 47 q = ggml_scale(ctx0, q, scale);48 49 cb(q, "q_in", il);50 cb(k, "k_in", il);51 cb(v, "v_in", il);52 cb(b, "b_in", il);53 cb(g, "g_in", il);54 55 q = ggml_permute(ctx0, q, 0, 2, 1, 3); // [S_k, n_tokens, H_k, n_seqs]56 k = ggml_permute(ctx0, k, 0, 2, 1, 3); // [S_k, n_tokens, H_k, n_seqs]57 v = ggml_permute(ctx0, v, 0, 2, 1, 3); // [S_v, n_tokens, H_v, n_seqs]58 g = ggml_permute(ctx0, g, 0, 2, 1, 3); // [g_0, n_tokens, H_v, n_seqs]59 b = ggml_permute(ctx0, b, 0, 2, 1, 3); // [ 1, n_tokens, H_v, n_seqs]60 61 const int CS = kda ? 16 : 64; // chunk size62 63 const int pad = (CS - n_tokens % CS) % CS;64 const int n_chunks = (n_tokens + pad) / CS;65 66 q = ggml_pad(ctx0, q, 0, pad, 0, 0);67 k = ggml_pad(ctx0, k, 0, pad, 0, 0);68 v = ggml_pad(ctx0, v, 0, pad, 0, 0);69 g = ggml_pad(ctx0, g, 0, pad, 0, 0);70 b = ggml_pad(ctx0, b, 0, pad, 0, 0);71 72 ggml_tensor * v_b = ggml_mul(ctx0, v, b);73 ggml_tensor * k_b = ggml_mul(ctx0, k, b);74 75 cb(v_b, "v_b", il);76 cb(k_b, "k_b", il);77 78 q = ggml_reshape_4d(ctx0, q, S_k, CS, n_chunks, H_k * n_seqs);79 k = ggml_reshape_4d(ctx0, k, S_k, CS, n_chunks, H_k * n_seqs);80 k_b = ggml_reshape_4d(ctx0, k_b, S_k, CS, n_chunks, H_v * n_seqs);81 v = ggml_reshape_4d(ctx0, v, S_v, CS, n_chunks, H_v * n_seqs);82 v_b = ggml_reshape_4d(ctx0, v_b, S_v, CS, n_chunks, H_v * n_seqs);83 84 g = ggml_reshape_4d(ctx0, g, g->ne[0], CS, n_chunks, H_v * n_seqs);85 b = ggml_reshape_4d(ctx0, b, 1, CS, n_chunks, H_v * n_seqs);86 87 // [CS, g_0, n_chunks, H_v * n_seqs]88 // TODO: extend ggml_cumsum with axis parameter to avoid transpose89 ggml_tensor * g_cs = ggml_cumsum(ctx0, ggml_cont(ctx0, ggml_transpose(ctx0, g)));90 cb(g_cs, "g_cs", il);91 92 ggml_tensor * kb = nullptr;93 ggml_tensor * kq = nullptr;94 if (kda) {95 const int64_t CHB = n_chunks * H_k * n_seqs;96 97 ggml_tensor * g_cs_i = ggml_reshape_4d(ctx0, g_cs, CS, 1, S_k, CHB); // [chunk_size, 1, S_k, CHB]98 ggml_tensor * g_cs_j = ggml_reshape_4d(ctx0, g_cs, 1, CS, S_k, CHB); // [1, chunk_size, S_k, CHB]99 100 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]101 102 // decay_mask [chunk_size,chunk_size,S_k,CHB]103 ggml_tensor * decay_mask;104 decay_mask = ggml_sub(ctx0, g_cs_j, g_cs_i);105 decay_mask = ggml_tri(ctx0, decay_mask, GGML_TRI_TYPE_LOWER_DIAG);106 decay_mask = ggml_exp(ctx0, decay_mask);107 cb(decay_mask, "decay_mask", il);108 109 // decay_mask [S_k,BT_j,BT_i,CHB] *Note* second and third chunk_sizes are switched110 decay_mask = ggml_cont_4d(ctx0, ggml_permute(ctx0, decay_mask, 2, 1, 0, 3), S_k, CS, CS, CHB);111 112 ggml_tensor * k_b_i = ggml_reshape_4d(ctx0, k_b, S_k, CS, 1, CHB);113 ggml_tensor * k_j = ggml_reshape_4d(ctx0, k, S_k, 1, CS, CHB);114 ggml_tensor * q_i = ggml_reshape_4d(ctx0, q, S_k, CS, 1, CHB);115 116 ggml_tensor * decay_k_b_i = ggml_mul(ctx0, decay_mask, k_b_i);117 ggml_tensor * decay_q_i = ggml_mul(ctx0, decay_mask, q_i);118 119 // decay_k_b_i [S,BT,BT,CHB] @ k_j [S,1,BT,CHB] = Akk [BT,1,BT,CHB]120 kb = ggml_mul_mat(ctx0, decay_k_b_i, k_j);121 kq = ggml_mul_mat(ctx0, decay_q_i, k_j);122 123 kb = ggml_cont(ctx0, ggml_transpose(ctx0, ggml_reshape_4d(ctx0, kb, CS, CS, n_chunks, H_v * n_seqs)));124 kq = ggml_cont(ctx0, ggml_transpose(ctx0, ggml_reshape_4d(ctx0, kq, CS, CS, n_chunks, H_v * n_seqs)));125 } else {126 ggml_tensor * g_cs_i = g_cs;127 ggml_tensor * g_cs_j = ggml_reshape_4d(ctx0, g_cs, 1, CS, n_chunks, H_v * n_seqs);128 129 g_cs_j = ggml_repeat_4d(ctx0, g_cs_j, CS, CS, n_chunks, H_v * n_seqs);130 131 // [CS, CS, n_chunks, H_v * n_seqs]132 ggml_tensor * decay_mask;133 decay_mask = ggml_sub(ctx0, g_cs_j, g_cs_i);134 decay_mask = ggml_tri(ctx0, decay_mask, GGML_TRI_TYPE_LOWER_DIAG);135 decay_mask = ggml_exp(ctx0, decay_mask);136 cb(decay_mask, "decay_mask", il);137 138 // [CS, CS, n_chunks, H_k * n_seqs]139 kb = ggml_mul_mat(ctx0, k, k_b);140 kb = ggml_mul (ctx0, kb, decay_mask);141 142 // [CS, CS, n_chunks, H_k * n_seqs]143 kq = ggml_mul_mat(ctx0, k, q);144 kq = ggml_mul(ctx0, kq, decay_mask);145 }146 147 kq = ggml_tri(ctx0, kq, GGML_TRI_TYPE_LOWER_DIAG);148 cb(kq, "kq", il);149 150 // [CS, CS, n_chunks, H_k * n_seqs]151 ggml_tensor * attn;152 attn = ggml_tri(ctx0, kb, GGML_TRI_TYPE_LOWER);153 cb(attn, "attn", il);154 155 ggml_tensor * identity;156 identity = ggml_view_1d(ctx0, attn, CS, 0);157 identity = ggml_fill (ctx0, identity, 1.0f);158 identity = ggml_diag (ctx0, identity);159 160 ggml_tensor * lhs = ggml_add(ctx0, attn, identity);161 cb(lhs, "dnet_add_ch_lhs", il);162 163 attn = ggml_neg(ctx0, attn);164 cb(attn, "attn_pre_solve", il);165 166 ggml_tensor * lin_solve = ggml_solve_tri(ctx0, lhs, attn, true, true, false);167 attn = ggml_add(ctx0, lin_solve, identity);168 cb(attn, "dnet_add_ch_attn_solved", il); // [CS, CS, n_chunks, H_k * n_seqs]169 170 // [S_v, CS, n_chunks, H_v * n_seqs]171 v = ggml_mul_mat(ctx0, ggml_cont(ctx0, ggml_transpose(ctx0, v_b)), attn);172 173 // [CS, 1, n_chunks, H_v * n_seqs] KDA: [CS, S_k, n_chunks, H_v * n_seqs]174 ggml_tensor * g_exp = ggml_exp(ctx0, g_cs);175 176 k_b = ggml_cont(ctx0, ggml_transpose(ctx0, k_b));177 178 // [CS, S_k, n_chunks, H_k * n_seqs]179 ggml_tensor * kbg = ggml_mul(ctx0, k_b, g_exp);180 cb(kbg, "k_beta_g_exp", il);181 182 // [S_k, CS, n_chunks, H_k * n_seqs]183 ggml_tensor * k_cd = ggml_mul_mat(ctx0, kbg, attn);184 cb(k_cd, "k_cumdecay", il);185 186 // [1, CS, n_chunks, H_k * n_seqs] KDA: [S_k, CS, n_chunks, H_k * n_seqs]187 ggml_tensor * g_exp_t = ggml_cont(ctx0, ggml_transpose(ctx0, g_exp));188 ggml_tensor * q_g_exp = ggml_mul(ctx0, q, g_exp_t);189 190 // vectorized calculation of key_gdiff191 // improved from the chunked version:192 // g_last = torch.clamp(g_cum[:, :, -1], max=50.0).exp().unsqueeze(-1).unsqueeze(-1)193 // g_diff = torch.clamp(g_cum[:, :, -1:] - g_cum, max=50.0).exp()194 // key_gdiff = key * g_diff.unsqueeze(-1)195 // kgdmulvnew = (key_gdiff).transpose(-1, -2) @ v_new196 // last_recurrent_state = last_recurrent_state * g_last + kgdmulvnew197 198 // get last element in g_cumsum along CS dimension (ne0)199 // example: [[x, y, z, ..., last], ...] -> [[last], ...]200 // [1, 1, n_chunks, H_v * n_seqs] KDA: [1, S_k, n_chunks, H_v * n_seqs]201 ggml_tensor * g_last = ggml_view_4d(ctx0, g_cs, 1, g_cs->ne[1], g_cs->ne[2], g_cs->ne[3],202 g_cs->nb[1],203 g_cs->nb[2],204 g_cs->nb[3],205 ggml_row_size(g_cs->type, g_cs->ne[0] - 1));206 cb(g_last, "g_last", il);207 208 // TODO: remove this cont when CUDA supports non-cont unary ops209 g_last = ggml_cont(ctx0, g_last);210 211 // [1, 1, n_chunks, H_v * n_seqs] KDA: [S_k, 1, n_chunks, H_v * n_seqs]212 ggml_tensor * g_last_exp_t = ggml_transpose(ctx0, ggml_exp(ctx0, g_last));213 cb(g_last_exp_t, "g_last_exp_t", il);214 215 // [CS, 1, n_chunks, H_v * n_seqs] KDA: [CS, S_k, n_chunks, H_v * n_seqs]216 ggml_tensor * g_diff = ggml_neg(ctx0, ggml_sub(ctx0, g_cs, g_last));217 cb(g_diff, "g_diff", il);218 219 ggml_tensor * g_diff_exp_t = ggml_cont(ctx0, ggml_transpose(ctx0, ggml_exp(ctx0, g_diff)));220 221 // [S_k, CS, n_chunks, H_v * n_seqs]222 ggml_tensor * kg = ggml_mul(ctx0, k, g_diff_exp_t);223 cb(kg, "key_gdiff", il);224 225 // [CS, S_k, n_chunks, H_v * n_seqs]226 ggml_tensor * kg_t = ggml_cont(ctx0, ggml_transpose(ctx0, kg));227 cb(kg_t, "key_gdiff_t", il);228 229 s = ggml_reshape_4d(ctx0, s, S_v, S_v, 1, H_v * n_seqs);230 cb(s, "dnet_add_ch_state", il);231 232 // [CS, S_v, n_chunks, H_v * n_seqs]233 ggml_tensor * v_t = ggml_cont(ctx0, ggml_transpose(ctx0, v));234 235 for (int64_t chunk = 0; chunk < n_chunks; chunk++) {236 ggml_tensor * ch_k_cd = get_slice_2d(ctx0, k_cd, chunk); // [S_k, CS, 1, H_k * n_seqs]237 ggml_tensor * ch_v_t = get_slice_2d(ctx0, v_t, chunk); // [ CS, S_v, 1, H_v * n_seqs]238 ggml_tensor * ch_kq = get_slice_2d(ctx0, kq, chunk); // [ CS, CS, 1, H_k * n_seqs]239 ggml_tensor * ch_q_g_exp = get_slice_2d(ctx0, q_g_exp, chunk); // [S_k, CS, 1, H_k * n_seqs]240 ggml_tensor * ch_kg_t = get_slice_2d(ctx0, kg_t, chunk); // [ CS, S_k, 1, H_v * n_seqs]241 242 // [CS, S_v, 1, H_v * n_seqs]243 ggml_tensor * v_t_p = ggml_mul_mat(ctx0, ch_k_cd, s);244 cb(v_t_p, "v_prime", il);245 246 // [CS, S_v, 1, H_v * n_seqs]247 ggml_tensor * v_t_new = ggml_sub(ctx0, ch_v_t, v_t_p);248 cb(v_t_new, "v_t_new", il);249 250 // [S_v, CS, 1, H_v * n_seqs]251 ggml_tensor * v_attn = ggml_mul_mat(ctx0, v_t_new, ch_kq);252 cb(v_attn, "v_attn", il);253 254 // [S_v, CS, 1, H_v * n_seqs]255 ggml_tensor * attn_inter = ggml_mul_mat(ctx0, s, ch_q_g_exp);256 cb(attn_inter, "attn_inter", il);257 258 // [S_v, CS, 1, H_v * n_seqs]259 ggml_tensor * o_ch = ggml_add(ctx0, attn_inter, v_attn);260 cb(o_ch, "dnet_add_ch_attn_out", il);261 262 v = ggml_set_inplace(ctx0, v, o_ch, v->nb[1], v->nb[2], v->nb[3], chunk * v->nb[2]);263 264 // kgdmulvnew = (key_gdiff).transpose(-1, -2) @ v_new265 // TODO: head broadcast might not work here - probably will need a transpose266 ggml_tensor * kgv = ggml_mul_mat(ctx0, ch_kg_t, v_t_new); // [S_k, S_v, 1, H_k * n_seqs]267 268 // last_recurrent_state = last_recurrent_state * g_last + kgdmulvnew269 ggml_tensor * ch_g_last_exp_t = get_slice_2d(ctx0, g_last_exp_t, chunk);270 271 s = ggml_mul(ctx0, s, ch_g_last_exp_t);272 s = ggml_add(ctx0, s, kgv);273 cb(s, "dnet_add_ch_state", il);274 }275 276 // truncate padded tokens277 ggml_tensor * o = ggml_view_4d(ctx0, v,278 S_v, n_tokens, H_v, n_seqs,279 ggml_row_size(v->type, S_v),280 ggml_row_size(v->type, S_v * CS * n_chunks),281 ggml_row_size(v->type, S_v * CS * n_chunks * H_v), 0);282 o = ggml_permute (ctx0, o, 0, 2, 1, 3); // [S_v, H_v, n_tokens, n_seqs]283 s = ggml_reshape_4d(ctx0, s, S_v, S_v, H_v, n_seqs);284 cb(s, "output_state", il);285 286 return {o, s};287}288 289std::pair<ggml_tensor *, ggml_tensor *> llm_build_delta_net_base::build_delta_net_autoregressive(290 ggml_tensor * q,291 ggml_tensor * k,292 ggml_tensor * v,293 ggml_tensor * g,294 ggml_tensor * b, // beta295 ggml_tensor * s, // state296 int il) {297 const int64_t S_k = q->ne[0];298 const int64_t H_k = q->ne[1];299 const int64_t n_tokens = q->ne[2];300 const int64_t n_seqs = q->ne[3];301 302 const int64_t S_v = v->ne[0];303 const int64_t H_v = v->ne[1];304 305 GGML_ASSERT(n_tokens == 1);306 307 GGML_ASSERT(S_k == S_v);308 GGML_ASSERT(H_v % H_k == 0);309 310 GGML_ASSERT(q->ne[0] == S_k && q->ne[1] == H_k && q->ne[2] == n_tokens && q->ne[3] == n_seqs);311 GGML_ASSERT(k->ne[0] == S_k && k->ne[1] == H_k && k->ne[2] == n_tokens && k->ne[3] == n_seqs);312 GGML_ASSERT(v->ne[0] == S_v && v->ne[1] == H_v && v->ne[2] == n_tokens && v->ne[3] == n_seqs);313 314 GGML_ASSERT(g->ne[0] == 1 || g->ne[0] == S_v);315 GGML_ASSERT( g->ne[1] == H_v && g->ne[2] == n_tokens && g->ne[3] == n_seqs);316 GGML_ASSERT(b->ne[0] == 1 && b->ne[1] == H_v && b->ne[2] == n_tokens && b->ne[3] == n_seqs);317 GGML_ASSERT(s->ne[0] == S_v && s->ne[1] == S_v && s->ne[2] == H_v && s->ne[3] == n_seqs);318 319 const float scale = 1.0f / sqrtf(S_k);320 321 q = ggml_scale(ctx0, q, scale);322 323 q = ggml_permute(ctx0, q, 0, 2, 1, 3); // [S_k, n_tokens, H_k, n_seqs]324 k = ggml_permute(ctx0, k, 0, 2, 1, 3); // [S_k, n_tokens, H_k, n_seqs]325 v = ggml_permute(ctx0, v, 0, 2, 1, 3); // [S_v, n_tokens, H_v, n_seqs]326 327 cb(q, "q_in", il);328 cb(k, "k_in", il);329 cb(v, "v_in", il);330 cb(b, "b_in", il);331 cb(g, "g_in", il);332 333 // GDA: [1, 1, H_v, n_seqs]334 // KDA: [1, S_k, H_v, n_seqs]335 g = ggml_reshape_4d(ctx0, g, 1, g->ne[0], H_v, n_seqs);336 b = ggml_reshape_4d(ctx0, b, 1, 1, H_v, n_seqs);337 338 // [S_v, S_v, H_v, n_seqs]339 g = ggml_exp(ctx0, g);340 s = ggml_mul(ctx0, s, g);341 342 // [1, S_v, H_v, n_seqs]343 ggml_tensor * sk;344 sk = ggml_mul (ctx0, s, k);345 sk = ggml_sum_rows(ctx0, sk);346 347 // [S_v, 1, H_v, n_seqs]348 ggml_tensor * d;349 d = ggml_sub(ctx0, v, ggml_transpose(ctx0, sk));350 d = ggml_mul(ctx0, d, b);351 352 // [1, S_v, H_v, n_seqs]353 ggml_tensor * d_t;354 d_t = ggml_transpose(ctx0, d);355 356 // [S_v, S_v, H_v, n_seqs]357 ggml_tensor * kd;358 k = ggml_repeat(ctx0, k, s);359 kd = ggml_mul (ctx0, k, d_t);360 361 s = ggml_add(ctx0, s, kd);362 363 cb(s, "dnet_add_ar_state", il);364 365 ggml_tensor * s_q = ggml_mul (ctx0, s, q);366 ggml_tensor * o = ggml_sum_rows(ctx0, s_q);367 368 o = ggml_permute (ctx0, o, 2, 0, 1, 3); // [S_v, H_v, n_tokens, n_seqs]369 370 return {o, s};371}372 373std::pair<ggml_tensor *, ggml_tensor *> llm_build_delta_net_base::build_delta_net_fused(374 ggml_tensor * q,375 ggml_tensor * k,376 ggml_tensor * v,377 ggml_tensor * g,378 ggml_tensor * b,379 ggml_tensor * s,380 int il) {381 const int64_t S_k = q->ne[0];382 const int64_t H_k = q->ne[1];383 const int64_t n_tokens = q->ne[2];384 const int64_t n_seqs = q->ne[3];385 386 const int64_t S_v = v->ne[0];387 const int64_t H_v = v->ne[1];388 389 GGML_ASSERT(S_k == S_v);390 GGML_ASSERT(H_v % H_k == 0);391 392 GGML_ASSERT(q->ne[0] == S_k && q->ne[1] == H_k && q->ne[2] == n_tokens && q->ne[3] == n_seqs);393 GGML_ASSERT(k->ne[0] == S_k && k->ne[1] == H_k && k->ne[2] == n_tokens && k->ne[3] == n_seqs);394 GGML_ASSERT(v->ne[0] == S_v && v->ne[1] == H_v && v->ne[2] == n_tokens && v->ne[3] == n_seqs);395 396 GGML_ASSERT(g->ne[0] == 1 || g->ne[0] == S_v);397 GGML_ASSERT( g->ne[1] == H_v && g->ne[2] == n_tokens && g->ne[3] == n_seqs);398 GGML_ASSERT(b->ne[0] == 1 && b->ne[1] == H_v && b->ne[2] == n_tokens && b->ne[3] == n_seqs);399 GGML_ASSERT(s->ne[0] == S_v && s->ne[1] == S_v && s->ne[2] == H_v && s->ne[3] == n_seqs);400 401 // K=1: output carries the final state only. state s is 4D [S_v, S_v, H_v, n_seqs].402 ggml_tensor * result = ggml_gated_delta_net(ctx0, q, k, v, g, b, s, /*K=*/1);403 if (n_tokens == 1) {404 res->add_fused_node({LLM_FUSED_OP_GDN_AR, result, il});405 } else {406 res->add_fused_node({LLM_FUSED_OP_GDN_CH, result, il});407 }408 409 ggml_tensor * output = ggml_view_4d(ctx0, result,410 S_v, H_v, n_tokens, n_seqs,411 ggml_row_size(result->type, S_v),412 ggml_row_size(result->type, S_v * H_v),413 ggml_row_size(result->type, S_v * H_v * n_tokens), 0);414 415 ggml_tensor * new_state = ggml_view_4d(ctx0, result,416 S_v, S_v, H_v, n_seqs,417 ggml_row_size(result->type, S_v),418 ggml_row_size(result->type, S_v * S_v),419 ggml_row_size(result->type, S_v * S_v * H_v),420 ggml_row_size(result->type, S_v * H_v * n_tokens * n_seqs));421 422 return {output, new_state};423}424 425std::pair<ggml_tensor *, ggml_tensor *> llm_build_delta_net_base::build_delta_net(426 ggml_tensor * q,427 ggml_tensor * k,428 ggml_tensor * v,429 ggml_tensor * g,430 ggml_tensor * b,431 ggml_tensor * s,432 int il) {433 const int64_t n_seq_tokens = q->ne[2];434 435 if (n_seq_tokens == 1) {436 if (cparams.fused_gdn_ar) {437 return build_delta_net_fused(q, k, v, g, b, s, il);438 }439 return build_delta_net_autoregressive(q, k, v, g, b, s, il);440 }441 442 if (cparams.fused_gdn_ch) {443 return build_delta_net_fused(q, k, v, g, b, s, il);444 }445 446 return build_delta_net_chunking(q, k, v, g, b, s, il);447}448 449ggml_tensor * llm_build_delta_net_base::build_conv_state(450 llm_graph_input_rs * inp,451 ggml_tensor * conv_states_all,452 ggml_tensor * qkv_mixed,453 int64_t conv_kernel_size,454 int64_t conv_channels,455 int il) {456 const auto * mctx_cur = inp->mctx;457 458 const auto kv_head = mctx_cur->get_head();459 const auto mem_size = mctx_cur->get_size();460 461 const int64_t n_seqs = ubatch.n_seqs;462 463 ggml_tensor * conv_states = build_rs(inp, conv_states_all, hparams.n_embd_r(), n_seqs);464 cb(conv_states, "conv_states", il);465 466 conv_states = ggml_reshape_3d(ctx0, conv_states, conv_kernel_size - 1, conv_channels, n_seqs);467 cb(conv_states, "conv_states_reshaped", il);468 469 qkv_mixed = ggml_transpose(ctx0, qkv_mixed);470 cb(qkv_mixed, "qkv_mixed_transposed", il);471 472 ggml_tensor * conv_input = ggml_concat(ctx0, conv_states, qkv_mixed, 0);473 cb(conv_input, "conv_input", il);474 475 const int64_t row_count = (conv_kernel_size - 1) * conv_channels;476 477 const size_t row_size = ggml_row_size(conv_states_all->type, row_count);478 479 if (cparams.n_rs_seq == 0) {480 const int64_t s_idx = conv_input->ne[0] - conv_states->ne[0];481 const int64_t s_slot = 0;482 483 ggml_tensor * conv_state_last =484 ggml_view_3d(ctx0, conv_input,485 conv_kernel_size - 1, conv_channels, n_seqs,486 conv_input->nb[1], conv_input->nb[2],487 ggml_row_size(conv_input->type, s_idx));488 cb(conv_state_last, "conv_state_last", il);489 490 ggml_tensor * conv_state_update =491 ggml_view_2d(ctx0, conv_states_all,492 row_count, n_seqs, conv_states_all->nb[1],493 (s_slot * mem_size + kv_head) * row_size);494 cb(conv_state_update, "conv_state_update", il);495 496 ggml_build_forward_expand(gf, ggml_cpy(ctx0, conv_state_last, conv_state_update));497 } else {498 // [TAG_RECURRENT_ROLLBACK_SPLITS]499 // this logic assumes that the last (n_rs_seq + 1) tokens of a sequence in a batch are inside500 // the same ubatch, which `split_equal()` guarantees via its n_keep_tail argument501 502 const int64_t K = (int64_t) cparams.n_rs_seq + 1;503 504 for (int64_t t = 1; t <= K; ++t) {505 const int64_t s_idx = std::max<int64_t>(0, conv_input->ne[0] - conv_states->ne[0] - K + t);506 const int64_t s_slot = K - t;507 508 ggml_tensor * conv_state_last =509 ggml_view_3d(ctx0, conv_input,510 conv_kernel_size - 1, conv_channels, n_seqs,511 conv_input->nb[1], conv_input->nb[2],512 ggml_row_size(conv_input->type, s_idx));513 514 ggml_tensor * conv_state_update =515 ggml_view_2d(ctx0,516 conv_states_all, row_count, n_seqs,517 conv_states_all->nb[1],518 (s_slot * mem_size + kv_head) * row_size);519 520 ggml_build_forward_expand(gf, ggml_cpy(ctx0, conv_state_last, conv_state_update));521 }522 }523 524 return conv_input;525}526 527ggml_tensor * llm_build_delta_net_base::build_recurrent_attn(528 llm_graph_input_rs * inp,529 ggml_tensor * ssm_states_all,530 ggml_tensor * q,531 ggml_tensor * k,532 ggml_tensor * v,533 ggml_tensor * g,534 ggml_tensor * b,535 ggml_tensor * s,536 int il) {537 const auto * mctx_cur = inp->mctx;538 const auto kv_head = mctx_cur->get_head();539 const uint32_t mem_size = mctx_cur->get_size();540 541 const int64_t S_v = s->ne[0];542 const int64_t H_v = s->ne[2];543 const int64_t n_seqs = s->ne[3];544 const int64_t n_seq_tokens = q->ne[2];545 546 const bool keep = cparams.n_rs_seq > 0;547 548 if (!keep) {549 auto attn_out = build_delta_net(q, k, v, g, b, s, il);550 ggml_tensor * output = attn_out.first;551 ggml_tensor * new_state = attn_out.second;552 cb(output, "attn_output", il);553 cb(new_state, "new_state", il);554 555 ggml_build_forward_expand(gf,556 ggml_cpy(ctx0, new_state,557 ggml_view_2d(ctx0, ssm_states_all, hparams.n_embd_s(), n_seqs, ssm_states_all->nb[1],558 kv_head * hparams.n_embd_s() * ggml_element_size(ssm_states_all))));559 560 return output;561 }562 563 const int64_t D = S_v * S_v * H_v;564 const int64_t K = cparams.n_rs_seq + 1;565 566 // state s is 4D [S_v, S_v, H_v, n_seqs]; K snapshot slots are written into the output.567 ggml_tensor * gdn_out = ggml_gated_delta_net(ctx0, q, k, v, g, b, s, K);568 if (n_seq_tokens > 1) {569 res->add_fused_node({LLM_FUSED_OP_GDN_CH, gdn_out, il});570 } else {571 res->add_fused_node({LLM_FUSED_OP_GDN_AR, gdn_out, il});572 }573 574 const int64_t attn_score_elems = S_v * H_v * n_seq_tokens * n_seqs;575 const int64_t state_size_per_snap = S_v * S_v * H_v * n_seqs;576 577 ggml_tensor * output = ggml_view_4d(ctx0, gdn_out,578 S_v, H_v, n_seq_tokens, n_seqs,579 ggml_row_size(gdn_out->type, S_v),580 ggml_row_size(gdn_out->type, S_v * H_v),581 ggml_row_size(gdn_out->type, S_v * H_v * n_seq_tokens),582 0);583 cb(output, "attn_output", il);584 585 const size_t row_size = hparams.n_embd_s() * ggml_element_size(ssm_states_all);586 587 // op writes the last min(n_seq_tokens, K) snapshots; trailing slots are left unwritten588 const int64_t n_written = std::min<int64_t>(n_seq_tokens, K);589 590 // write the produced snapshots into the recurrent cache (snapshot slot i -> rollback group i)591 ggml_tensor * src = ggml_view_3d(ctx0, gdn_out,592 D, n_seqs, n_written,593 ggml_row_size(gdn_out->type, D),594 ggml_row_size(gdn_out->type, state_size_per_snap),595 ggml_row_size(gdn_out->type, attn_score_elems));596 597 ggml_tensor * dst = ggml_view_3d(ctx0, ssm_states_all,598 D, n_seqs, n_written,599 ssm_states_all->nb[1],600 (size_t) mem_size * row_size,601 (size_t) kv_head * row_size);602 603 ggml_build_forward_expand(gf, ggml_cpy(ctx0, src, dst));604 605 return output;606}607 