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

# -------------------------------------------------------------------
# 1. nRAN class (from nran-theory-implementation.md)
# -------------------------------------------------------------------
Number = Union[int, float]

class nRAN:
    """
    Multi-term Rational-Addition Number: sum_i (a_i/b_i) + c.
    Keeps up to max_n rational terms; compacts if exceeded.
    """

    def __init__(self, terms: List[Tuple[Number, Number]] = None,
                 c: Number = 0, max_n: int = None):
        self.terms = list(terms) if terms else []
        self.c = c
        self.max_n = max_n
        for a, b in self.terms:
            if b == 0:
                raise ValueError("b cannot be zero")
        if max_n is not None and len(self.terms) > max_n:
            self._compact()

    def collapse(self) -> float:
        """Evaluate to a single float."""
        return sum(a / b for a, b in self.terms) + self.c

    def __float__(self):
        return self.collapse()

    def __repr__(self):
        return f"nRAN({self.terms}, c={self.c!r})"

    def __str__(self):
        if not self.terms:
            return str(self.c)
        parts = " + ".join(f"{a}/{b}" for a, b in self.terms)
        return f"{parts} + {self.c}"

    def _combine_rational_terms(self):
        """Combine all integer terms into one exact fraction."""
        int_terms = [(a, b) for a, b in self.terms
                     if isinstance(a, int) and isinstance(b, int)]
        float_terms = [(a, b) for a, b in self.terms
                       if not (isinstance(a, int) and isinstance(b, int))]
        if not int_terms:
            return float_terms
        total = sum(Fraction(a, b) for a, b in int_terms)
        return [(total.numerator, total.denominator)] + float_terms

    def _compact(self):
        """Reduce term count to respect max_n."""
        if self.max_n is None:
            return
        self.terms = self._combine_rational_terms()
        if len(self.terms) > self.max_n:
            rational_value = sum(a / b for a, b in self.terms)
            self.terms = []
            self.c += rational_value

    def _after_op(self, result):
        if self.max_n is not None:
            result._compact()
        return result

    def __add__(self, other):
        if isinstance(other, (int, float)):
            return nRAN(self.terms, self.c + other, self.max_n)
        if not isinstance(other, nRAN):
            return NotImplemented
        return self._after_op(
            nRAN(self.terms + other.terms, self.c + other.c, self.max_n)
        )

    def __sub__(self, other):
        if isinstance(other, (int, float)):
            return nRAN(self.terms, self.c - other, self.max_n)
        if not isinstance(other, nRAN):
            return NotImplemented
        neg_terms = [(-a, b) for a, b in other.terms]
        return self._after_op(
            nRAN(self.terms + neg_terms, self.c - other.c, self.max_n)
        )

    def __mul__(self, other):
        if isinstance(other, (int, float)):
            new_terms = [(a * other, b) for a, b in self.terms]
            return nRAN(new_terms, self.c * other, self.max_n)
        if not isinstance(other, nRAN):
            return NotImplemented
        new_terms = []
        for a, b in self.terms:
            for d, e in other.terms:
                new_terms.append((a * d, b * e))
        for a, b in self.terms:
            new_terms.append((a * other.c, b))
        for d, e in other.terms:
            new_terms.append((self.c * d, e))
        return self._after_op(
            nRAN(new_terms, self.c * other.c, self.max_n)
        )

    def __truediv__(self, other):
        if isinstance(other, (int, float)):
            new_terms = [(a, b * other) for a, b in self.terms]
            return nRAN(new_terms, self.c / other, self.max_n)
        if not isinstance(other, nRAN):
            return NotImplemented
        combined = other._combine_rational_terms()
        if len(combined) != 1:
            return self * (1 / other.collapse())
        d, e = combined[0]
        denom = d + other.c * e
        reciprocal = nRAN([(e, denom)], 0, self.max_n)
        return self * reciprocal

    __radd__ = __add__
    __rmul__ = __mul__

    def __rsub__(self, other):
        return (-self) + other

    def __neg__(self):
        return nRAN([(-a, b) for a, b in self.terms], -self.c, self.max_n)


# -------------------------------------------------------------------
# 2. Optimized MLP using nRAN parameters (flat lists)
# -------------------------------------------------------------------
class nRAN_MLP:
    def __init__(self, input_size: int, hidden_size: int, output_size: int,
                 max_n: int = 8, init_scale: float = 0.01):
        self.input_size = input_size
        self.hidden_size = hidden_size
        self.output_size = output_size
        self.max_n = max_n

        # Flat weight lists: W1 length = input_size * hidden_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)]

    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
        z1 = []
        for j in range(self.hidden_size):
            s = nRAN([], 0, self.max_n)
            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

    def compute_loss(self, probs: List[float], y_true: int) -> float:
        return -math.log(probs[y_true] + 1e-12)

    def backward(self, x, y_true, probs, z1, a1, z2):
        # dz2 = probs - one_hot
        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 = dz2[:]

        # 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
        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 = dz1[:]

        return dW1, db1, dW2, db2

    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)


# -------------------------------------------------------------------
# 3. Training loop
# -------------------------------------------------------------------
def main():
    # Load MNIST
    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)

    # Create model
    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
            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)
                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 after each epoch
        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()