import java.lang.foreign.Arena;
import java.lang.foreign.MemorySegment;
import java.lang.foreign.ValueLayout;
import java.util.ArrayList;
import java.util.List;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
import java.util.concurrent.Future;
import java.util.random.RandomGenerator;

import jdk.incubator.vector.FloatVector;

/**
 * Roofline-meting voor lokale LLM-inferentie in pure Java.
 * Geen Maven, geen dependencies:
 *   java --add-modules jdk.incubator.vector Roofline.java
 * Kernels.java wordt door de multi-file source launcher (JEP 458) meegecompileerd.
 */
public class Roofline {

    static final int K = 2048;                    // hidden dim, ~Llama-3.2-1B
    static final long TOTAL_BYTES = 512L << 20;   // 512 MiB "gewichten" per sweep

    interface RowKernel { float dot(MemorySegment w, long rowBase, float[] x, int k); }

    public static void main(String[] args) {
        System.out.println("== omgeving ==");
        System.out.println("java.version      : " + System.getProperty("java.version"));
        System.out.println("os/arch           : " + System.getProperty("os.name") + " / "
                + System.getProperty("os.arch"));
        System.out.println("beschikbare cores : " + Runtime.getRuntime().availableProcessors());
        System.out.println("SPECIES_PREFERRED : " + FloatVector.SPECIES_PREFERRED.vectorBitSize()
                + " bit, " + FloatVector.SPECIES_PREFERRED.length() + " float-lanes");

        try (Arena arena = Arena.ofShared()) {
            float[] x = new float[K];
            RandomGenerator rnd = RandomGenerator.getDefault();
            for (int i = 0; i < K; i++) x[i] = rnd.nextFloat() - 0.5f;

            run(arena, rnd, x, "Q8_0", Kernels.rowBytes(K), Kernels.BLOCK_BYTES,
                    Kernels::dotRowScalar, Kernels::dotRowSimd, 34.0 / 32.0);
            run(arena, rnd, x, "Q4_0", Kernels.rowBytesQ4(K), Kernels.Q4_BLOCK_BYTES,
                    Kernels::dotRowQ4Scalar, Kernels::dotRowQ4Simd, 18.0 / 32.0);
        }
    }

    static void run(Arena arena, RandomGenerator rnd, float[] x, String naam,
                    long rowBytes, int blockBytes, RowKernel scalarK, RowKernel simdK,
                    double bytesPerParam) {

        int rows = (int) (TOTAL_BYTES / rowBytes);
        long bytes = rows * rowBytes;
        MemorySegment w = arena.allocate(bytes, 64);

        byte[] chunk = new byte[1 << 16];
        for (long off = 0; off < bytes; off += chunk.length) {
            rnd.nextBytes(chunk);
            MemorySegment.copy(chunk, 0, w, ValueLayout.JAVA_BYTE, off,
                    (int) Math.min(chunk.length, bytes - off));
        }
        for (long b = 0, n = bytes / blockBytes; b < n; b++) {   // geldige fp16-schalen
            w.set(ValueLayout.JAVA_SHORT_UNALIGNED, b * blockBytes,
                    Float.floatToFloat16(0.005f + rnd.nextFloat() * 0.02f));
        }
        float[] out = new float[rows];

        System.out.printf("%n=========== %s ===========%n", naam);
        System.out.printf("K=%d, rij=%d bytes, rijen=%d, sweep=%.1f MiB, %.3f byte/gewicht%n",
                K, rowBytes, rows, bytes / 1048576.0, bytesPerParam);

        double maxRel = 0;
        for (int r = 0; r < 64; r++) {
            float a = scalarK.dot(w, r * rowBytes, x, K);
            float b = simdK.dot(w, r * rowBytes, x, K);
            maxRel = Math.max(maxRel, Math.abs(a - b) / Math.max(1e-6, Math.abs(a)));
        }
        System.out.printf("SIMD vs scalair   : max. rel. drift %.2e %s%n",
                maxRel, maxRel < 1e-3 ? "(OK - enkel sommatievolgorde)" : "(FOUT!)");

        double best = 0; int bestT = 1;
        double sc = bestOf(3, 2, bytes, () -> sweep(scalarK, w, rowBytes, x, out, 0, rows));
        System.out.printf("%-30s %7.1f GB/s%n", "scalair (1 thread)", sc);

        for (int t : new int[]{1, 2, 4, 6, 8}) {
            double r;
            if (t == 1) {
                r = bestOf(5, 3, bytes, () -> sweep(simdK, w, rowBytes, x, out, 0, rows));
            } else {
                ExecutorService pool = Executors.newFixedThreadPool(t);
                final int tt = t;
                r = bestOf(5, 3, bytes, () -> parallel(pool, tt, simdK, w, rowBytes, rows, x, out));
                pool.shutdown();
            }
            System.out.printf("%-30s %7.1f GB/s%s%n", "Vector API (" + t + " thread"
                    + (t == 1 ? "" : "s") + ")", r, t == 1 ? String.format("   (%.1fx t.o.v. scalair)", r / sc) : "");
            if (r > best) { best = r; bestT = t; }
        }

        System.out.printf("-> beste: %.1f GB/s met %d threads   [checksum %.4f]%n", best, bestT, out[rows / 2]);
        System.out.printf("%-24s %10s %10s%n", "model", "gewichten", "tokens/s");
        long[] params = {135_000_000L, 500_000_000L, 1_000_000_000L, 3_000_000_000L, 8_000_000_000L};
        String[] namen = {"SmolLM2-135M", "Qwen2.5-0.5B", "Llama-3.2-1B", "Llama-3.2-3B", "Llama-3.1-8B"};
        for (int i = 0; i < params.length; i++) {
            double gb = params[i] * bytesPerParam / 1e9;
            System.out.printf("%-24s %7.2f GB %9.1f%n", namen[i], gb, best / gb);
        }
    }

    interface Sweep { void run(); }

    static double bestOf(int n, int warmup, long bytes, Sweep s) {
        double best = 0;
        for (int i = 0; i < n + warmup; i++) {
            long t0 = System.nanoTime();
            s.run();
            long dt = System.nanoTime() - t0;
            if (i >= warmup) best = Math.max(best, bytes / (dt / 1e9) / 1e9);
        }
        return best;
    }

    static void sweep(RowKernel kern, MemorySegment w, long rowBytes, float[] x, float[] out,
                      int from, int to) {
        for (int r = from; r < to; r++) out[r] = kern.dot(w, r * rowBytes, x, K);
    }

    static void parallel(ExecutorService pool, int threads, RowKernel kern, MemorySegment w,
                         long rowBytes, int rows, float[] x, float[] out) {
        List<Future<?>> fs = new ArrayList<>(threads);
        int chunk = (rows + threads - 1) / threads;
        for (int t = 0; t < threads; t++) {
            int from = t * chunk, to = Math.min(rows, from + chunk);
            if (from >= to) break;
            fs.add(pool.submit(() -> sweep(kern, w, rowBytes, x, out, from, to)));
        }
        try { for (Future<?> f : fs) f.get(); }
        catch (Exception e) { throw new RuntimeException(e); }
    }
}
