import torch
import torchvision
import torchvision.transforms as transforms
import math
import random
from typing import List, Tuple, Union

# -------------------------------------------------------------
# nRAN class (same as before, but we add a __float__ for convenience)
# -------------------------------------------------------------
Number = Union[int, float]

class nRAN:
    # ... (exactly the same implementation as before) ...
    # I'll omit the full class here to keep it concise; assume it's identical.
    # The only addition is:
    def __float__(self):
        return self.collapse()

# -------------------------------------------------------------
# Optimized MLP with nRAN parameters
# -------------------------------------------------------------
class nRAN_MLP:
    def __init__(self, input_size, hidden_size, output_size, max_n=8, init_scale=0.01):
        self.max_n = max_n
        # Use flat lists for weights: W1_flat of length input_size * hidden_size
        # and W2_flat of length hidden_size * output_size
        self.W1 = [nRAN([(random.gauss(0, init_scale), 1)], 0, max_n)
                   for _ in range(input_size * hidden_size)]
        self.b1 = [nRAN([], 0, max_n) for _ in range(hidden_size)]
        self.W2 = [nRAN([(random.gauss(0, init_scale), 1)], 0, max_n)
                   for _ in range(hidden_size * output_size)]
        self.b2 = [nRAN([], 0, max_n) for _ in range(output_size)]

        self.input_size = input_size
        self.hidden_size = hidden_size
        self.output_size = output_size

    def _relu(self, z: nRAN) -> nRAN:
        return z if float(z) > 0 else nRAN([], 0, self.max_n)

    def _softmax(self, logits: List[nRAN]) -> List[float]:
        vals = [float(z) for z in logits]
        maxv = max(vals)
        exp_vals = [math.exp(v - maxv) for v in vals]
        s = sum(exp_vals)
        return [e / s for e in exp_vals]

    def forward(self, x: List[float]):
        # Hidden layer: compute z1 = x * W1 + b1
        z1 = []
        for j in range(self.hidden_size):
            s = nRAN([], 0, self.max_n)
            # Use local variables for speed
            base = j * self.input_size
            for i, xi in enumerate(x):
                s = s + (xi * self.W1[base + i])
            s = s + self.b1[j]
            z1.append(s)

        a1 = [self._relu(z) for z in z1]

        # Output layer
        z2 = []
        for k in range(self.output_size):
            s = nRAN([], 0, self.max_n)
            base = k * self.hidden_size
            for j, aj in enumerate(a1):
                s = s + (aj * self.W2[base + j])
            s = s + self.b2[k]
            z2.append(s)

        probs = self._softmax(z2)
        return probs, z1, a1, z2

    # ---------------------------------------------------------
    # Backward pass – returns flat gradients
    # ---------------------------------------------------------
    def backward(self, x, y_true, probs, z1, a1, z2):
        # Output gradients
        dz2 = [probs[k] - (1.0 if k == y_true else 0.0) for k in range(self.output_size)]

        # dW2 (flat)
        dW2 = [0.0] * (self.hidden_size * self.output_size)
        for k, dz in enumerate(dz2):
            base = k * self.hidden_size
            for j, aj in enumerate(a1):
                dW2[base + j] = float(aj) * dz

        # db2
        db2 = dz2[:]  # copy

        # da1
        da1 = [0.0] * self.hidden_size
        for j in range(self.hidden_size):
            s = 0.0
            for k, dz in enumerate(dz2):
                s += dz * float(self.W2[k * self.hidden_size + j])
            da1[j] = s

        # dz1 = da1 * relu'(z1)
        dz1 = [da1[j] if float(z1[j]) > 0 else 0.0 for j in range(self.hidden_size)]

        # dW1 (flat)
        dW1 = [0.0] * (self.input_size * self.hidden_size)
        for j, dz in enumerate(dz1):
            base = j * self.input_size
            for i, xi in enumerate(x):
                dW1[base + i] = xi * dz

        # db1
        db1 = dz1[:]

        return dW1, db1, dW2, db2

    # ---------------------------------------------------------
    # Update – applies gradients scaled by lr/batch_size
    # ---------------------------------------------------------
    def update(self, dW1, db1, dW2, db2, lr_scaled):
        # lr_scaled = learning_rate / batch_size
        for i, grad in enumerate(dW1):
            self.W1[i] = self.W1[i] - (lr_scaled * grad)
        for j, grad in enumerate(db1):
            self.b1[j] = self.b1[j] - (lr_scaled * grad)
        for i, grad in enumerate(dW2):
            self.W2[i] = self.W2[i] - (lr_scaled * grad)
        for k, grad in enumerate(db2):
            self.b2[k] = self.b2[k] - (lr_scaled * grad)

    def predict(self, x):
        probs, _, _, _ = self.forward(x)
        return max(range(self.output_size), key=lambda i: probs[i])

    def accuracy(self, X, y):
        correct = 0
        for x, label in zip(X, y):
            if self.predict(x) == label:
                correct += 1
        return correct / len(X)


# -------------------------------------------------------------
# Training – now much faster
# -------------------------------------------------------------
def main():
    transform = transforms.Compose([transforms.ToTensor(),
                                    transforms.Normalize((0.1307,), (0.3081,))])
    train_set = torchvision.datasets.MNIST(root='../data', train=True,
                                           download=True, transform=transform)
    test_set = torchvision.datasets.MNIST(root='../data', train=False,
                                          download=True, transform=transform)
    train_loader = torch.utils.data.DataLoader(train_set, batch_size=64, shuffle=True)
    test_loader = torch.utils.data.DataLoader(test_set, batch_size=1000, shuffle=False)

    model = nRAN_MLP(784, 128, 10, max_n=8, init_scale=0.01)
    lr = 0.01
    epochs = 3

    for epoch in range(epochs):
        total_loss = 0.0
        for batch_idx, (data, targets) in enumerate(train_loader):
            X_batch = data.view(data.size(0), -1).tolist()
            y_batch = targets.tolist()
            B = len(X_batch)

            # Initialize flat gradient accumulators with zeros
            dW1_acc = [0.0] * (784 * 128)
            db1_acc = [0.0] * 128
            dW2_acc = [0.0] * (128 * 10)
            db2_acc = [0.0] * 10

            for x, y in zip(X_batch, y_batch):
                probs, z1, a1, z2 = model.forward(x)
                total_loss += model.compute_loss(probs, y)  # compute_loss omitted for brevity, but it's just -log(probs[y])
                dW1, db1, dW2, db2 = model.backward(x, y, probs, z1, a1, z2)
                # Accumulate
                for i, g in enumerate(dW1):
                    dW1_acc[i] += g
                for j, g in enumerate(db1):
                    db1_acc[j] += g
                for i, g in enumerate(dW2):
                    dW2_acc[i] += g
                for k, g in enumerate(db2):
                    db2_acc[k] += g

            # Update with scaled learning rate
            lr_scaled = lr / B
            model.update(dW1_acc, db1_acc, dW2_acc, db2_acc, lr_scaled)

            if batch_idx % 100 == 0:
                avg_loss = total_loss / ((batch_idx + 1) * B)
                print(f"Epoch {epoch+1}, Batch {batch_idx}, Loss {avg_loss:.4f}")

        # Evaluate
        test_X = []
        test_y = []
        for data, targets in test_loader:
            test_X.extend(data.view(data.size(0), -1).tolist())
            test_y.extend(targets.tolist())
        acc = model.accuracy(test_X, test_y)
        print(f"Epoch {epoch+1} test accuracy: {acc:.4f}")

if __name__ == "__main__":
    main()
