import java.util.stream.IntStream;

/**
 * De tweede motor &mdash; mijlpaal M9: de hybride qwen35-graf (Qwen3.5 / Qwen3.8).
 *
 * <p>Waar {@link Qwen2} een klassieke transformer is (elke laag aandacht met een
 * groeiende KV-cache), wisselt qwen35 twee laagsoorten af in een vast patroon van
 * drie-om-&eacute;&eacute;n:
 * <ol>
 *   <li><b>Gated DeltaNet-lagen</b> (driekwart): geen terugkijken naar alle vorige
 *       tokens, maar een geheugenmatrix van v&aacute;ste grootte per kop die per token
 *       vervalt (exp(g)), wordt uitgelezen en met een delta-regel wordt bijgewerkt.
 *       De invoer gaat eerst door een korte causale convolutie (kernel 4) met SiLU,
 *       q en k worden per kop L2-genormaliseerd, en de uitvoer gaat door een
 *       ge-gate RMSNorm (&middot;silu(z)).</li>
 *   <li><b>Volledige-aandachtslagen</b> (een kwart): zoals hoofdstuk 14, maar met
 *       drie nieuwigheden: RMSNorm op q en k per kop v&oacute;&oacute;r RoPE, RoPE op maar
 *       een kwart van de kopdimensies (64 van de 256), en een sigmoid-uitgangspoort
 *       die met de q-projectie is meegebakken.</li>
 * </ol>
 *
 * <p>De opbouw volgt de llama.cpp-referentie (src/models/qwen35.cpp en
 * delta-net-base.cpp, het autoregressieve pad) operatie voor operatie; het
 * gereedschap van de eerdere mijlpalen (mmap-lader, kernen, sampler, sjabloon)
 * wordt ongewijzigd hergebruikt. Multimodale tensoren ({@code v.*}) en een
 * eventueel MTP-blok worden genegeerd: dit is de tekstmotor.
 */
public final class Qwen35 implements Motor {

    private static final int THREADS = Integer.getInteger("qllm.threads",
            Math.min(8, Runtime.getRuntime().availableProcessors()));

    // ---------------------------------------------------------------- configuratie
    final int nLayers, dim, nHeads, headDim, ffnDim, vocab;
    public final int ctx;
    final float ropeBase, rmsEps;
    final int ropeDims;                              // 64: RoPE op een kwart van de kop
    final int[] kvKoppen;                            // per laag: 0 = DeltaNet, >0 = aandacht
    // DeltaNet-maten
    final int dConv, nK, dState, nV, headV, dInner, kanalen;

    private final Gguf g;
    private final boolean snel;

    // ---------------------------------------------------------------- gewichten
    private final Tensor embd;
    private final Tensor lmHead;                     // los output.weight, of de gebonden embd
    private final float[] outNorm;
    private final Laag[] lagen;

    private final class Laag {
        final float[] attnNorm, postNorm;
        final Tensor ffnGate, ffnUp, ffnDown;
        // aandacht (null bij DeltaNet)
        Tensor wq, wk, wv, wo;
        float[] qNorm, kNorm;
        int kvH;
        // DeltaNet (null bij aandacht)
        Tensor wqkv, wz, wBeta, wAlpha, wOut;
        float[] a, dtBias, conv, ssmNorm;

        Laag(int i, boolean aandacht) {
            String p = "blk." + i + ".";
            attnNorm = vec(p + "attn_norm.weight", dim);
            postNorm = vec(p + "post_attention_norm.weight", dim);
            ffnGate = mat(p + "ffn_gate.weight", ffnDim, dim);
            ffnUp   = mat(p + "ffn_up.weight", ffnDim, dim);
            ffnDown = mat(p + "ffn_down.weight", dim, ffnDim);
            if (aandacht) {
                kvH = kvKoppen[i];
                if (nHeads % kvH != 0) {
                    throw new IllegalArgumentException("laag " + i + ": " + nHeads
                            + " koppen niet deelbaar door " + kvH + " kv-koppen");
                }
                wq = mat(p + "attn_q.weight", 2 * headDim * nHeads, dim);
                wk = mat(p + "attn_k.weight", kvH * headDim, dim);
                wv = mat(p + "attn_v.weight", kvH * headDim, dim);
                wo = mat(p + "attn_output.weight", dim, headDim * nHeads);
                qNorm = vec(p + "attn_q_norm.weight", headDim);
                kNorm = vec(p + "attn_k_norm.weight", headDim);
            } else {
                wqkv   = mat(p + "attn_qkv.weight", kanalen, dim);
                wz     = mat(p + "attn_gate.weight", dInner, dim);
                wBeta  = mat(p + "ssm_beta.weight", nV, dim);
                wAlpha = mat(p + "ssm_alpha.weight", nV, dim);
                wOut   = mat(p + "ssm_out.weight", dim, dInner);
                a      = vec(p + "ssm_a", nV);
                dtBias = vec(p + "ssm_dt", nV);
                ssmNorm = vec(p + "ssm_norm.weight", headV);
                conv = new float[dConv * kanalen];   // [tap][kanaal], tap 0 = oudste
                Tensor c = mat(p + "ssm_conv1d.weight", kanalen, dConv);
                float[] rij = new float[dConv];
                for (int ch = 0; ch < kanalen; ch++) {
                    Dequant.row(c, ch, rij);
                    for (int t = 0; t < dConv; t++) conv[t * kanalen + ch] = rij[t];
                }
            }
        }
    }

