import java.lang.foreign.MemorySegment;
import java.lang.foreign.ValueLayout;
import java.nio.ByteOrder;

import jdk.incubator.vector.ByteVector;
import jdk.incubator.vector.FloatVector;
import jdk.incubator.vector.VectorMask;
import jdk.incubator.vector.VectorOperators;
import jdk.incubator.vector.VectorShuffle;
import jdk.incubator.vector.VectorSpecies;

/**
 * De snelle kernen &mdash; mijlpaal M4/M5.
 *
 * <p>E&eacute;n taak: het inwendig product van &eacute;&eacute;n gekwantiseerde rij met een
 * float-vector, rechtstreeks uit het gemapte bestand, zonder de rij ooit naar een
 * tussenbuffer uit te pakken. Dit is waar een taalmodel &gt;90&nbsp;% van zijn tijd zit,
 * dus dit is de enige plek waar SIMD de moeite loont. Alles daarbuiten blijft de
 * leesbare referentiecode.
 *
 * <p>De opbouw per blok is overal dezelfde: laad de ruwe bytes, wring de quants met
 * vectorbewerkingen in bytelanes, til ze met {@code castShape} naar floats en doe
 * fused multiply-add tegen x. De schaalfactor van het blok wordt als broadcast-fma in
 * &eacute;&eacute;n lopende accumulatorvector gevouwen, zodat de dure laan-reductie maar
 * &eacute;&eacute;n keer per rij gebeurt in plaats van per blok.
 *
 * <p>Twee trucs verdienen uitleg:
 * <ul>
 *   <li><b>Q5_0:</b> de 32 losse vijfde bits liggen als uint32 v&oacute;&oacute;r de nibbles.
 *       In plaats van ze bit voor bit uit te pakken worden de vier qh-bytes met een
 *       shuffle over de lanes uitgesmeerd en tegen een bitmasker getest; het
 *       resulterende masker telt er 16 bij op precies waar de vijfde bit aan staat.</li>
 *   <li><b>Q4_K en Q6_K:</b> die trekken per groepje een minimum of offset af
 *       (dmin&middot;m, respectievelijk &minus;32). Omdat &Sigma;(q&minus;c)&middot;x =
 *       &Sigma;q&middot;x &minus; c&middot;&Sigma;x hoeft dat niet per element: de bloksommen
 *       van x worden &eacute;&eacute;n keer per matvec voorberekend ({@link #blokSommen16})
 *       en gelden voor &aacute;lle rijen; de correctietermen lopen scalair mee.</li>
 * </ul>
 *
 * <p>Vereist {@code --add-modules jdk.incubator.vector}. De correctheid wordt bewaakt
 * door {@code KernelTest} (vergelijking met de scalaire referentie op echte
 * modelrijen) en uiteindelijk door {@code ForwardVsOllama}.
 */
final class Kernel {

    static final VectorSpecies<Float> FS = FloatVector.SPECIES_128;   // NEON: 4 float-lanes
    static final VectorSpecies<Byte>  BS = ByteVector.SPECIES_128;    // 16 bytes per load

    private static final ValueLayout.OfShort F16 =
            ValueLayout.JAVA_SHORT_UNALIGNED.withOrder(ByteOrder.LITTLE_ENDIAN);

    // Q5_0: smeer qh-byte 0 over lanes 0..7 en byte 1 over 8..15 (en idem 2/3).
    private static final VectorShuffle<Byte> SPREID_01 = VectorShuffle.fromValues(BS,
            0,0,0,0,0,0,0,0, 1,1,1,1,1,1,1,1);
    private static final VectorShuffle<Byte> SPREID_23 = VectorShuffle.fromValues(BS,
            2,2,2,2,2,2,2,2, 3,3,3,3,3,3,3,3);
    private static final ByteVector BITS = ByteVector.fromArray(BS, new byte[]{
            1,2,4,8,16,32,64,-128, 1,2,4,8,16,32,64,-128}, 0);

    /**
     * Laan-reductie in v&aacute;ste volgorde. {@code reduceLanes} mag per JIT-tier anders
     * associ&euml;ren (sequentieel in de interpreter, paarsgewijs in C2), waardoor de
     * allereerste doorloop minimaal zou afwijken van latere. Vier lanes scalair optellen
     * is exact even snel genoeg &eacute;n bit-deterministisch.
     */
    private static float som4(FloatVector v) {
        return ((v.lane(0) + v.lane(1)) + v.lane(2)) + v.lane(3);
    }

    private Kernel() {}

    /** Kan dit tensortype door een snelle kern? (Anders valt de matvec terug op de referentie.) */
    static boolean kan(GgmlType t) {
        return switch (t) {
            case Q8_0, Q5_0, Q4_K, Q6_K -> true;
            default -> false;
        };
    }

