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

/**
 * RekenNet — een dieper neuraal netwerk dat vier bewerkingen leert:
 * optellen (+), aftrekken (-), vermenigvuldigen (x) en gehele deling (/),
 * voor twee getallen a, b uit {0,1,2,3,4}. Pure Java, geen bibliotheken.
 *
 *   Invoer (14):  one-hot a(5) + one-hot b(5) + one-hot bewerking(4)
 *   Verborgen 1:  32 neuronen (ReLU)
 *   Verborgen 2:  32 neuronen (ReLU)
 *   Uitvoer (21): mogelijke resultaten -4..16   (klasse = resultaat + 4)
 *
 * Deling is GEHELE deling (zoals Java's / op int), en b = 0 laten we weg
 * (delen door nul is ongedefinieerd).
 *
 * Uitvoeren:  java RekenNet.java     (of: javac RekenNet.java && java RekenNet)
 */
public class RekenNet {

    static final int INPUTS  = 14;
    static final int H1      = 32;
    static final int H2      = 32;
    static final int OUTPUTS = 21;   // resultaten -4..16
    static final int OFFSET  = 4;    // klasse-index = resultaat + 4

    static final double LEARNING_RATE = 0.2;
    static final int    EPOCHS        = 3000;
    static final long   SEED          = 42;

    static final char[] SYM = {'+', '-', 'x', '/'};

    // gewichten en biassen van de drie lagen
    static double[][] W1 = new double[H1][INPUTS];
    static double[]   b1 = new double[H1];
    static double[][] W2 = new double[H2][H1];
    static double[]   b2 = new double[H2];
    static double[][] W3 = new double[OUTPUTS][H2];
    static double[]   b3 = new double[OUTPUTS];

    static final Random rng = new Random(SEED);

    public static void main(String[] args) {

        // 1) Bouw de dataset: alle geldige (a, b, bewerking)-combinaties.
        int cap = 4 * 5 * 5;
        double[][] X   = new double[cap][];
        double[][] T   = new double[cap][];
        int[] aOf = new int[cap], bOf = new int[cap], opOf = new int[cap], cls = new int[cap];
        int n = 0;
        for (int op = 0; op < 4; op++) {
            for (int a = 0; a < 5; a++) {
                for (int b = 0; b < 5; b++) {
                    if (op == 3 && b == 0) continue;        // geen deling door 0
                    int r = compute(a, b, op);
                    X[n]   = encodeInput(a, b, op);
                    T[n]   = oneHot(r + OFFSET, OUTPUTS);
                    aOf[n] = a; bOf[n] = b; opOf[n] = op; cls[n] = r + OFFSET;
                    n++;
                }
            }
        }

        // 2) He-initialisatie (geschikt voor ReLU en diepere netwerken).
        initWeights();

        // 3) Trainen met volledige-batch gradient descent.
        for (int epoch = 1; epoch <= EPOCHS; epoch++) {
            double[][] gW1 = new double[H1][INPUTS]; double[] gb1 = new double[H1];
            double[][] gW2 = new double[H2][H1];     double[] gb2 = new double[H2];
            double[][] gW3 = new double[OUTPUTS][H2];double[] gb3 = new double[OUTPUTS];
            double totalLoss = 0;

            for (int s = 0; s < n; s++) {
                double[] x = X[s], t = T[s];

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

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

                // --- achterwaarts ---
                double[] dz3 = new double[OUTPUTS];
                for (int k = 0; k < OUTPUTS; k++) dz3[k] = y[k] - t[k];
                for (int k = 0; k < OUTPUTS; k++) {
                    gb3[k] += dz3[k];
                    for (int j = 0; j < H2; j++) gW3[k][j] += dz3[k] * h2[j];
                }

                double[] dz2 = new double[H2];
                for (int j = 0; j < H2; j++) {
                    double sum = 0;
                    for (int k = 0; k < OUTPUTS; k++) sum += W3[k][j] * dz3[k];
                    dz2[j] = (z2[j] > 0) ? sum : 0;          // door ReLU van laag 2
                }
                for (int j = 0; j < H2; j++) {
                    gb2[j] += dz2[j];
                    for (int i = 0; i < H1; i++) gW2[j][i] += dz2[j] * h1[i];
                }

                double[] dz1 = new double[H1];
                for (int j = 0; j < H1; j++) {
                    double sum = 0;
                    for (int k = 0; k < H2; k++) sum += W2[k][j] * dz2[k];
                    dz1[j] = (z1[j] > 0) ? sum : 0;          // door ReLU van laag 1
                }
                for (int j = 0; j < H1; j++) {
                    gb1[j] += dz1[j];
                    for (int i = 0; i < INPUTS; i++) gW1[j][i] += dz1[j] * x[i];
                }
            }

            // --- bijsturen met de gemiddelde gradient ---
            double sc = LEARNING_RATE / n;
            for (int k = 0; k < OUTPUTS; k++) { b3[k] -= sc*gb3[k]; for (int j = 0; j < H2; j++) W3[k][j] -= sc*gW3[k][j]; }
            for (int j = 0; j < H2; j++)      { b2[j] -= sc*gb2[j]; for (int i = 0; i < H1; i++) W2[j][i] -= sc*gW2[j][i]; }
            for (int j = 0; j < H1; j++)      { b1[j] -= sc*gb1[j]; for (int i = 0; i < INPUTS; i++) W1[j][i] -= sc*gW1[j][i]; }

            if (epoch == 1 || epoch % 500 == 0) {
                System.out.printf(Locale.ROOT, "epoch %4d   verlies %.4f   juistheid %3.0f%%%n",
                        epoch, totalLoss / n, 100.0 * accuracy(X, cls, n));
            }
        }

        // 4) Resultaat: totaal en per bewerking.
        System.out.printf(Locale.ROOT, "%nJuistheid op alle %d sommen: %d/%d%n", n, (int)Math.round(accuracy(X, cls, n)*n), n);
        System.out.println("\nPer bewerking:");
        for (int op = 0; op < 4; op++) {
            int ok = 0, tot = 0;
            for (int s = 0; s < n; s++) if (opOf[s] == op) {
                tot++;
                if (indexOfMax(predict(X[s])) == cls[s]) ok++;
            }
            System.out.printf(Locale.ROOT, "  %c   %2d/%2d%n", SYM[op], ok, tot);
        }

        // 5) Een paar voorbeelden, één per bewerking plus enkele willekeurige.
        System.out.println("\nEnkele voorbeelden:");
        int[][] demo = { {2,3,0}, {4,1,1}, {3,2,2}, {4,2,3} };
        for (int[] d : demo) showExample(d[0], d[1], d[2]);
        for (int t = 0; t < 4; t++) {
            int op = rng.nextInt(4);
            int a = rng.nextInt(5);
            int b = (op == 3) ? 1 + rng.nextInt(4) : rng.nextInt(5);
            showExample(a, b, op);
        }
    }