    // ---------------------------------------------------------------- toestand
    private int pos;
    private final float[][] kCache, vCache;          // alleen voor aandachtslagen
    private final float[][] convStaat;               // [laag][(dConv-1) * kanalen]
    private final float[][] S;                       // [laag][nV * headV * dState], S[i][j] met i contigu

    // scratch
    private final float[] x, xb, xb2, q4, kv1, att2, ffnA, ffnB, mix, z1, o2, logits;

    public Qwen35(Gguf gg, int maxContext) {
        this(gg, maxContext, !"uit".equals(System.getProperty("qllm.kernels")));
    }

    public Qwen35(Gguf gg, int maxContext, boolean snelleKernen) {
        this.g = gg;
        this.snel = snelleKernen;
        String arch = g.architecture();
        if (!arch.equals("qwen35")) {
            throw new IllegalArgumentException("dit is een '" + arch + "'-model; deze motor kent "
                    + "alleen qwen35 (gebruik Qwen2 voor qwen2-modellen).");
        }
        nLayers = g.getInt("qwen35.block_count");
        dim     = g.getInt("qwen35.embedding_length");
        nHeads  = g.getInt("qwen35.attention.head_count");
        headDim = g.getInt("qwen35.attention.key_length");
        ffnDim  = g.getInt("qwen35.feed_forward_length");
        ropeBase = g.getFloat("qwen35.rope.freq_base", 10000f);
        rmsEps  = g.getFloat("qwen35.attention.layer_norm_rms_epsilon");
        ropeDims = g.getInt("qwen35.rope.dimension_count");
        ctx     = Math.min(maxContext, g.getInt("qwen35.context_length"));
        if (ctx < 1) throw new IllegalArgumentException("context " + ctx + " is kleiner dan 1");
        if (g.has("qwen35.rope.scaling.type")) {
            String st = g.getString("qwen35.rope.scaling.type");
            if (!st.isEmpty() && !st.equals("none")) {
                throw new IllegalArgumentException("dit model vraagt rope-scaling '" + st
                        + "'; die is hier niet gebouwd — negeren zou stil verkeerde rotaties geven");
            }
        }
        long[] kvArr = g.getLongArray("qwen35.attention.head_count_kv");
        if (kvArr.length != nLayers) {
            throw new IllegalArgumentException("head_count_kv heeft " + kvArr.length
                    + " waarden voor " + nLayers + " lagen");
        }
        kvKoppen = new int[nLayers];
        for (int i = 0; i < nLayers; i++) kvKoppen[i] = (int) kvArr[i];

        dConv  = g.getInt("qwen35.ssm.conv_kernel");
        nK     = g.getInt("qwen35.ssm.group_count");
        dState = g.getInt("qwen35.ssm.state_size");
        nV     = g.getInt("qwen35.ssm.time_step_rank");
        dInner = g.getInt("qwen35.ssm.inner_size");
        headV  = dInner / nV;
        kanalen = dInner + 2 * nK * dState;
        if (headV != dState) {
            throw new IllegalArgumentException("v-kopdimensie " + headV
                    + " != toestandsdimensie " + dState + "; deze motor gaat van gelijke maten uit");
        }
        if (nV % nK != 0) {
            throw new IllegalArgumentException("v-koppen " + nV + " niet deelbaar door k-koppen " + nK);
        }

        embd = mat("token_embd.weight", -1, dim);
        vocab = (int) embd.rows();
        // kleine qwen3.5-modellen delen de inbeddingen met de lm-head; het 27B-bestand
        // heeft een losse output.weight — beide smaken, net als bij qwen2
        Tensor los = g.tensorOrNull("output.weight");
        lmHead = los != null ? mat("output.weight", vocab, dim) : embd;
        outNorm = vec("output_norm.weight", dim);

        lagen = new Laag[nLayers];
        kCache = new float[nLayers][];
        vCache = new float[nLayers][];
        convStaat = new float[nLayers][];
        S = new float[nLayers][];
        for (int i = 0; i < nLayers; i++) {
            boolean aandacht = kvKoppen[i] > 0;
            lagen[i] = new Laag(i, aandacht);
            if (aandacht) {
                kCache[i] = new float[ctx * kvKoppen[i] * headDim];
                vCache[i] = new float[ctx * kvKoppen[i] * headDim];
            } else {
                convStaat[i] = new float[(dConv - 1) * kanalen];
                S[i] = new float[nV * headV * dState];
            }
        }

        x = new float[dim]; xb = new float[dim]; xb2 = new float[dim];
        q4 = new float[2 * headDim * nHeads];
        kv1 = new float[nHeads * headDim];           // ruim genoeg voor k- of v-projectie
        att2 = new float[nHeads * headDim];
        ffnA = new float[ffnDim]; ffnB = new float[ffnDim];
        mix = new float[kanalen]; z1 = new float[dInner]; o2 = new float[dInner];
        logits = new float[vocab];
        bsum16Scratch = new float[(Math.max(Math.max(dim, ffnDim), Math.max(dInner, nHeads * headDim)) + 15) / 16];
    }