    /**
     * Inwendig product van rij {@code row} van tensor {@code t} met {@code x}.
     * {@code bsum16} zijn de bloksommen van x per 16 elementen; alleen nodig
     * voor de K-formaten (mag anders null zijn).
     */
    static float dot(Tensor t, long row, float[] x, float[] bsum16) {
        long base = row * t.rowBytes();
        int k = (int) t.rowLength();
        MemorySegment d = t.data();
        return switch (t.type()) {
            case Q8_0 -> dotQ8_0(d, base, x, k);
            case Q5_0 -> dotQ5_0(d, base, x, k);
            case Q4_K -> dotQ4_K(d, base, x, k, bsum16);
            case Q6_K -> dotQ6_K(d, base, x, k, bsum16);
            default -> throw new UnsupportedOperationException(t.type() + " heeft geen snelle kern");
        };
    }

    /** Bloksommen van x per 16 elementen: de gedeelde voorberekening voor de K-kernen. */
    static void blokSommen16(float[] x, int k, float[] uit) {
        for (int b = 0; b < k / 16; b++) {
            FloatVector s = FloatVector.zero(FS);
            int o = b * 16;
            for (int i = 0; i < 16; i += 4) {
                s = s.add(FloatVector.fromArray(FS, x, o + i));
            }
            uit[b] = som4(s);
        }
    }

    // ---------------------------------------------------------------- Q8_0: 34 bytes per 32

    static float dotQ8_0(MemorySegment d, long base, float[] x, int k) {
        FloatVector totaal = FloatVector.zero(FS);
        for (int b = 0; b < k / 32; b++) {
            long p = base + (long) b * 34;
            float s = Float.float16ToFloat(d.get(F16, p));
            FloatVector acc = FloatVector.zero(FS);
            int xo = b * 32;
            for (int helft = 0; helft < 2; helft++) {
                ByteVector q = ByteVector.fromMemorySegment(BS, d, p + 2 + helft * 16L, ByteOrder.LITTLE_ENDIAN);
                int xh = xo + helft * 16;
                for (int deel = 0; deel < 4; deel++) {
                    acc = ((FloatVector) q.castShape(FS, deel))
                            .fma(FloatVector.fromArray(FS, x, xh + deel * 4), acc);
                }
            }
            totaal = acc.fma(FloatVector.broadcast(FS, s), totaal);
        }
        return som4(totaal);
    }

    // ---------------------------------------------------------------- Q5_0: 22 bytes per 32

    static float dotQ5_0(MemorySegment d, long base, float[] x, int k) {
        FloatVector totaal = FloatVector.zero(FS);
        for (int b = 0; b < k / 32; b++) {
            long p = base + (long) b * 22;
            float s = Float.float16ToFloat(d.get(F16, p));

            // qh-bytes 0..3 zitten in de eerste vier lanes van deze load
            ByteVector kop = ByteVector.fromMemorySegment(BS, d, p + 2, ByteOrder.LITTLE_ENDIAN);
            ByteVector qs  = ByteVector.fromMemorySegment(BS, d, p + 6, ByteOrder.LITTLE_ENDIAN);

            ByteVector lo = qs.and((byte) 0x0F);
            ByteVector hi = qs.lanewise(VectorOperators.LSHR, 4).and((byte) 0x0F);

            // vijfde bits: bit j voor element j, bit j+16 voor element j+16
            VectorMask<Byte> bit0 = kop.rearrange(SPREID_01).and(BITS).eq((byte) 0).not();
            VectorMask<Byte> bit1 = kop.rearrange(SPREID_23).and(BITS).eq((byte) 0).not();
            ByteVector q0 = lo.add((byte) 16, bit0).sub((byte) 16);   // (nibble | bit<<4) - 16
            ByteVector q1 = hi.add((byte) 16, bit1).sub((byte) 16);

            FloatVector acc = FloatVector.zero(FS);
            int xo = b * 32;
            for (int deel = 0; deel < 4; deel++) {
                acc = ((FloatVector) q0.castShape(FS, deel))
                        .fma(FloatVector.fromArray(FS, x, xo + deel * 4), acc);
                acc = ((FloatVector) q1.castShape(FS, deel))
                        .fma(FloatVector.fromArray(FS, x, xo + 16 + deel * 4), acc);
            }
            totaal = acc.fma(FloatVector.broadcast(FS, s), totaal);
        }
        return som4(totaal);
    }

    // ---------------------------------------------------------------- Q4_K: 144 bytes per 256

