import java.util.stream.IntStream;

/**
 * De forward pass van een Qwen2-model, volledig in fp32 &mdash; mijlpaal M3.
 *
 * <p>Dit is de referentie-implementatie: bewust eenvoudig, zonder SIMD, zonder listige
 * geheugentrucs. Elke bewerking staat er zoals ze in de formule staat. Traag mag &mdash;
 * het doel van M3 is <em>juist</em>, en juistheid wordt extern getoetst door greedy
 * decodering te vergelijken met llama.cpp (via ollama) op hetzelfde modelbestand.
 *
 * <p>De opbouw per laag (de qwen2-graf van llama.cpp):
 * <ol>
 *   <li>RMSNorm met {@code attn_norm.weight};</li>
 *   <li>Q, K, V als matvec m&eacute;t bias (qwen2 heeft biassen op precies deze drie);</li>
 *   <li>RoPE in <b>NEOX-stijl</b> op Q en K: element i draait tegen element i + kop/2 &mdash;
 *       n&iacute;et tegen zijn buurman. Dit is h&eacute;t verschil met de llama-architectuur,
 *       en de klassieke stille fout;</li>
 *   <li>aandacht met GQA: 14 vraagkoppen delen 2 sleutel/waarde-koppen (7 op 1),
 *       schaal 1/&radic;64, causale softmax over de KV-cache;</li>
 *   <li>uitgangsprojectie (zonder bias) en residu;</li>
 *   <li>RMSNorm met {@code ffn_norm.weight}, dan SwiGLU:
 *       {@code down( silu(gate(x)) ⊙ up(x) )}, en weer residu.</li>
 * </ol>
 * Na de laatste laag volgt RMSNorm met {@code output_norm.weight} en de eindprojectie.
 * Dit model heeft gebonden inbeddingen: de logits komen uit dezelfde matrix als de
 * inbedding ({@code token_embd.weight}), er is geen {@code output.weight}.
 *
 * <p>De gewichten blijven gekwantiseerd in het gemapte bestand liggen; {@link Dequant}
 * pakt per matvec &eacute;&eacute;n rij tegelijk uit in een scratchbuffer. Zo kost het
 * model vrijwel geen geheugen (alleen de KV-cache groeit met de context).
 */
final class Qwen2 implements AutoCloseable, Motor {

    final Gguf g;
    final int nLayers, dim, nHeads, nKvHeads, headDim, kvDim, gqa, ffnDim, vocab, ctx;
    final float ropeBase, rmsEps;

    private record Laag(float[] attnNorm, Tensor wq, float[] bq, Tensor wk, float[] bk,
                        Tensor wv, float[] bv, Tensor wo,
                        float[] ffnNorm, Tensor wGate, Tensor wUp, Tensor wDown) {}

    private final Laag[] lagen;
    private final Tensor tokenEmbd, lmHead;
    private final float[] outputNorm;
    private final float[] lmHeadBias;             // qwen2 staat een optionele output.bias toe

    // scratch: één keer aangemaakt, elke stap hergebruikt — niets alloceren in de hete lus
    private final float[] x, xb, q, k, v, attnOut, gate, up, logits;
    private final float[][] kCache, vCache;      // per laag: ctx * kvDim
    private final float[] bsum16;                // bloksommen van de matvec-invoer, voor de K-kernen
    private int pos = 0;

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

    /** De snelle kernen staan standaard aan; {@code -Dqllm.kernels=uit} dwingt de referentie af. */
    final boolean snel;

    Qwen2(Gguf g, int maxContext) {
        this(g, maxContext, !"uit".equals(System.getProperty("qllm.kernels")));
    }

