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