Team Ai
Datasetpublic

echodict/llama.cpp

version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes479downloads
delta-net-base.cpp446 linesDownload Raw Back to models
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