Van neuron tot taalmodel

Hoofdstuk 7

De code van RekenNet: Java en Python

Dezelfde aanpak als bij het optelnetwerk, maar nu met een grotere invoer, twee verborgen lagen en een ruimere uitvoer. We bekijken eerst de Java-versie — vooral wat er veranderde ten opzichte van AddNet — en daarna precies hetzelfde in Python.

7.1 — draaien

Zo voer je het uit

Bewaar de code als RekenNet.java. Net als eerder: pure Java, geen bibliotheken, versie 11 of nieuwer.

java RekenNet.java
7.2 — wat veranderde

Het verschil met AddNet

De structuur is herkenbaar, maar op zes punten is het netwerk meegegroeid met de zwaardere taak:

  • Invoer 10 → 14: er kwam een one-hot blok van 4 bij voor de bewerking.
  • Eén → twee verborgen lagen van elk 32 neuronen (was één laag van 16).
  • Uitvoer 9 → 21: resultaten van −4 tot 16.
  • He-initialisatie van de gewichten in plaats van kleine willekeurige waarden — diepere netwerken starten daarmee betrouwbaarder.
  • Leersnelheid 0,1 → 0,2: het grotere netwerk verdraagt — en beloont — een wat grotere stap per update.
  • Vier bewerkingen in compute(), met gehele deling en b = 0 overgeslagen.
  • Backpropagatie door twee lagen: het foutsignaal reist nu dz3 → dz2 → dz1.
7.3 — de invoer

De bewerking mee coderen

De codering krijgt er een derde blok bij. De bewerking op (0=+, 1=−, 2=×, 3=÷) zet een 1 op positie 10 tot 13.

invoercoderingRekenNet.java
static double[] encodeInput(int a, int b, int op) {
    double[] x = new double[INPUTS];     // 14 getallen
    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;
}
7.4 — de bewerkingen

Rekenen en de dataset bouwen

Eén functie rekent het juiste antwoord uit, afhankelijk van de bewerking. Let op de gehele deling (a / b op gehele getallen).

het juiste antwoordRekenNet.java
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
    }
}

Bij het bouwen van de dataset lopen we over alle bewerkingen en getallen, en slaan we de combinaties met b = 0 voor deling over.

alle 95 voorbeeldenRekenNet.java
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);
        cls[n] = r + OFFSET;
        n++;
    }
7.5 — betere start

He-initialisatie

Een diep netwerk met willekeurige startgewichten kan moeilijk op gang komen: signalen doven uit of ontploffen. He-initialisatie kiest de spreiding van de startgewichten op maat van het aantal ingangen van een laag — dat houdt de signalen mooi in balans.

gewichten initialiserenRekenNet.java
static void fill(double[][] W, double[] b, int fanIn) {
    double std = Math.sqrt(2.0 / fanIn);    // He-initialisatie
    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;
    }
}
7.6 — voorwaarts

Door twee verborgen lagen

De voorwaartse doorgang krijgt een extra trap. Laag 2 neemt de uitvoer van laag 1 (h1) als invoer; daarna volgt pas de uitvoerlaag.

twee lagen na elkaarRekenNet.java
// laag 1
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);
}
// laag 2 neemt h1 als invoer
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);
}
7.7 — achterwaarts

Het foutsignaal twee lagen terug

De backpropagatie spiegelt de voorwaartse weg, maar achterstevoren. Het foutsignaal start bij de uitvoer en reist via laag 2 naar laag 1. Bij elke laag zorgt de ReLU-poort dat alleen actieve neuronen meetellen.

terug door beide lagenRekenNet.java
// foutsignaal bij de uitvoer
for (int k = 0; k < OUTPUTS; k++) dz3[k] = y[k] - t[k];

// terug naar laag 2  (via W3 en de ReLU van laag 2)
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;
}

// terug naar laag 1  (via W2 en de ReLU van laag 1)
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;
}

Daarna sturen we, net als bij AddNet, alle gewichten bij met de gemiddelde gradient maal de leersnelheid.

7.8 — uitvoer

Wat het programma toont

Dit is de echte uitvoer. Met twee lagen van 32 zit het netwerk al rond epoch 500 op 100%, en het verlies blijft daarna dalen. Onderaan zie je per bewerking de score en een paar opgeloste sommen.

