Van neuron tot taalmodel

Hoofdstuk 3

De code: Java en Python

Hier staat het recept uit hoofdstuk 2 als een echt, draaiend programma — eerst in pure Java, daarna precies hetzelfde in pure Python. Beide zonder externe bibliotheken. We lopen de Java-versie stuk voor stuk door, en bekijken daarna wat er in Python verandert.

3.1 — draaien

Zo voer je het uit

Bewaar de code (onderaan, of via de knop) als een bestand met de naam AddNet.java. Je hebt enkel Java nodig, versie 11 of nieuwer — verder niets te installeren.

Zonder apart te compileren (Java 11+)

java AddNet.java

Of klassiek: eerst compileren, dan draaien

javac AddNet.java
java AddNet

Getest met OpenJDK 21. Het programma gebruikt enkel java.util.Random en wat wiskunde uit java.lang.Math — allemaal standaard aanwezig.

3.2 — de opzet

Hyperparameters en geheugen

Bovenaan leggen we de vorm van het netwerk en de leerinstellingen vast, en maken we de tabellen voor de gewichten. De vaste SEED zorgt dat je elke keer exact hetzelfde resultaat krijgt.

opzetAddNet.java
static final int INPUTS  = 10;   // 2 getallen, elk one-hot(5)
static final int HIDDEN  = 16;   // verborgen laag
static final int OUTPUTS = 9;    // sommen 0..8

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

static double[][] W1 = new double[HIDDEN][INPUTS];    // verborgen gewichten
static double[]   b1 = new double[HIDDEN];
static double[][] W2 = new double[OUTPUTS][HIDDEN];   // uitvoergewichten
static double[]   b2 = new double[OUTPUTS];
3.3 — de data

Getallen coderen en de dataset bouwen

encodeInput zet twee getallen om in de invoervector van lengte 10: een 1 in de eerste helft voor a, een 1 in de tweede helft voor b.

one-hot coderingAddNet.java
static double[] encodeInput(int a, int b) {
    double[] x = new double[INPUTS];
    x[a]     = 1.0;   // getal a -> posities 0..4
    x[5 + b] = 1.0;   // getal b -> posities 5..9
    return x;
}

Daarna bouwen we de volledige dataset: alle 25 combinaties, met telkens de invoer, de gewenste uitvoer (one-hot van de som) en de juiste som apart bewaard voor de controle.

alle 25 voorbeeldenAddNet.java
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++;
    }
3.4 — voorwaarts

De voorwaartse doorgang

De verborgen laag: voor elk neuron de gewogen som van de invoer, plus bias, dan ReLU (Math.max(0, ·)).

verborgen laag + ReLUAddNet.java
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
}

De uitvoerlaag levert 9 scores, en softmax maakt er kansen van. Het aftrekken van max verandert het resultaat niet, maar voorkomt te grote getallen in Math.exp.

softmaxAddNet.java
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);   // altijd positief
        sum += out[k];
    }
    for (int k = 0; k < z.length; k++) out[k] /= sum;   // samen = 1
    return out;
}
3.5 — achterwaarts

Backpropagatie

Het hart van het leren. Het foutsignaal bij de uitvoer is simpelweg y - t. Dat sturen we terug naar de verborgen laag, waar de ReLU-poort bepaalt of een neuron meetelt.

gradiëntenAddNet.java
// foutsignaal bij de uitvoer:  kans - gewenst
for (int k = 0; k < OUTPUTS; k++) dz2[k] = y[k] - t[k];

// uitvoergewichten
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
for (int j = 0; j < HIDDEN; j++) {
    double sum = 0;
    for (int k = 0; k < OUTPUTS; k++) sum += W2[k][j] * dz2[k];
    dz1[j] = (z1[j] > 0) ? sum : 0;      // poort van ReLU
}

De accumulatoren gW1, gW2, ... tellen de gradiënten van alle 25 voorbeelden op. Pas daarna sturen we bij met het gemiddelde:

bijsturen (gradient descent)AddNet.java
double scale = LEARNING_RATE / n;          // gemiddelde gradient
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];
}
3.6 — uitvoer

Wat het programma toont

Tijdens het trainen drukt het elke 500 epochs de voortgang af. Daarna test het op alle 25 combinaties en lost het een paar willekeurige sommen op. Dit is de echte uitvoer:

epoch    1   verlies 2.2043   juistheid  12%
epoch  500   verlies 0.6139   juistheid  92%
epoch 1000   verlies 0.1057   juistheid 100%
epoch 1500   verlies 0.0410   juistheid 100%
epoch 2000   verlies 0.0235   juistheid 100%
epoch 2500   verlies 0.0159   juistheid 100%
epoch 3000   verlies 0.0118   juistheid 100%

