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

/**
 * Optelnet — een klein neuraal netwerk dat leert om twee getallen (0..4) op te tellen.
 * Geschreven in pure Java, zonder externe bibliotheken.
 *
 *   Architectuur:  10 invoer  ->  16 verborgen (ReLU)  ->  9 uitvoer (softmax)
 *                  (2x one-hot)                            (sommen 0..8)
 *
 * Uitvoeren met Java 11 of nieuwer, zonder apart te compileren:
 *     java AddNet.java
 *
 * Of klassiek compileren en draaien:
 *     javac AddNet.java
 *     java AddNet
 */
public class AddNet {

    // ---- Vorm van het netwerk ----
    static final int INPUTS  = 10;   // 2 getallen, elk one-hot over {0,1,2,3,4}
    static final int HIDDEN  = 16;   // grootte van de verborgen laag
    static final int OUTPUTS = 9;    // mogelijke sommen: 0..8

    // ---- Leerinstellingen ----
    static final double LEARNING_RATE = 0.1;
    static final int    EPOCHS        = 3000;
    static final long   SEED          = 42;   // vaste seed = telkens hetzelfde resultaat

    // ---- Gewichten en biassen ----
    static double[][] W1 = new double[HIDDEN][INPUTS];    // verborgen laag
    static double[]   b1 = new double[HIDDEN];
    static double[][] W2 = new double[OUTPUTS][HIDDEN];   // uitvoerlaag
    static double[]   b2 = new double[OUTPUTS];

    static final Random rng = new Random(SEED);

    public static void main(String[] args) {

        // 1) Bouw de volledige dataset: alle 25 combinaties van a en b.
        int n = 5 * 5;
        double[][] X    = new double[n][];   // invoer (one-hot, lengte 10)
        double[][] T    = new double[n][];   // gewenste uitvoer (one-hot van de som, lengte 9)
        int[]      sums = new int[n];        // de juiste som, voor controle
        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++;
            }
        }

        // 2) Geef de gewichten kleine willekeurige startwaarden.
        initWeights();

        // 3) Trainen: per epoch alle voorbeelden bekijken, gradiënten optellen, gewichten bijstellen.
        for (int epoch = 1; epoch <= EPOCHS; epoch++) {

            // accumulatoren voor de gradiënten over alle voorbeelden
            double[][] gW1 = new double[HIDDEN][INPUTS];
            double[]   gb1 = new double[HIDDEN];
            double[][] gW2 = new double[OUTPUTS][HIDDEN];
            double[]   gb2 = new double[OUTPUTS];
            double totalLoss = 0;

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

                // --- voorwaartse doorgang ---
                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);                 // ReLU
                }
                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);                     // kansen over de 9 sommen

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

                // --- achterwaartse doorgang (gradiënten) ---
                // uitvoer: dz2 = y - t  (de mooie eigenschap van softmax + kruisentropie)
                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];
                }

                // terug naar de verborgen laag
                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;         // afgeleide van ReLU
                }
                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];
                }
            }

            // --- gewichten bijstellen met de gemiddelde gradient ---
            double scale = LEARNING_RATE / n;
            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];
            }

            // --- af en toe de voortgang tonen ---
            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, sums));
            }
        }

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

        // 5) Een paar willekeurige sommen laten oplossen.
        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"));
        }
    }

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

    /** Codeer twee getallen (0..4) als een invoervector van lengte 10 (twee one-hots na elkaar). */
    static double[] encodeInput(int a, int b) {
        double[] x = new double[INPUTS];
        x[a]     = 1.0;    // eerste getal  -> posities 0..4
        x[5 + b] = 1.0;    // tweede getal  -> posities 5..9
        return x;
    }

    /** Maak een one-hot vector van lengte len met een 1 op positie value. */
    static double[] oneHot(int value, int len) {
        double[] v = new double[len];
        v[value] = 1.0;
        return v;
    }

    /** Geef alle gewichten kleine willekeurige startwaarden, biassen op 0. */
    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;
        }
    }

    /** Klein willekeurig getal rond 0 (tussen -0.25 en 0.25). */
    static double small() {
        return (rng.nextDouble() - 0.5) * 0.5;
    }

    /** Zet 9 scores om in 9 kansen die samen 1 zijn. */
    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);   // -max voor numerieke stabiliteit
            sum += out[k];
        }
        for (int k = 0; k < z.length; k++) out[k] /= sum;
        return out;
    }

    /** Volledige voorwaartse doorgang: invoer -> kansen over de 9 sommen. */
    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);
    }

    /** Positie van de grootste waarde in een vector. */
    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;
    }

    /** Aandeel juiste voorspellingen over de gegeven voorbeelden. */
    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;
    }
}
