"""
AddNetMiniBatch — dezelfde Optelnet als in hoofdstuk 1-4, maar nu getraind met
MINI-BATCH stochastische gradient descent in plaats van volledige-batch.

Het enige wat verandert ten opzichte van AddNet.py zit in de trainingslus:
we schudden elke epoch de dataset, knippen ze in kleine batches, en sturen
de gewichten na ELKE batch bij (in plaats van één keer per epoch over alles).

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

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

INPUTS, HIDDEN, OUTPUTS = 10, 16, 9
LEARNING_RATE = 0.1
EPOCHS = 400
BATCH_SIZE = 5          # <-- de kern van dit hoofdstuk: kleine batches
SEED = 42

rng = random.Random(SEED)


def small():
    return (rng.random() - 0.5) * 0.5


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
    x[5 + b] = 1.0
    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):
    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)
    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


def train_on_batch(batch):
    """Eén bijstuurstap op basis van alleen de voorbeelden in 'batch'."""
    gW1 = [[0.0] * INPUTS for _ in range(HIDDEN)]
    gb1 = [0.0] * HIDDEN
    gW2 = [[0.0] * HIDDEN for _ in range(OUTPUTS)]
    gb2 = [0.0] * OUTPUTS
    batch_loss = 0.0

    for s in batch:
        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)

        batch_loss += -math.log(max(y[index_of_max(t)], 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 over deze batch
    scale = LEARNING_RATE / len(batch)
    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]
    return batch_loss


# trainen met mini-batch SGD
order = list(range(n))
for epoch in range(1, EPOCHS + 1):
    rng.shuffle(order)                                   # 1) elke epoch opnieuw schudden
    epoch_loss = 0.0
    for start in range(0, n, BATCH_SIZE):                # 2) in batches doorlopen
        batch = order[start:start + BATCH_SIZE]
        epoch_loss += train_on_batch(batch)              # 3) na elke batch bijsturen

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

batches_per_epoch = (n + BATCH_SIZE - 1) // BATCH_SIZE
print(f"\n{BATCH_SIZE} voorbeelden per batch  ->  {batches_per_epoch} bijstuurstappen per epoch "
      f"(volledige-batch zou er 1 doen)")

goed = round(accuracy() * n)
print(f"Resultaat 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}")
