import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, random_split
from torchvision import datasets, transforms

# -------------------------------
# 1. Load MNIST
# -------------------------------
transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Lambda(lambda x: x.view(-1))
])

train_full = datasets.MNIST('../data', train=True, download=True, transform=transform)
test_set = datasets.MNIST('../data', train=False, download=True, transform=transform)
train_set, _ = random_split(train_full, [60000, len(train_full)-60000])

batch_size = 100
train_loader = DataLoader(train_set, batch_size=batch_size, shuffle=True)
test_loader = DataLoader(test_set, batch_size=batch_size, shuffle=False)

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

# -------------------------------
# 3. MLP classifier (same as before)
# -------------------------------
class MLP(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(784, 100),
            nn.ReLU(),
            nn.Linear(100, 10)
        )
    def forward(self, x):
        return self.net(x)

mlp = MLP().to(device)
opt_classifier = optim.SGD(mlp.parameters(), lr=0.001)
mse = nn.MSELoss()
criterion = nn.CrossEntropyLoss()

# Helper: pseudo-inverse
def pinv(A):
    return torch.linalg.pinv(A)

W1 = nn.Parameter(0.01 * torch.randn(784, 80, device=device))
W2 = nn.Parameter(0.01 * torch.randn(10, 80, device=device))
W3 = nn.Parameter(0.01 * torch.randn(784, 60, device=device))
W4 = nn.Parameter(0.01 * torch.randn(10, 60, device=device))

opt_w12 = optim.SGD([W1,W2,W3,W4], lr=0.001)
opt_w34 = optim.SGD([W3,W2,W3,W4], lr=0.001)

i = 0
while True:
    for X, yt in train_loader:
        X = X.to(device)
        yt = yt.to(device)
        
        opt_w12.zero_grad()
        for _ in range(10):
            a = X @ W1
            b = torch.eye(10)[yt] @ W2
            
            loss = mse(a,b)
            loss.backward()
            opt_w12.step()

            y = X @ W1 @ W2.T

            err = criterion(y, yt)
            err.backward()
            opt_w12.step()
                        
        score = sum(y.argmax(1)==yt)
        print(loss.item(), score)