    static float dotQ4_K(MemorySegment d, long base, float[] x, int k, float[] bsum16) {
        byte[] sc = new byte[12];
        FloatVector totaal = FloatVector.zero(FS);
        double minTerm = 0;
        for (int sb = 0; sb < k / 256; sb++) {
            long p = base + (long) sb * 144;
            float dd = Float.float16ToFloat(d.get(F16, p));
            float dm = Float.float16ToFloat(d.get(F16, p + 2));
            MemorySegment.copy(d, ValueLayout.JAVA_BYTE, p + 4, sc, 0, 12);
            long q = p + 16;
            int xo = sb * 256;
            int is = 0;
            for (int c = 0; c < 4; c++) {                   // vier chunks van 64
                int s1 = Dequant.scaleK4(sc, is),     m1 = Dequant.minK4(sc, is);
                int s2 = Dequant.scaleK4(sc, is + 1), m2 = Dequant.minK4(sc, is + 1);

                FloatVector accLo = FloatVector.zero(FS);
                FloatVector accHi = FloatVector.zero(FS);
                for (int helft = 0; helft < 2; helft++) {
                    ByteVector bytes = ByteVector.fromMemorySegment(BS, d, q + helft * 16L, ByteOrder.LITTLE_ENDIAN);
                    ByteVector lo = bytes.and((byte) 0x0F);
                    ByteVector hi = bytes.lanewise(VectorOperators.LSHR, 4).and((byte) 0x0F);
                    int xlo = xo + helft * 16;
                    for (int deel = 0; deel < 4; deel++) {
                        accLo = ((FloatVector) lo.castShape(FS, deel))
                                .fma(FloatVector.fromArray(FS, x, xlo + deel * 4), accLo);
                        accHi = ((FloatVector) hi.castShape(FS, deel))
                                .fma(FloatVector.fromArray(FS, x, xlo + 32 + deel * 4), accHi);
                    }
                }
                int g = xo / 16;                            // bloksom-index van dit chunk
                totaal = accLo.fma(FloatVector.broadcast(FS, dd * s1), totaal);
                totaal = accHi.fma(FloatVector.broadcast(FS, dd * s2), totaal);
                minTerm += dm * m1 * (double) (bsum16[g] + bsum16[g + 1])
                         + dm * m2 * (double) (bsum16[g + 2] + bsum16[g + 3]);
                q += 32;
                is += 2;
                xo += 64;
            }
        }
        return (float) (som4(totaal) - minTerm);
    }

    // ---------------------------------------------------------------- Q6_K: 210 bytes per 256

    static float dotQ6_K(MemorySegment d, long base, float[] x, int k, float[] bsum16) {
        FloatVector totaal = FloatVector.zero(FS);
        double corr = 0;
        for (int sb = 0; sb < k / 256; sb++) {
            long p = base + (long) sb * 210;
            float dd = Float.float16ToFloat(d.get(F16, p + 208));
            int xo = sb * 256;
            for (int helft = 0; helft < 2; helft++) {
                long ql = p + helft * 64;
                long qh = p + 128 + helft * 32;
                long sc = p + 192 + helft * 8;
                int xh = xo + helft * 128;
                for (int lb = 0; lb < 2; lb++) {            // twee groepen van 16 lanes
                    ByteVector la = ByteVector.fromMemorySegment(BS, d, ql + lb * 16L, ByteOrder.LITTLE_ENDIAN);
                    ByteVector lc = ByteVector.fromMemorySegment(BS, d, ql + 32 + lb * 16L, ByteOrder.LITTLE_ENDIAN);
                    ByteVector h  = ByteVector.fromMemorySegment(BS, d, qh + lb * 16L, ByteOrder.LITTLE_ENDIAN);

                    ByteVector q1 = la.and((byte) 0x0F).or(h.and((byte) 3).lanewise(VectorOperators.LSHL, 4));
                    ByteVector q2 = lc.and((byte) 0x0F).or(h.lanewise(VectorOperators.LSHR, 2).and((byte) 3).lanewise(VectorOperators.LSHL, 4));
                    ByteVector q3 = la.lanewise(VectorOperators.LSHR, 4).and((byte) 0x0F).or(h.lanewise(VectorOperators.LSHR, 4).and((byte) 3).lanewise(VectorOperators.LSHL, 4));
                    ByteVector q4 = lc.lanewise(VectorOperators.LSHR, 4).and((byte) 0x0F).or(h.lanewise(VectorOperators.LSHR, 6).and((byte) 3).lanewise(VectorOperators.LSHL, 4));

                    // vier kwadranten: elementen l, l+32, l+64, l+96 — elk 16 lanes = één schaalgroep
                    for (int kw = 0; kw < 4; kw++) {
                        ByteVector q = switch (kw) { case 0 -> q1; case 1 -> q2; case 2 -> q3; default -> q4; };
                        int xg = xh + 32 * kw + lb * 16;
                        float schaal = dd * d.get(ValueLayout.JAVA_BYTE, sc + 2L * kw + lb);   // ondertekend
                        FloatVector acc = FloatVector.zero(FS);
                        for (int deel = 0; deel < 4; deel++) {
                            acc = ((FloatVector) q.castShape(FS, deel))
                                    .fma(FloatVector.fromArray(FS, x, xg + deel * 4), acc);
                        }
                        totaal = acc.fma(FloatVector.broadcast(FS, schaal), totaal);
                        corr += schaal * 32.0 * bsum16[xg / 16];
                    }
                }
            }
        }
        return (float) (som4(totaal) - corr);
    }
}