    Qwen2(Gguf g, int maxContext, boolean snelleKernen) {
        this.g = g;
        String arch = g.architecture();
        if (!arch.equals("qwen2")) {
            throw new IllegalArgumentException("dit is een '" + arch + "'-model; deze forward pass "
                    + "kent alleen qwen2. De graf (rope-stijl, biassen) verschilt per architectuur, "
                    + "dus raden zou stil verkeerde antwoorden geven.");
        }
        nLayers  = g.getInt("qwen2.block_count");
        dim      = g.getInt("qwen2.embedding_length");
        nHeads   = g.getInt("qwen2.attention.head_count");
        nKvHeads = g.getInt("qwen2.attention.head_count_kv");
        if (nHeads <= 0 || nKvHeads <= 0 || nHeads % nKvHeads != 0 || dim % nHeads != 0) {
            throw new IllegalArgumentException("onbruikbare kopgeometrie: head_count=" + nHeads
                    + ", head_count_kv=" + nKvHeads + ", embedding_length=" + dim
                    + " (dim moet deelbaar zijn door de koppen, koppen door de kv-koppen)");
        }
        headDim  = dim / nHeads;
        kvDim    = nKvHeads * headDim;
        gqa      = nHeads / nKvHeads;
        ffnDim   = g.getInt("qwen2.feed_forward_length");
        ropeBase = g.getFloat("qwen2.rope.freq_base", 10000f);   // llama.cpp-standaard
        rmsEps   = g.getFloat("qwen2.attention.layer_norm_rms_epsilon");
        ctx      = Math.min(maxContext, g.getInt("qwen2.context_length"));
        if (ctx < 1) throw new IllegalArgumentException("context moet minstens 1 zijn, niet " + ctx);

        // Lange-context-varianten dragen rope-scaling (yarn/linear) mee. Die negeren zou
        // op elke positie stil verkeerde rotaties geven — dus weiger luid.
        if (g.has("qwen2.rope.scaling.type")) {
            String st = g.getString("qwen2.rope.scaling.type");
            if (!st.equals("none")) {
                throw new IllegalArgumentException("dit model vraagt rope-scaling '" + st
                        + "'; die is hier bewust niet geïmplementeerd (M3 is de fp32-referentie). "
                        + "Negeren zou stil verkeerde antwoorden geven.");
            }
        }
        // Afwijkende kopafmetingen in de metadata: alleen aanvaarden als ze met onze
        // afgeleide waarde overeenkomen, anders luid weigeren (llama.cpp doet hetzelfde).
        for (String sleutel : new String[]{"qwen2.rope.dimension_count",
                "qwen2.attention.key_length", "qwen2.attention.value_length"}) {
            if (g.has(sleutel) && g.getInt(sleutel) != headDim) {
                throw new IllegalArgumentException(sleutel + " = " + g.getInt(sleutel)
                        + " maar de kopdimensie is " + headDim + "; dat geval is niet ondersteund");
            }
        }

        tokenEmbd  = g.tensor("token_embd.weight");
        outputNorm = vector(g.tensor("output_norm.weight"));
        lmHead     = g.hasTiedEmbeddings() ? tokenEmbd : g.tensor("output.weight");
        Tensor ob  = g.tensorOrNull("output.bias");
        lmHeadBias = ob == null ? null : vector(ob);
        vocab      = (int) tokenEmbd.dims()[1];

        // Normgewichten en biassen zijn kleine F32-vectoren: één keer uitpakken bij het
        // laden in plaats van bij elke stap opnieuw. De Q/K/V-biassen zijn in de
        // architectuur optioneel (Qwen2.5 heeft ze; matvec aanvaardt null). Elke
        // gewichtsmatrix wordt op vorm gecontroleerd, zodat een afwijkend bestand hier
        // luid faalt in plaats van stil verkeerd te rekenen.
        lagen = new Laag[nLayers];
        for (int l = 0; l < nLayers; l++) {
            String p = "blk." + l + ".";
            lagen[l] = new Laag(
                vector(g.tensor(p + "attn_norm.weight"), dim),
                mat(p + "attn_q.weight", dim, dim),      biasOfNull(p + "attn_q.bias", dim),
                mat(p + "attn_k.weight", kvDim, dim),    biasOfNull(p + "attn_k.bias", kvDim),
                mat(p + "attn_v.weight", kvDim, dim),    biasOfNull(p + "attn_v.bias", kvDim),
                mat(p + "attn_output.weight", dim, dim),
                vector(g.tensor(p + "ffn_norm.weight"), dim),
                mat(p + "ffn_gate.weight", ffnDim, dim), mat(p + "ffn_up.weight", ffnDim, dim),
                mat(p + "ffn_down.weight", dim, ffnDim));
        }

        x = new float[dim]; xb = new float[dim];
        q = new float[dim]; k = new float[kvDim]; v = new float[kvDim];
        attnOut = new float[dim];
        gate = new float[ffnDim]; up = new float[ffnDim];
        logits = new float[vocab];
        kCache = new float[nLayers][ctx * kvDim];
        vCache = new float[nLayers][ctx * kvDim];
        bsum16 = new float[Math.max(dim, ffnDim) / 16];
        this.snel = snelleKernen;
    }