Resultaat op alle 25 combinaties: 25/25 juist

Enkele willekeurige sommen:
  3 + 2 = 5   (zekerheid  98%)
  1 + 4 = 5   (zekerheid  99%)
  3 + 4 = 7   (zekerheid  99%)
  3 + 1 = 4   (zekerheid  99%)
  3 + 3 = 6   (zekerheid  99%)
  3 + 0 = 3   (zekerheid  99%)
  4 + 3 = 7   (zekerheid  99%)
  2 + 2 = 4   (zekerheid  98%)

Mooi om te zien: na ongeveer 1000 epochs zit het netwerk al op 100%, en daarna blijft het verlies nog dalen — het wordt steeds zekerder van zijn juiste antwoorden (van zo’n 92% naar 98–99% zekerheid).

3.7 — alles samen
De Locale.ROOT bij elke printf is geen versiering: zonder die zou Java op een Nederlandstalig systeem komma's afdrukken (“verlies 2,2043”) en zou jouw uitvoer nét afwijken van wat hierboven staat.

De volledige code

Hier is het complete bestand. Het is precies de code die de uitvoer hierboven produceerde.

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

Nu in Python

Exact hetzelfde netwerk en dezelfde wiskunde, maar in Python — de taal die in machine learning het vaakst gebruikt wordt. Het algoritme is identiek; alleen de schrijfwijze verschilt. De belangrijkste verschillen met Java:

  • Geen types: geen int of double — Python leidt het type zelf af.
  • Inspringing in plaats van accolades: blokken worden door witruimte afgebakend, niet door { }.
  • Lijsten in plaats van arrays, handig aangemaakt met “list comprehensions” zoals [0.0] * HIDDEN.
  • Standaardmodules: random.Random(SEED) en math.exp / math.log in plaats van Java's Random en Math.

De codering en de dataset, regel voor regel herkenbaar:

codering + datasetAddNet.py
def encode_input(a, b):
    x = [0.0] * INPUTS
    x[a] = 1.0        # getal a: posities 0..4
    x[5 + b] = 1.0    # getal b: posities 5..9
    return x

# alle 25 combinaties
X, T, sums = [], [], []
for a in range(5):
    for b in range(5):
        X.append(encode_input(a, b))
        T.append(one_hot(a + b, OUTPUTS))
        sums.append(a + b)

De voorwaartse doorgang, met dezelfde dubbele lus als in Java:

voorwaartsAddNet.py
def predict(x):
    h = [0.0] * HIDDEN
    for j in range(HIDDEN):
        s = b1[j]
        for i in range(INPUTS):
            s += W1[j][i] * x[i]
        h[j] = max(0.0, s)        # ReLU
    z2 = [0.0] * OUTPUTS
    for k in range(OUTPUTS):
        s = b2[k]
        for j in range(HIDDEN):
            s += W2[k][j] * h[j]
        z2[k] = s
    return softmax(z2)

En de kern van het leren — backpropagatie en het bijsturen:

achterwaartsAddNet.py
# foutsignaal bij de uitvoer:  kans - gewenst
dz2 = [y[k] - t[k] for k in range(OUTPUTS)]

# terug naar de verborgen laag
dh = [0.0] * HIDDEN
for j in range(HIDDEN):
    ssum = 0.0
    for k in range(OUTPUTS):
        ssum += W2[k][j] * dz2[k]
    dh[j] = ssum
dz1 = [dh[j] if z1[j] > 0 else 0.0 for j in range(HIDDEN)]   # ReLU-poort

# bijsturen met de gemiddelde gradient
scale = LEARNING_RATE / n
3.9 — draaien & uitvoer

De Python-versie uitvoeren

Bewaar als AddNet.py en draai met Python 3. Geen installatie nodig — enkel de standaardmodules random en math.

python3 AddNet.py

De echte uitvoer (door een andere toevalsgenerator verschillen de exacte getallen van de Java-versie, maar het verloop is hetzelfde):

epoch    1   verlies 2.1881   juistheid  20%
epoch  500   verlies 0.6015   juistheid 100%
epoch 1000   verlies 0.0946   juistheid 100%
epoch 1500   verlies 0.0371   juistheid 100%
epoch 2000   verlies 0.0214   juistheid 100%
epoch 2500   verlies 0.0145   juistheid 100%
epoch 3000   verlies 0.0108   juistheid 100%

Resultaat op alle 25 combinaties: 25/25 juist

Enkele willekeurige sommen:
  1 + 3 = 4   (zekerheid  99%)
  1 + 0 = 1   (zekerheid  99%)
  2 + 2 = 4   (zekerheid  99%)
  0 + 2 = 2   (zekerheid  99%)
  1 + 1 = 2   (zekerheid  98%)
  0 + 2 = 2   (zekerheid  99%)
  4 + 3 = 7   (zekerheid  99%)
  4 + 1 = 5   (zekerheid  99%)
