import java.util.Arrays;
import java.util.Locale;
import java.util.Random;

/**
 * AddNetMiniBatch — dezelfde Optelnet als in hoofdstuk 1-4, maar getraind met
 * MINI-BATCH stochastische gradient descent in plaats van volledige-batch.
 *
 * Het enige wat verandert ten opzichte van AddNet.java zit in de trainingslus:
 * we schudden elke epoch de dataset, knippen ze in kleine batches, en sturen
 * de gewichten na ELKE batch bij (in plaats van één keer per epoch over alles).
 *
 *   Architectuur:  10 invoer  ->  16 verborgen (ReLU)  ->  9 uitvoer (softmax)
 *
 * Uitvoeren met Java 11 of nieuwer, zonder apart te compileren:
 *     java AddNetMiniBatch.java
 */
public class AddNetMiniBatch {

    static final int INPUTS  = 10;
    static final int HIDDEN  = 16;
    static final int OUTPUTS = 9;

    static final double LEARNING_RATE = 0.1;
    static final int    EPOCHS        = 400;
    static final int    BATCH_SIZE    = 5;     // <-- de kern van dit hoofdstuk: kleine batches
    static final long   SEED          = 42;

    static double[][] W1 = new double[HIDDEN][INPUTS];
    static double[]   b1 = new double[HIDDEN];
    static double[][] W2 = new double[OUTPUTS][HIDDEN];
    static double[]   b2 = new double[OUTPUTS];

    static final Random rng = new Random(SEED);

    public static void main(String[] args) {

        // 1) Volledige dataset: alle 25 combinaties.
        int n = 5 * 5;
        double[][] X    = new double[n][];
        double[][] T    = new double[n][];
        int[]      sums = new int[n];
        int idx = 0;
        for (int a = 0; a <= 4; a++) {
            for (int b = 0; b <= 4; b++) {
                X[idx]    = encodeInput(a, b);
                T[idx]    = oneHot(a + b, OUTPUTS);
                sums[idx] = a + b;
                idx++;
            }
        }

        initWeights();

        // 2) Trainen met mini-batch SGD.
        int[] order = new int[n];
        for (int i = 0; i < n; i++) order[i] = i;

        for (int epoch = 1; epoch <= EPOCHS; epoch++) {
            shuffle(order);                                      // 1) elke epoch opnieuw schudden
            double epochLoss = 0;
            for (int start = 0; start < n; start += BATCH_SIZE) {   // 2) in batches doorlopen
                int end = Math.min(start + BATCH_SIZE, n);
                int[] batch = Arrays.copyOfRange(order, start, end);
                epochLoss += trainOnBatch(batch, X, T);             // 3) na elke batch bijsturen
            }
            if (epoch == 1 || epoch % 50 == 0) {
                System.out.printf(Locale.ROOT, "epoch %4d   verlies %.4f   juistheid %3.0f%%%n",
                        epoch, epochLoss / n, 100.0 * accuracy(X, sums));
            }
        }

        int batchesPerEpoch = (n + BATCH_SIZE - 1) / BATCH_SIZE;
        System.out.printf(Locale.ROOT, "%n%d voorbeelden per batch  ->  %d bijstuurstappen per epoch "
                + "(volledige-batch zou er 1 doen)%n", BATCH_SIZE, batchesPerEpoch);

        int goed = (int) Math.round(accuracy(X, sums) * n);
        System.out.println("Resultaat op alle 25 combinaties: " + goed + "/" + n + " juist\n");

        System.out.println("Enkele willekeurige sommen:");
        for (int test = 0; test < 8; test++) {
            int a = rng.nextInt(5);
            int b = rng.nextInt(5);
            double[] y = predict(encodeInput(a, b));
            int guess = indexOfMax(y);
            System.out.printf(Locale.ROOT, "  %d + %d = %d   (zekerheid %3.0f%%)%s%n",
                    a, b, guess, 100.0 * y[guess], (guess == a + b ? "" : "   <-- FOUT"));
        }
    }