epoch    1   verlies 3.3744   juistheid   9%
epoch  500   verlies 0.0437   juistheid 100%
epoch 1000   verlies 0.0111   juistheid 100%
epoch 1500   verlies 0.0056   juistheid 100%
epoch 2000   verlies 0.0036   juistheid 100%
epoch 2500   verlies 0.0026   juistheid 100%
epoch 3000   verlies 0.0020   juistheid 100%

Juistheid op alle 95 sommen: 95/95

Per bewerking:
  +   25/25
  -   25/25
  x   25/25
  /   20/20

Enkele voorbeelden:
  2 + 3 = 5   (zekerheid  99%)
  4 - 1 = 3   (zekerheid 100%)
  3 x 2 = 6   (zekerheid 100%)
  4 / 2 = 2   (zekerheid 100%)
  0 / 1 = 0   (zekerheid 100%)
  2 + 0 = 2   (zekerheid 100%)
  1 + 3 = 4   (zekerheid 100%)
  0 / 2 = 0   (zekerheid 100%)
7.9 — alles samen

De volledige code

Het complete bestand — precies de code die bovenstaande uitvoer produceerde.

volledig programmaRekenNet.java
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;
    }
}
7.10 — hetzelfde in Python

Nu in Python

Hetzelfde diepere netwerk, dezelfde wiskunde, in pure Python. De verschillen met Java zijn van dezelfde aard als in hoofdstuk 3 (geen types, inspringing, lijsten, de modules random en math). Specifiek voor dit programma: math.sqrt voor de He-initialisatie, rng.gauss voor een normale verdeling, en a // b voor de gehele deling.

De bewerkingen en de codering met de bewerkingscode:

bewerkingen + coderingRekenNet.py
def compute(a, b, op):
    if op == 0: return a + b
    if op == 1: return a - b
    if op == 2: return a * b
    return a // b          # gehele deling

def encode_input(a, b, op):
    x = [0.0] * 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

He-initialisatie, kort en krachtig met een list comprehension:

He-initialisatieRekenNet.py
def he(rows, cols, fan_in):
    std = math.sqrt(2.0 / fan_in)     # spreiding op maat van de ingangen
    return [[rng.gauss(0.0, std) for _ in range(cols)]
            for _ in range(rows)]

De voorwaartse doorgang door beide verborgen lagen:

voorwaarts (twee lagen)RekenNet.py
def predict(x):
    h1 = [0.0] * H1
    for j in range(H1):
        s = b1[j]
        for i in range(INPUTS):
            s += W1[j][i] * x[i]
        h1[j] = s if s > 0 else 0.0      # ReLU
    h2 = [0.0] * H2
    for j in range(H2):
        s = b2[j]
        for i in range(H1):
            s += W2[j][i] * h1[i]
        h2[j] = s if s > 0 else 0.0
    z3 = [0.0] * OUTPUTS
    for k in range(OUTPUTS):
        s = b3[k]
        for j in range(H2):
            s += W3[k][j] * h2[j]
        z3[k] = s
    return softmax(z3)

En het foutsignaal twee lagen terug, net als in Java:

achterwaarts (twee lagen)RekenNet.py
# foutsignaal bij de uitvoer
dz3 = [y[k] - t[k] for k in range(OUTPUTS)]

# terug naar laag 2  (via W3 en de ReLU van laag 2)
dz2 = [0.0] * H2
for j in range(H2):
    ssum = 0.0
    for k in range(OUTPUTS):
        ssum += W3[k][j] * dz3[k]
    dz2[j] = ssum if z2[j] > 0 else 0.0

# terug naar laag 1  (via W2 en de ReLU van laag 1)
dz1 = [0.0] * H1
for j in range(H1):
    ssum = 0.0
    for k in range(H2):
        ssum += W2[k][j] * dz2[k]
    dz1[j] = ssum if z1[j] > 0 else 0.0
7.11 — draaien & uitvoer

De Python-versie uitvoeren

Bewaar als RekenNet.py en draai met Python 3 (enkel random en math). Let op: in pure Python met twee lagen duurt het trainen merkbaar langer dan in Java — reken op een halve tot enkele minuten.

python3 RekenNet.py

De echte uitvoer — ook hier 100% op alle 95 sommen:

epoch    1   verlies 3.3242   juistheid   3%
epoch  500   verlies 0.0430   juistheid 100%
epoch 1000   verlies 0.0109   juistheid 100%
epoch 1500   verlies 0.0055   juistheid 100%
epoch 2000   verlies 0.0035   juistheid 100%
epoch 2500   verlies 0.0026   juistheid 100%
epoch 3000   verlies 0.0020   juistheid 100%

Juistheid op alle 95 sommen: 95/95

Per bewerking:
  +   25/25
  -   25/25
  x   25/25
  /   20/20

Enkele voorbeelden:
  2 + 3 = 5  (zekerheid 100%)
  4 - 1 = 3  (zekerheid 100%)
  3 x 2 = 6  (zekerheid 100%)
  4 / 2 = 2  (zekerheid 100%)
  0 + 0 = 0  (zekerheid 100%)
  3 / 4 = 0  (zekerheid 100%)
  0 / 1 = 0  (zekerheid 100%)
  0 - 3 = -3 (zekerheid 100%)
7.12 — volledige Python-code

Het hele Python-programma

Het complete bestand — precies de code die de uitvoer hierboven produceerde.

volledig programmaRekenNet.py
"""
RekenNet — een dieper netwerk dat vier bewerkingen leert, in pure Python.
Zelfde opzet als RekenNet.java.

  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): resultaten -4..16   (klasse = resultaat + 4)

Deling is gehele deling; b = 0 valt weg. Geen externe bibliotheken.
Uitvoeren:  python3 RekenNet.py
"""
import random
import math

INPUTS = 14
H1 = 32
H2 = 32
OUTPUTS = 21          # resultaten -4..16
OFFSET = 4            # klasse-index = resultaat + 4
LEARNING_RATE = 0.2
EPOCHS = 3000
SEED = 42
SYM = ['+', '-', 'x', '/']

rng = random.Random(SEED)


def he(rows, cols, fan_in):
    """He-initialisatie: spreiding sqrt(2 / aantal-ingangen)."""
    std = math.sqrt(2.0 / fan_in)
    return [[rng.gauss(0.0, std) for _ in range(cols)] for _ in range(rows)]


W1 = he(H1, INPUTS, INPUTS)
b1 = [0.0] * H1
W2 = he(H2, H1, H1)
b2 = [0.0] * H2
W3 = he(OUTPUTS, H2, H2)
b3 = [0.0] * OUTPUTS


def compute(a, b, op):
    if op == 0:
        return a + b
    if op == 1:
        return a - b
    if op == 2:
        return a * b
    return a // b          # gehele deling


def encode_input(a, b, op):
    x = [0.0] * 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


def one_hot(value, length):
    v = [0.0] * length
    v[value] = 1.0
    return v


def softmax(z):
    m = max(z)
    exps = [math.exp(v - m) for v in z]
    s = sum(exps)
    return [e / s for e in exps]


def predict(x):
    h1 = [0.0] * H1
    for j in range(H1):
        s = b1[j]
        for i in range(INPUTS):
            s += W1[j][i] * x[i]
        h1[j] = s if s > 0 else 0.0
    h2 = [0.0] * H2
    for j in range(H2):
        s = b2[j]
        for i in range(H1):
            s += W2[j][i] * h1[i]
        h2[j] = s if s > 0 else 0.0
    z3 = [0.0] * OUTPUTS
    for k in range(OUTPUTS):
        s = b3[k]
        for j in range(H2):
            s += W3[k][j] * h2[j]
        z3[k] = s
    return softmax(z3)


def index_of_max(v):
    best = 0
    for i in range(1, len(v)):
        if v[i] > v[best]:
            best = i
    return best


# dataset: alle geldige (a, b, bewerking)
X, T, cls, opi = [], [], [], []
for op in range(4):
    for a in range(5):
        for b in range(5):
            if op == 3 and b == 0:
                continue              # geen deling door 0
            r = compute(a, b, op)
            X.append(encode_input(a, b, op))
            T.append(one_hot(r + OFFSET, OUTPUTS))
            cls.append(r + OFFSET)
            opi.append(op)
n = len(X)


def accuracy():
    ok = 0
    for s in range(n):
        if index_of_max(predict(X[s])) == cls[s]:
            ok += 1
    return ok / n