3.10 — volledige Python-code

Het hele Python-programma

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

volledig programmaAddNet.py
"""
AddNet — een klein neuraal netwerk dat leert optellen, in pure Python.
Zelfde opzet als AddNet.java: lijsten in plaats van arrays, lussen in
plaats van matrixwiskunde, geen externe bibliotheken.

  Invoer (10):  one-hot a(5) + one-hot b(5)
  Verborgen:    16 neuronen (ReLU)
  Uitvoer (9):  de mogelijke sommen 0..8 (softmax)

Uitvoeren:  python3 AddNet.py
"""
import random
import math

INPUTS, HIDDEN, OUTPUTS = 10, 16, 9
LEARNING_RATE = 0.1
EPOCHS = 3000
SEED = 42

rng = random.Random(SEED)


def small():
    return (rng.random() - 0.5) * 0.5   # tussen -0.25 en 0.25


# gewichten en biassen
W1 = [[small() for _ in range(INPUTS)] for _ in range(HIDDEN)]
b1 = [0.0] * HIDDEN
W2 = [[small() for _ in range(HIDDEN)] for _ in range(OUTPUTS)]
b2 = [0.0] * OUTPUTS


def encode_input(a, b):
    x = [0.0] * INPUTS
    x[a] = 1.0          # eerste getal  -> posities 0..4
    x[5 + b] = 1.0      # tweede getal  -> posities 5..9
    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]   # -m voor stabiliteit
    s = sum(exps)
    return [e / s for e in exps]


def predict(x):
    h = [0.0] * HIDDEN
    for j in range(HIDDEN):
        s = b1[j]
        for i in range(INPUTS):
            s += W1[j][i] * x[i]
        h[j] = max(0.0, s)              # ReLU
    z2 = [0.0] * OUTPUTS
    for k in range(OUTPUTS):
        s = b2[k]
        for j in range(HIDDEN):
            s += W2[k][j] * h[j]
        z2[k] = s
    return softmax(z2)


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 25 combinaties
X, T, sums = [], [], []
for a in range(5):
    for b in range(5):
        X.append(encode_input(a, b))
        T.append(one_hot(a + b, OUTPUTS))
        sums.append(a + b)
n = len(X)


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


# trainen met volledige-batch gradient descent
for epoch in range(1, EPOCHS + 1):
    gW1 = [[0.0] * INPUTS for _ in range(HIDDEN)]
    gb1 = [0.0] * HIDDEN
    gW2 = [[0.0] * HIDDEN for _ in range(OUTPUTS)]
    gb2 = [0.0] * OUTPUTS
    total_loss = 0.0

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

        # --- voorwaarts ---
        z1 = [0.0] * HIDDEN
        h = [0.0] * HIDDEN
        for j in range(HIDDEN):
            ssum = b1[j]
            for i in range(INPUTS):
                ssum += W1[j][i] * x[i]
            z1[j] = ssum
            h[j] = max(0.0, ssum)
        z2 = [0.0] * OUTPUTS
        for k in range(OUTPUTS):
            ssum = b2[k]
            for j in range(HIDDEN):
                ssum += W2[k][j] * h[j]
            z2[k] = ssum
        y = softmax(z2)

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

        # --- achterwaarts ---
        dz2 = [y[k] - t[k] for k in range(OUTPUTS)]
        for k in range(OUTPUTS):
            gb2[k] += dz2[k]
            for j in range(HIDDEN):
                gW2[k][j] += dz2[k] * h[j]

        dh = [0.0] * HIDDEN
        for j in range(HIDDEN):
            ssum = 0.0
            for k in range(OUTPUTS):
                ssum += W2[k][j] * dz2[k]
            dh[j] = ssum
        dz1 = [dh[j] if z1[j] > 0 else 0.0 for j in range(HIDDEN)]
        for j in range(HIDDEN):
            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):
        b2[k] -= scale * gb2[k]
        for j in range(HIDDEN):
            W2[k][j] -= scale * gW2[k][j]
    for j in range(HIDDEN):
        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}%")


goed = round(accuracy() * n)
print(f"\nResultaat op alle 25 combinaties: {goed}/{n} juist\n")

print("Enkele willekeurige sommen:")
for _ in range(8):
    a = rng.randint(0, 4)
    b = rng.randint(0, 4)
    y = predict(encode_input(a, b))
    guess = index_of_max(y)
    flag = "" if guess == a + b else "   <-- FOUT"
    print(f"  {a} + {b} = {guess}   (zekerheid {100 * y[guess]:3.0f}%){flag}")