    public int position() { return pos; }
    public int vocab()    { return vocab; }
    public String omschrijving() {
        int recurrent = 0;
        for (int kv : kvKoppen) if (kv == 0) recurrent++;
        return "qwen35-hybride: %d lagen (%d DeltaNet + %d aandacht), dim %d, vocab %,d, %s".formatted(
                nLayers, recurrent, nLayers - recurrent, dim, vocab,
                snel ? "snelle kernen" : "fp32-referentie");
    }
    public int ctx() { return ctx; }

    public void reset() {
        pos = 0;
        for (int i = 0; i < nLayers; i++) {
            if (convStaat[i] != null) {
                java.util.Arrays.fill(convStaat[i], 0f);
                java.util.Arrays.fill(S[i], 0f);
            }
        }
    }

    // ---------------------------------------------------------------- de forward pass

    public float[] forward(int token) {
        if (pos >= ctx) throw new IllegalStateException("context vol (" + ctx + " tokens)");
        if (token < 0 || token >= vocab) throw new IllegalArgumentException("token " + token);
        Dequant.row(embd, token, x);

        for (int l = 0; l < nLayers; l++) {
            Laag laag = lagen[l];
            rmsnorm(x, laag.attnNorm, xb);
            if (laag.wq != null) aandacht(l, laag); else deltaNet(l, laag);
            for (int i = 0; i < dim; i++) x[i] += xb2[i];
            rmsnorm(x, laag.postNorm, xb);
            matvec(laag.ffnGate, xb, ffnA);
            matvec(laag.ffnUp, xb, ffnB);
            for (int i = 0; i < ffnDim; i++) ffnA[i] = silu(ffnA[i]) * ffnB[i];
            matvec(laag.ffnDown, ffnA, xb2);
            for (int i = 0; i < dim; i++) x[i] += xb2[i];
        }

        rmsnorm(x, outNorm, xb);
        matvec(lmHead, xb, logits);
        pos++;
        return logits;
    }

