import java.util.random.RandomGenerator;

/**
 * De sampler &mdash; mijlpaal M6/M7. Kiest het volgende token uit de logits.
 *
 * <p>Tot nu toe koos het project altijd het hoogste hokje (greedy, temperatuur 0):
 * perfect om te bew&iacute;jzen dat de forward pass klopt, maar dodelijk voor een
 * gesprek &mdash; hetzelfde begin geeft altijd hetzelfde vervolg, en bij twijfel blijft
 * het model in kringetjes draaien. Een echt gesprek vraagt gecontroleerd gokken:
 * <ol>
 *   <li><b>temperatuur</b> &mdash; deel de logits door T v&oacute;&oacute;r de softmax.
 *       T &lt; 1 maakt de verdeling scherper, T &gt; 1 vlakker, T &rarr; 0 is greedy;</li>
 *   <li><b>top-k</b> &mdash; kijk alleen naar de k beste kandidaten;</li>
 *   <li><b>top-p</b> &mdash; knip daarbinnen de staart af: houd de kleinste verzameling
 *       die samen minstens kans p draagt. Zonder afknippen dragen de ~150&nbsp;000
 *       onzinkandidaten samen alsnog een paar procent kans, en &eacute;&eacute;n gek token
 *       kan een gesprek laten ontsporen.</li>
 * </ol>
 *
 * <p>De standaardwaarden (0,8 / 40 / 0,9) zijn dezelfde als die van ollama, zodat een
 * gesprek hier "hetzelfde soort toeval" heeft als daar. De toevalsbron is gezaaid:
 * met dezelfde seed is elke run exact herhaalbaar &mdash; determinisme blijft ook hier
 * het gereedschap.
 */
final class Sampler {

    static final float STANDAARD_TEMPERATUUR = 0.8f;
    static final int   STANDAARD_TOP_K = 40;
    static final float STANDAARD_TOP_P = 0.9f;

    final float temperatuur;
    final int topK;
    final float topP;
    final long seed;
    private final RandomGenerator rnd;

    Sampler(float temperatuur, int topK, float topP, long seed) {
        // genegeerde predicaten in plaats van omgekeerde vergelijkingen: NaN faalt elke
        // vergelijking, dus "< 0" zou NaN stil doorlaten — "!(>= 0)" weigert hem wél
        if (!(temperatuur >= 0)) throw new IllegalArgumentException("temperatuur mag niet negatief of NaN zijn, kreeg " + temperatuur);
        if (!(topP > 0 && topP <= 1)) throw new IllegalArgumentException("top-p hoort in (0, 1], niet " + topP);
        if (temperatuur > 0 && topP < 1 && topK < 1) {
            throw new IllegalArgumentException("bij top-p < 1 hoort een top-k van minstens 1 "
                    + "(de kandidaten moeten gesorteerd worden; standaard is 40)");
        }
        this.temperatuur = temperatuur;
        this.topK = topK;
        this.topP = topP;
        this.seed = seed;
        this.rnd = new java.util.Random(seed);
    }

    /** Greedy sampler: altijd het hoogste hokje, geen toeval. */
    static Sampler greedy() {
        return new Sampler(0f, 1, 1f, 0);
    }

    /** Kiest het volgende token. Bij temperatuur 0 is dit exact argmax. */
    int kies(float[] logits) {
        if (temperatuur == 0) return argmax(logits);
        if (topK < 1 && topP >= 1) return trekVolledig(logits);

        // top-k selectie: de k hoogste logits, aflopend gesorteerd (k is klein)
        int k = Math.min(topK, logits.length);
        int[] idx = new int[k];
        float[] val = new float[k];
        int n = 0;
        for (int i = 0; i < logits.length; i++) {
            float v = logits[i];
            if (n == k && v <= val[n - 1]) continue;
            int j = Math.min(n, k - 1);
            while (j > 0 && val[j - 1] < v) {
                val[j] = val[j - 1];
                idx[j] = idx[j - 1];
                j--;
            }
            val[j] = v;
            idx[j] = i;
            if (n < k) n++;
        }

        // temperatuur + softmax over de kandidaten (de volgorde verandert daar niet door)
        double[] p = new double[n];
        double som = 0;
        for (int j = 0; j < n; j++) {
            p[j] = Math.exp((val[j] - val[0]) / temperatuur);
            som += p[j];
        }

        // top-p: de kleinste kopgroep die samen minstens kans topP draagt
        int m = n;
        if (topP < 1) {
            double cum = 0;
            for (int j = 0; j < n; j++) {
                cum += p[j];
                if (cum >= topP * som) { m = j + 1; break; }
            }
        }

        // trekken uit de (impliciet hernormaliseerde) kopgroep
        double binnen = 0;
        for (int j = 0; j < m; j++) binnen += p[j];
        double doel = rnd.nextDouble() * binnen;
        double cum = 0;
        for (int j = 0; j < m; j++) {
            cum += p[j];
            if (doel <= cum) return idx[j];
        }
        return idx[m - 1];
    }

    /** Pure temperatuur-sampling over de hele woordenschat (top-k en top-p uit). */
    private int trekVolledig(float[] logits) {
        float max = logits[argmax(logits)];
        double som = 0;
        double[] p = new double[logits.length];
        for (int i = 0; i < logits.length; i++) {
            p[i] = Math.exp((logits[i] - max) / temperatuur);
            som += p[i];
        }
        double doel = rnd.nextDouble() * som;
        double cum = 0;
        for (int i = 0; i < logits.length; i++) {
            cum += p[i];
            if (doel <= cum) return i;
        }
        return logits.length - 1;
    }

    static int argmax(float[] a) {
        int best = 0;
        for (int i = 1; i < a.length; i++) if (a[i] > a[best]) best = i;
        return best;
    }
}