    /** Pakt een kleine 1D-tensor volledig uit naar een float-array. */
    private static float[] vector(Tensor t) {
        float[] v = new float[(int) t.elements()];
        Dequant.row(t, 0, v);
        return v;
    }

    private static float[] vector(Tensor t, int verwacht) {
        if (t.elements() != verwacht) {
            throw new IllegalArgumentException(t.name() + " heeft " + t.elements()
                    + " elementen, verwacht " + verwacht);
        }
        return vector(t);
    }

    /** Gewichtsmatrix met vormcontrole: {@code rijen} uitgangen van lengte {@code kolommen}. */
    private Tensor mat(String naam, long rijen, long kolommen) {
        Tensor t = g.tensor(naam);
        if (t.rows() != rijen || t.rowLength() != kolommen) {
            throw new IllegalArgumentException(naam + " heeft vorm " + t.shape()
                    + ", verwacht [" + kolommen + ", " + rijen + "]");
        }
        return t;
    }

    /** Optionele bias: null als de tensor ontbreekt, fout als hij de verkeerde lengte heeft. */
    private float[] biasOfNull(String naam, int lengte) {
        Tensor t = g.tensorOrNull(naam);
        return t == null ? null : vector(t, lengte);
    }

    public int position()  { return pos; }
    public int vocab()     { return vocab; }
    public String omschrijving() {
        return "%d lagen, dim %d, %d koppen, vocab %,d, %s".formatted(
                nLayers, dim, nHeads, vocab, snel ? "snelle kernen" : "fp32-referentie");
    }
    public int ctx()       { return ctx; }
    public void reset()    { pos = 0; }

    @Override public void close() { g.close(); }

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