    /** De volledige-aandachtslaag: QK-norm, partiële RoPE, GQA, en de sigmoid-uitgangspoort. */
    private void aandacht(int l, Laag laag) {
        int kvH = laag.kvH;
        int groep = nHeads / kvH;

        // q en poort zitten samen in één projectie: per kop eerst q, dan de poort
        matvec(laag.wq, xb, q4);
        matvec(laag.wk, xb, kv1);
        for (int h = 0; h < kvH; h++) {
            int ko = h * headDim;
            rmsnormKop(kv1, ko, laag.kNorm);
            ropeDeel(kv1, ko);
        }
        System.arraycopy(kv1, 0, kCache[l], pos * kvH * headDim, kvH * headDim);
        matvec(laag.wv, xb, kv1);
        System.arraycopy(kv1, 0, vCache[l], pos * kvH * headDim, kvH * headDim);

        float schaal = (float) (1.0 / Math.sqrt(headDim));
        float[] kc = kCache[l], vc = vCache[l];
        for (int h = 0; h < nHeads; h++) {
            int qo = h * 2 * headDim;                // q van kop h
            rmsnormKop(q4, qo, laag.qNorm);
            ropeDeel(q4, qo);

            int kvKop = h / groep;
            float[] scores = new float[pos + 1];
            float max = Float.NEGATIVE_INFINITY;
            for (int p = 0; p <= pos; p++) {
                int basis = p * kvH * headDim + kvKop * headDim;
                double s = 0;
                for (int i = 0; i < headDim; i++) s += (double) q4[qo + i] * kc[basis + i];
                scores[p] = (float) (s * schaal);
                if (scores[p] > max) max = scores[p];
            }
            double som = 0;
            for (int p = 0; p <= pos; p++) { scores[p] = (float) Math.exp(scores[p] - max); som += scores[p]; }
            float inv = (float) (1.0 / som);
            int uo = h * headDim;
            java.util.Arrays.fill(att2, uo, uo + headDim, 0f);
            for (int p = 0; p <= pos; p++) {
                float w = scores[p] * inv;
                int basis = p * kvH * headDim + kvKop * headDim;
                for (int i = 0; i < headDim; i++) att2[uo + i] += w * vc[basis + i];
            }
            // de uitgangspoort: de tweede helft van het q-blok, sigmoid, vóór de projectie
            int go = qo + headDim;
            for (int i = 0; i < headDim; i++) att2[uo + i] *= sigmoid(q4[go + i]);
        }
        matvec(laag.wo, att2, xb2);
    }

    /** De Gated DeltaNet-laag: conv → SiLU → L2-norm → delta-regel → ge-gate norm. */
    private void deltaNet(int l, Laag laag) {
        matvec(laag.wqkv, xb, mix);
        matvec(laag.wz, xb, z1);

        // g en beta per v-kop: g = -exp(A_log) · softplus(alpha + dt_bias); beta = sigmoid
        float[] gKop = new float[nV], bKop = new float[nV];
        float[] proj = new float[nV];
        matvec(laag.wAlpha, xb, proj);
        for (int h = 0; h < nV; h++) gKop[h] = laag.a[h] * softplus(proj[h] + laag.dtBias[h]);
        matvec(laag.wBeta, xb, proj);
        for (int h = 0; h < nV; h++) bKop[h] = sigmoid(proj[h]);

        // causale conv (kernel dConv) per kanaal over [staat..., dit token], daarna SiLU;
        // de staat bewaart de rúwe projecties van de laatste dConv-1 tokens
        float[] staat = convStaat[l];
        float[] w = laag.conv;
        int st = dConv - 1;
        for (int c = 0; c < kanalen; c++) {
            double s = w[st * kanalen + c] * mix[c];
            for (int t = 0; t < st; t++) s += (double) w[t * kanalen + c] * staat[t * kanalen + c];
            float uit = silu((float) s);
            // staat doorschuiven en de nieuwe rauwe waarde achteraan
            for (int t = 0; t < st - 1; t++) staat[t * kanalen + c] = staat[(t + 1) * kanalen + c];
            staat[(st - 1) * kanalen + c] = mix[c];
            mix[c] = uit;                            // mix bevat nu de conv+silu-uitkomst
        }

        // splitsen: q | k | v, L2-norm per kop op q en k, schaal q met 1/sqrt(dState)
        int qOff = 0, kOff = nK * dState, vOff = 2 * nK * dState;
        for (int h = 0; h < nK; h++) {
            l2norm(mix, qOff + h * dState, dState);
            l2norm(mix, kOff + h * dState, dState);
        }
        float qSchaal = (float) (1.0 / Math.sqrt(dState));

        // de delta-regel per v-kop (k/q-kop h % nK, zoals de cyclische herhaling in llama.cpp)
        float[] Sl = S[l];
        for (int h = 0; h < nV; h++) {
            int kh = h % nK;
            int qb = qOff + kh * dState, kb = kOff + kh * dState, vb = vOff + h * headV;
            int sb = h * headV * dState;
            float verval = (float) Math.exp(gKop[h]);
            float beta = bKop[h];
            for (int j = 0; j < headV; j++) {
                int rij = sb + j * dState;
                double sk = 0;
                for (int i = 0; i < dState; i++) {
                    Sl[rij + i] *= verval;
                    sk += (double) Sl[rij + i] * mix[kb + i];
                }
                float d = (float) ((mix[vb + j] - sk) * beta);
                double o = 0;
                for (int i = 0; i < dState; i++) {
                    Sl[rij + i] += mix[kb + i] * d;
                    o += (double) Sl[rij + i] * mix[qb + i];
                }
                o2[h * headV + j] = (float) (o * qSchaal);
            }
        }

        // ge-gate norm: RMSNorm per kop (gedeeld gewicht) en · silu(z), dan de uitprojectie
        for (int h = 0; h < nV; h++) {
            rmsnormDeel(o2, h * headV, headV, laag.ssmNorm);
            for (int j = 0; j < headV; j++) {
                int i = h * headV + j;
                o2[i] *= silu(z1[i]);
            }
        }
        matvec(laag.wOut, o2, xb2);
    }