    /** Eén bijstuurstap op basis van alleen de voorbeelden in 'batch'. */
    static double trainOnBatch(int[] batch, double[][] X, double[][] T) {
        double[][] gW1 = new double[HIDDEN][INPUTS];
        double[]   gb1 = new double[HIDDEN];
        double[][] gW2 = new double[OUTPUTS][HIDDEN];
        double[]   gb2 = new double[OUTPUTS];
        double batchLoss = 0;

        for (int s : batch) {
            double[] x = X[s];
            double[] t = T[s];

            // --- voorwaarts ---
            double[] z1 = new double[HIDDEN];
            double[] h  = new double[HIDDEN];
            for (int j = 0; j < HIDDEN; j++) {
                double sum = b1[j];
                for (int i = 0; i < INPUTS; i++) sum += W1[j][i] * x[i];
                z1[j] = sum;
                h[j]  = Math.max(0, sum);
            }
            double[] z2 = new double[OUTPUTS];
            for (int k = 0; k < OUTPUTS; k++) {
                double sum = b2[k];
                for (int j = 0; j < HIDDEN; j++) sum += W2[k][j] * h[j];
                z2[k] = sum;
            }
            double[] y = softmax(z2);

            batchLoss += -Math.log(Math.max(y[indexOfMax(t)], 1e-12));

            // --- achterwaarts ---
            double[] dz2 = new double[OUTPUTS];
            for (int k = 0; k < OUTPUTS; k++) dz2[k] = y[k] - t[k];
            for (int k = 0; k < OUTPUTS; k++) {
                gb2[k] += dz2[k];
                for (int j = 0; j < HIDDEN; j++) gW2[k][j] += dz2[k] * h[j];
            }
            double[] dh = new double[HIDDEN];
            for (int j = 0; j < HIDDEN; j++) {
                double sum = 0;
                for (int k = 0; k < OUTPUTS; k++) sum += W2[k][j] * dz2[k];
                dh[j] = sum;
            }
            double[] dz1 = new double[HIDDEN];
            for (int j = 0; j < HIDDEN; j++) dz1[j] = (z1[j] > 0) ? dh[j] : 0;
            for (int j = 0; j < HIDDEN; j++) {
                gb1[j] += dz1[j];
                for (int i = 0; i < INPUTS; i++) gW1[j][i] += dz1[j] * x[i];
            }
        }

        // bijsturen met de GEMIDDELDE gradient over deze batch
        double scale = LEARNING_RATE / batch.length;
        for (int k = 0; k < OUTPUTS; k++) {
            b2[k] -= scale * gb2[k];
            for (int j = 0; j < HIDDEN; j++) W2[k][j] -= scale * gW2[k][j];
        }
        for (int j = 0; j < HIDDEN; j++) {
            b1[j] -= scale * gb1[j];
            for (int i = 0; i < INPUTS; i++) W1[j][i] -= scale * gW1[j][i];
        }
        return batchLoss;
    }

    /** Fisher-Yates: schud de volgorde van de indices door elkaar. */
    static void shuffle(int[] a) {
        for (int i = a.length - 1; i > 0; i--) {
            int j = rng.nextInt(i + 1);
            int tmp = a[i]; a[i] = a[j]; a[j] = tmp;
        }
    }

    // ---------- hulpfuncties (identiek aan AddNet.java) ----------

    static double[] encodeInput(int a, int b) {
        double[] x = new double[INPUTS];
        x[a]     = 1.0;
        x[5 + b] = 1.0;
        return x;
    }

    static double[] oneHot(int value, int len) {
        double[] v = new double[len];
        v[value] = 1.0;
        return v;
    }

    static void initWeights() {
        for (int j = 0; j < HIDDEN; j++) {
            for (int i = 0; i < INPUTS; i++) W1[j][i] = small();
            b1[j] = 0;
        }
        for (int k = 0; k < OUTPUTS; k++) {
            for (int j = 0; j < HIDDEN; j++) W2[k][j] = small();
            b2[k] = 0;
        }
    }

    static double small() {
        return (rng.nextDouble() - 0.5) * 0.5;
    }

    static double[] softmax(double[] z) {
        double max = Double.NEGATIVE_INFINITY;
        for (double v : z) max = Math.max(max, v);
        double sum = 0;
        double[] out = new double[z.length];
        for (int k = 0; k < z.length; k++) {
            out[k] = Math.exp(z[k] - max);
            sum += out[k];
        }
        for (int k = 0; k < z.length; k++) out[k] /= sum;
        return out;
    }

    static double[] predict(double[] x) {
        double[] h = new double[HIDDEN];
        for (int j = 0; j < HIDDEN; j++) {
            double sum = b1[j];
            for (int i = 0; i < INPUTS; i++) sum += W1[j][i] * x[i];
            h[j] = Math.max(0, sum);
        }
        double[] z2 = new double[OUTPUTS];
        for (int k = 0; k < OUTPUTS; k++) {
            double sum = b2[k];
            for (int j = 0; j < HIDDEN; j++) sum += W2[k][j] * h[j];
            z2[k] = sum;
        }
        return softmax(z2);
    }

    static int indexOfMax(double[] v) {
        int best = 0;
        for (int i = 1; i < v.length; i++) if (v[i] > v[best]) best = i;
        return best;
    }

    static double accuracy(double[][] X, int[] sums) {
        int ok = 0;
        for (int s = 0; s < X.length; s++) {
            if (indexOfMax(predict(X[s])) == sums[s]) ok++;
        }
        return (double) ok / X.length;
    }
}