    // ---------- hulpfuncties ----------

    static int compute(int a, int b, int op) {
        switch (op) {
            case 0: return a + b;
            case 1: return a - b;
            case 2: return a * b;
            default: return a / b;     // gehele deling
        }
    }

    static void showExample(int a, int b, int op) {
        double[] y = predict(encodeInput(a, b, op));
        int g = indexOfMax(y);
        int result = g - OFFSET;
        boolean ok = (result == compute(a, b, op));
        System.out.printf(Locale.ROOT, "  %d %c %d = %-3d (zekerheid %3.0f%%)%s%n",
                a, SYM[op], b, result, 100.0 * y[g], ok ? "" : "   <-- FOUT");
    }

    /** one-hot a(5) + one-hot b(5) + one-hot bewerking(4) = 14 getallen. */
    static double[] encodeInput(int a, int b, int op) {
        double[] x = new double[INPUTS];
        x[a]      = 1.0;   // getal a       -> posities 0..4
        x[5 + b]  = 1.0;   // getal b       -> posities 5..9
        x[10 + op]= 1.0;   // bewerking     -> posities 10..13
        return x;
    }

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

    /** He-initialisatie: standaarddeviatie sqrt(2 / aantal-ingangen). */
    static void initWeights() {
        fill(W1, b1, INPUTS);
        fill(W2, b2, H1);
        fill(W3, b3, H2);
    }
    static void fill(double[][] W, double[] b, int fanIn) {
        double std = Math.sqrt(2.0 / fanIn);
        for (int r = 0; r < W.length; r++) {
            for (int c = 0; c < W[r].length; c++) W[r][c] = rng.nextGaussian() * std;
            b[r] = 0;
        }
    }

    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;
    }

    /** Voorwaartse doorgang door beide verborgen lagen -> kansen. */
    static double[] predict(double[] x) {
        double[] h1 = new double[H1];
        for (int j = 0; j < H1; j++) {
            double sum = b1[j];
            for (int i = 0; i < INPUTS; i++) sum += W1[j][i] * x[i];
            h1[j] = Math.max(0, sum);
        }
        double[] h2 = new double[H2];
        for (int j = 0; j < H2; j++) {
            double sum = b2[j];
            for (int i = 0; i < H1; i++) sum += W2[j][i] * h1[i];
            h2[j] = Math.max(0, sum);
        }
        double[] z3 = new double[OUTPUTS];
        for (int k = 0; k < OUTPUTS; k++) {
            double sum = b3[k];
            for (int j = 0; j < H2; j++) sum += W3[k][j] * h2[j];
            z3[k] = sum;
        }
        return softmax(z3);
    }

    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[] cls, int n) {
        int ok = 0;
        for (int s = 0; s < n; s++) if (indexOfMax(predict(X[s])) == cls[s]) ok++;
        return (double) ok / n;
    }
}