    /**
     * Verwerkt één token op de huidige positie en geeft de logits voor het volgende terug.
     * De teruggegeven array is intern en wordt bij de volgende aanroep overschreven.
     */
    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 + " valt buiten de woordenschat (0.." + (vocab - 1) + ")");
        }

        Dequant.row(tokenEmbd, token, x);        // de inbedding: één rij opzoeken

        for (int l = 0; l < nLayers; l++) {
            Laag L = lagen[l];

            // -- aandacht --
            rmsnorm(xb, x, L.attnNorm(), rmsEps);
            matvec(L.wq(), xb, L.bq(), q);
            matvec(L.wk(), xb, L.bk(), k);
            matvec(L.wv(), xb, L.bv(), v);

            for (int h = 0; h < nHeads; h++)   ropeNeox(q, h * headDim, headDim, pos, ropeBase);
            for (int h = 0; h < nKvHeads; h++) ropeNeox(k, h * headDim, headDim, pos, ropeBase);

            System.arraycopy(k, 0, kCache[l], pos * kvDim, kvDim);
            System.arraycopy(v, 0, vCache[l], pos * kvDim, kvDim);

            attention(l);
            matvecAdd(L.wo(), attnOut, x);           // uitgangsprojectie + residu ineen

            // -- feed-forward --
            rmsnorm(xb, x, L.ffnNorm(), rmsEps);
            matvec(L.wGate(), xb, (float[]) null, gate);
            matvec(L.wUp(),   xb, (float[]) null, up);
            for (int i = 0; i < ffnDim; i++) {
                gate[i] = silu(gate[i]) * up[i];
            }
            matvecAdd(L.wDown(), gate, x);           // terugprojectie + residu ineen
        }

        rmsnorm(xb, x, outputNorm, rmsEps);
        matvec(lmHead, xb, lmHeadBias, logits);
        pos++;
        return logits;
    }

    /** Aandacht voor de huidige positie over alles in de cache, kop voor kop. */
    private void attention(int l) {
        float[] kc = kCache[l], vc = vCache[l];
        float schaal = (float) (1.0 / Math.sqrt(headDim));
        float[] scores = new float[pos + 1];
        for (int h = 0; h < nHeads; h++) {
            int qo = h * headDim;
            int kvo = (h / gqa) * headDim;           // GQA: 7 vraagkoppen per kv-kop
            for (int t = 0; t <= pos; t++) {
                double s = 0;
                int ko = t * kvDim + kvo;
                for (int i = 0; i < headDim; i++) s += q[qo + i] * kc[ko + i];
                scores[t] = (float) (s * schaal);
            }
            softmax(scores, pos + 1);
            for (int i = 0; i < headDim; i++) {
                double s = 0;
                for (int t = 0; t <= pos; t++) s += scores[t] * vc[t * kvDim + kvo + i];
                attnOut[qo + i] = (float) s;
            }
        }
    }

    // ---------------------------------------------------------------- bouwstenen

    /** RMSNorm: out = x / sqrt(gemiddelde(x²) + eps), elementsgewijs maal het gewicht. */
    static void rmsnorm(float[] out, float[] x, float[] w, float eps) {
        double ss = 0;
        for (float v : x) ss += (double) v * v;
        float schaal = (float) (1.0 / Math.sqrt(ss / x.length + eps));
        for (int i = 0; i < x.length; i++) out[i] = x[i] * schaal * w[i];
    }

    /** Softmax ter plekke over de eerste n elementen, numeriek stabiel. */
    static void softmax(float[] a, int n) {
        float max = a[0];
        for (int i = 1; i < n; i++) if (a[i] > max) max = a[i];
        double som = 0;
        for (int i = 0; i < n; i++) { a[i] = (float) Math.exp(a[i] - max); som += a[i]; }
        float inv = (float) (1.0 / som);
        for (int i = 0; i < n; i++) a[i] *= inv;
    }

    static float silu(float z) {
        return (float) (z / (1.0 + Math.exp(-z)));
    }

    /**
     * RoPE in NEOX-stijl ("rotate half"): binnen één kop draait element i tegen element
     * i + headDim/2, over hoek pos·base^(−2i/headDim). Toegepast op Q en K, nooit op V.
     */
    static void ropeNeox(float[] a, int off, int headDim, int pos, float base) {
        int half = headDim / 2;
        for (int i = 0; i < half; i++) {
            double freq = Math.pow(base, -2.0 * i / headDim);
            double hoek = pos * freq;
            float cos = (float) Math.cos(hoek), sin = (float) Math.sin(hoek);
            float a0 = a[off + i], a1 = a[off + i + half];
            a[off + i]        = a0 * cos - a1 * sin;
            a[off + i + half] = a0 * sin + a1 * cos;
        }
    }

    /**
     * out = W·x (+ bias). Rijen parallel over een handvol threads; wiskundig deterministisch.
     * Met de snelle kernen aan rekent elke rij rechtstreeks op de gekwantiseerde blokken
     * ({@link Kernel}); anders wordt de rij eerst uitgepakt (de referentieweg).
     */
    private void matvec(Tensor w, float[] x, float[] bias, float[] out) {
        int rows = (int) w.rows();
        int kk = (int) w.rowLength();
        boolean kern = snel && Kernel.kan(w.type());
        float[] sommen = kernVoorbereiding(w, x, kk, kern);
        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++) {
                    float s = Kernel.dot(w, r, x, sommen);
                    out[r] = bias == null ? s : s + bias[r];
                }
            } 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] * x[i];
                    out[r] = (float) (bias == null ? s : s + bias[r]);
                }
            }
        });
    }

    /** out += W·x — de matvec voor de twee residuverbindingen. */
    private void matvecAdd(Tensor w, float[] x, float[] out) {
        int rows = (int) w.rows();
        int kk = (int) w.rowLength();
        boolean kern = snel && Kernel.kan(w.type());
        float[] sommen = kernVoorbereiding(w, x, kk, kern);
        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++) out[r] += Kernel.dot(w, r, x, 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] * x[i];
                    out[r] += (float) s;
                }
            }
        });
    }

    /** De K-kernen hebben de bloksommen van x nodig; één keer per matvec volstaat. */
    private float[] kernVoorbereiding(Tensor w, float[] x, int kk, boolean kern) {
        if (!kern || (w.type() != GgmlType.Q4_K && w.type() != GgmlType.Q6_K)) return null;
        Kernel.blokSommen16(x, kk, bsum16);
        return bsum16;
    }
}
