"""
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}")
