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.VectorOperators;
import jdk.incubator.vector.VectorSpecies;

/**
 * De enige kernel die er echt toe doet in een LLM-decoder:
 * gekwantiseerd matrix-vector product (GEMV) met dequantisatie ter plaatse.
 *
 * Gewichtsformaat = Q8_0 van GGUF: blok van 32 gewichten =
 *   2 bytes fp16 schaal (d) + 32 bytes int8 quants  ->  34 bytes.
 */
final class Kernels {

    static final int QK = 32;              // blokgrootte
    static final int BLOCK_BYTES = 34;     // 2 (fp16 d) + 32 (int8 qs)

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

    private Kernels() {}

    /** Scalair: geen SIMD, puur java. Referentie + ondergrens. */
    static float dotRowScalar(MemorySegment w, long rowBase, float[] x, int k) {
        int nb = k / QK;
        float total = 0f;
        for (int b = 0; b < nb; b++) {
            long bb = rowBase + (long) b * BLOCK_BYTES;
            float d = Float.float16ToFloat(w.get(ValueLayout.JAVA_SHORT_UNALIGNED, bb));
            float acc = 0f;
            int xo = b * QK;
            for (int j = 0; j < QK; j++) {
                acc += w.get(ValueLayout.JAVA_BYTE, bb + 2 + j) * x[xo + j];
            }
            total += d * acc;
        }
        return total;
    }

    /** Vector API: byte-lanes uit de mmap, castShape naar float, fma. */
    static float dotRowSimd(MemorySegment w, long rowBase, float[] x, int k) {
        int nb = k / QK;
        FloatVector total = FloatVector.zero(FS);
        for (int b = 0; b < nb; b++) {
            long bb = rowBase + (long) b * BLOCK_BYTES;
            float d = Float.float16ToFloat(w.get(ValueLayout.JAVA_SHORT_UNALIGNED, bb));
            FloatVector acc = FloatVector.zero(FS);
            int xo = b * QK;
            for (int half = 0; half < 2; half++) {
                ByteVector bv = ByteVector.fromMemorySegment(
                        BS, w, bb + 2 + half * 16L, ByteOrder.LITTLE_ENDIAN);
                for (int part = 0; part < 4; part++) {
                    FloatVector wv = (FloatVector) bv.castShape(FS, part);
                    FloatVector xv = FloatVector.fromArray(FS, x, xo + half * 16 + part * 4);
                    acc = wv.fma(xv, acc);
                }
            }
            total = acc.fma(FloatVector.broadcast(FS, d), total);
        }
        return total.reduceLanes(VectorOperators.ADD);
    }

    // ---------------- Q4_0: 2 bytes fp16 schaal + 16 bytes gepakte nibbles = 18 bytes ----------------

    static final int Q4_BLOCK_BYTES = 18;

    static float dotRowQ4Scalar(MemorySegment w, long rowBase, float[] x, int k) {
        int nb = k / QK;
        float total = 0f;
        for (int b = 0; b < nb; b++) {
            long bb = rowBase + (long) b * Q4_BLOCK_BYTES;
            float d = Float.float16ToFloat(w.get(ValueLayout.JAVA_SHORT_UNALIGNED, bb));
            float acc = 0f;
            int xo = b * QK;
            for (int j = 0; j < 16; j++) {
                int q = w.get(ValueLayout.JAVA_BYTE, bb + 2 + j) & 0xFF;
                acc += ((q & 0x0F) - 8) * x[xo + j];          // lage nibble -> gewicht j
                acc += ((q >>> 4) - 8) * x[xo + j + 16];      // hoge nibble -> gewicht j+16
            }
            total += d * acc;
        }
        return total;
    }

    static float dotRowQ4Simd(MemorySegment w, long rowBase, float[] x, int k) {
        int nb = k / QK;
        FloatVector total = FloatVector.zero(FS);
        for (int b = 0; b < nb; b++) {
            long bb = rowBase + (long) b * Q4_BLOCK_BYTES;
            float d = Float.float16ToFloat(w.get(ValueLayout.JAVA_SHORT_UNALIGNED, bb));
            ByteVector packed = ByteVector.fromMemorySegment(BS, w, bb + 2, ByteOrder.LITTLE_ENDIAN);
            ByteVector lo = packed.and((byte) 0x0F).sub((byte) 8);
            ByteVector hi = packed.lanewise(VectorOperators.LSHR, 4).and((byte) 0x0F).sub((byte) 8);
            FloatVector acc = FloatVector.zero(FS);
            int xo = b * QK;
            for (int part = 0; part < 4; part++) {
                acc = ((FloatVector) lo.castShape(FS, part))
                        .fma(FloatVector.fromArray(FS, x, xo + part * 4), acc);
                acc = ((FloatVector) hi.castShape(FS, part))
                        .fma(FloatVector.fromArray(FS, x, xo + 16 + part * 4), acc);
            }
            total = acc.fma(FloatVector.broadcast(FS, d), total);
        }
        return total.reduceLanes(VectorOperators.ADD);
    }

    static long rowBytes(int k) {
        return (long) (k / QK) * BLOCK_BYTES;
    }

    static long rowBytesQ4(int k) {
        return (long) (k / QK) * Q4_BLOCK_BYTES;
    }
}