    // ---------------------------------------------------------------- rekenhulpjes

    private void rmsnorm(float[] in, float[] w, float[] uit) {
        double som = 0;
        for (int i = 0; i < dim; i++) som += (double) in[i] * in[i];
        float schaal = (float) (1.0 / Math.sqrt(som / dim + rmsEps));
        for (int i = 0; i < dim; i++) uit[i] = in[i] * schaal * w[i];
    }

    private void rmsnormKop(float[] a, int off, float[] w) {
        rmsnormDeel(a, off, headDim, w);
    }

    private void rmsnormDeel(float[] a, int off, int n, float[] w) {
        double som = 0;
        for (int i = 0; i < n; i++) som += (double) a[off + i] * a[off + i];
        float schaal = (float) (1.0 / Math.sqrt(som / n + rmsEps));
        for (int i = 0; i < n; i++) a[off + i] = a[off + i] * schaal * w[i];
    }

    /** L2-normalisatie zoals ggml_l2_norm: x / max(|x|, eps). */
    private void l2norm(float[] a, int off, int n) {
        double som = 0;
        for (int i = 0; i < n; i++) som += (double) a[off + i] * a[off + i];
        float schaal = (float) (1.0 / Math.max(Math.sqrt(som), rmsEps));
        for (int i = 0; i < n; i++) a[off + i] *= schaal;
    }

    /** RoPE op alleen de eerste {@code ropeDims} dimensies van een kop (NEOX-paren). */
    private void ropeDeel(float[] a, int off) {
        Qwen2.ropeNeox(a, off, ropeDims, pos, ropeBase);
    }

    private static float silu(float v)    { return (float) (v / (1.0 + Math.exp(-v))); }
    private static float sigmoid(float v) { return (float) (1.0 / (1.0 + Math.exp(-v))); }
    private static float softplus(float v) { return v > 20f ? v : (float) Math.log1p(Math.exp(v)); }

    // ---------------------------------------------------------------- gewichten & matvec

    private float[] vec(String naam, int lengte) {
        Tensor t = g.tensor(naam);
        if (t.rows() != 1 || t.rowLength() != lengte) {
            throw new IllegalArgumentException(naam + " heeft lengte " + t.rowLength() + ", verwacht " + lengte);
        }
        float[] uit = new float[lengte];
        Dequant.row(t, 0, uit);
        return uit;
    }

    private Tensor mat(String naam, long rows, long cols) {
        Tensor t = g.tensor(naam);
        if ((rows >= 0 && t.rows() != rows) || t.rowLength() != cols) {
            throw new IllegalArgumentException(naam + " is " + t.rows() + "×" + t.rowLength()
                    + ", verwacht " + rows + "×" + cols);
        }
        return t;
    }

    private final float[] bsum16Scratch;

    private void matvec(Tensor w, float[] in, float[] uit) {
        int rows = (int) w.rows();
        int kk = (int) w.rowLength();
        boolean kern = snel && Kernel.kan(w.type());
        final float[] sommen;
        if (kern && (w.type() == GgmlType.Q4_K || w.type() == GgmlType.Q6_K)) {
            Kernel.blokSommen16(in, kk, bsum16Scratch);
            sommen = bsum16Scratch;
        } else {
            sommen = null;
        }
        int chunk = (rows + THREADS - 1) / THREADS;
        IntStream.range(0, THREADS).parallel().forEach(t -> {
            int van = t * chunk, tot = Math.min(rows, van + chunk);
            if (van >= tot) return;
            if (kern) {
                for (int r = van; r < tot; r++) uit[r] = Kernel.dot(w, r, in, sommen);
            } else {
                float[] rij = new float[kk];
                for (int r = van; r < tot; r++) {
                    Dequant.row(w, r, rij);
                    double s = 0;
                    for (int i = 0; i < kk; i++) s += (double) rij[i] * in[i];
                    uit[r] = (float) s;
                }
            }
        });
    }
}