for epoch in range(1, EPOCHS + 1):
    gW1 = [[0.0] * INPUTS for _ in range(H1)]
    gb1 = [0.0] * H1
    gW2 = [[0.0] * H1 for _ in range(H2)]
    gb2 = [0.0] * H2
    gW3 = [[0.0] * H2 for _ in range(OUTPUTS)]
    gb3 = [0.0] * OUTPUTS
    total_loss = 0.0

    for s in range(n):
        x, t = X[s], T[s]

        # --- voorwaarts ---
        z1 = [0.0] * H1
        h1 = [0.0] * H1
        for j in range(H1):
            ssum = b1[j]
            for i in range(INPUTS):
                ssum += W1[j][i] * x[i]
            z1[j] = ssum
            h1[j] = ssum if ssum > 0 else 0.0
        z2 = [0.0] * H2
        h2 = [0.0] * H2
        for j in range(H2):
            ssum = b2[j]
            for i in range(H1):
                ssum += W2[j][i] * h1[i]
            z2[j] = ssum
            h2[j] = ssum if ssum > 0 else 0.0
        z3 = [0.0] * OUTPUTS
        for k in range(OUTPUTS):
            ssum = b3[k]
            for j in range(H2):
                ssum += W3[k][j] * h2[j]
            z3[k] = ssum
        y = softmax(z3)

        correct = index_of_max(t)
        total_loss += -math.log(max(y[correct], 1e-12))

        # --- achterwaarts ---
        dz3 = [y[k] - t[k] for k in range(OUTPUTS)]
        for k in range(OUTPUTS):
            gb3[k] += dz3[k]
            for j in range(H2):
                gW3[k][j] += dz3[k] * h2[j]

        dz2 = [0.0] * H2
        for j in range(H2):
            ssum = 0.0
            for k in range(OUTPUTS):
                ssum += W3[k][j] * dz3[k]
            dz2[j] = ssum if z2[j] > 0 else 0.0
        for j in range(H2):
            gb2[j] += dz2[j]
            for i in range(H1):
                gW2[j][i] += dz2[j] * h1[i]

        dz1 = [0.0] * H1
        for j in range(H1):
            ssum = 0.0
            for k in range(H2):
                ssum += W2[k][j] * dz2[k]
            dz1[j] = ssum if z1[j] > 0 else 0.0
        for j in range(H1):
            gb1[j] += dz1[j]
            for i in range(INPUTS):
                gW1[j][i] += dz1[j] * x[i]

    # --- bijsturen met de gemiddelde gradient ---
    scale = LEARNING_RATE / n
    for k in range(OUTPUTS):
        b3[k] -= scale * gb3[k]
        for j in range(H2):
            W3[k][j] -= scale * gW3[k][j]
    for j in range(H2):
        b2[j] -= scale * gb2[j]
        for i in range(H1):
            W2[j][i] -= scale * gW2[j][i]
    for j in range(H1):
        b1[j] -= scale * gb1[j]
        for i in range(INPUTS):
            W1[j][i] -= scale * gW1[j][i]

    if epoch == 1 or epoch % 500 == 0:
        print(f"epoch {epoch:4d}   verlies {total_loss / n:.4f}   "
              f"juistheid {100 * accuracy():3.0f}%")


print(f"\nJuistheid op alle {n} sommen: {round(accuracy() * n)}/{n}")
print("\nPer bewerking:")
for op in range(4):
    ok = tot = 0
    for s in range(n):
        if opi[s] == op:
            tot += 1
            if index_of_max(predict(X[s])) == cls[s]:
                ok += 1
    print(f"  {SYM[op]}   {ok:2d}/{tot:2d}")


def show(a, b, op):
    y = predict(encode_input(a, b, op))
    g = index_of_max(y)
    result = g - OFFSET
    flag = "" if result == compute(a, b, op) else "   <-- FOUT"
    print(f"  {a} {SYM[op]} {b} = {result:<3d}(zekerheid {100 * y[g]:3.0f}%){flag}")


print("\nEnkele voorbeelden:")
for a, b, op in [(2, 3, 0), (4, 1, 1), (3, 2, 2), (4, 2, 3)]:
    show(a, b, op)
for _ in range(4):
    op = rng.randint(0, 3)
    a = rng.randint(0, 4)
    b = rng.randint(1, 4) if op == 3 else rng.randint(0, 4)
    show(a, b, op)