import torch
import torch.nn as nn
import torch.fft
import torchvision
import torchvision.transforms as transforms
import matplotlib.pyplot as plt

torch.manual_seed(42)

# ------------------- SPDE Weight Field (same as before) -------------------
class SPDEWeightField(nn.Module):
    def __init__(self, H, W, initial_kappa=10.0):
        super().__init__()
        self.H, self.W = H, W
        self.log_kappa = nn.Parameter(torch.tensor(initial_kappa).log())
        # trainable complex noise in Fourier domain
        self.noise_hat = nn.Parameter(torch.randn(H, W, dtype=torch.complex64) * 0.1)
        # fixed wave‑number grid
        kx = torch.fft.fftfreq(H) * 2 * torch.pi
        ky = torch.fft.fftfreq(W) * 2 * torch.pi
        KX, KY = torch.meshgrid(kx, ky, indexing='ij')
        self.register_buffer('k_sq', KX**2 + KY**2)

    def forward(self):
        kappa = torch.exp(self.log_kappa)
        f_hat = self.noise_hat / (kappa**2 + self.k_sq)
        return torch.fft.ifft2(f_hat).real   # shape (H, W)

# ------------------- SPDE Linear Layer -------------------
class SPDELinear(nn.Module):
    def __init__(self, in_features, out_features, initial_kappa=10.0):
        super().__init__()
        # Weight field grid = (out_features, in_features)
        self.field = SPDEWeightField(out_features, in_features, initial_kappa)

    def forward(self, x):
        W = self.field()            # (out_f, in_f)
        return nn.functional.linear(x, W)

# ------------------- Data Loading -------------------
transform = transforms.Compose([transforms.ToTensor(),
                                transforms.Normalize((0.1307,), (0.3081,))])
batch_size = 128
train_loader = torch.utils.data.DataLoader(
    torchvision.datasets.MNIST(root='../data', train=True, download=True, transform=transform),
    batch_size=batch_size, shuffle=True)
test_loader = torch.utils.data.DataLoader(
    torchvision.datasets.MNIST(root='../data', train=False, download=True, transform=transform),
    batch_size=batch_size, shuffle=False)

# ------------------- Model, Loss, Optimizer -------------------
model = nn.Sequential(
    nn.Flatten(),
    SPDELinear(28*28, 10, initial_kappa=10.0)
)

criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)

# ------------------- Training Loop -------------------
epochs = 5
for epoch in range(epochs):
    model.train()
    running_loss = 0.0
    for images, labels in train_loader:
        optimizer.zero_grad()
        outputs = model(images)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
        running_loss += loss.item() * images.size(0)
    avg_loss = running_loss / len(train_loader.dataset)
    print(f"Epoch {epoch+1}/{epochs} - Loss: {avg_loss:.4f}")

    # Quick test evaluation
    model.eval()
    correct = 0
    total = 0
    with torch.no_grad():
        for images, labels in test_loader:
            outputs = model(images)
            _, predicted = torch.max(outputs, 1)
            total += labels.size(0)
            correct += (predicted == labels).sum().item()
    print(f"  Test Accuracy: {100 * correct / total:.2f}%")

# ------------------- Visualisation -------------------
with torch.no_grad():
    spde_layer = model[1]                          # SPDELinear
    weight_field = spde_layer.field()              # shape (10, 784)
    kappa = torch.exp(spde_layer.field.log_kappa).item()

    # Plot original 10x784 field as 28x28 digit maps
    fig, axes = plt.subplots(1, 10, figsize=(15, 2))
    for i in range(10):
        w = weight_field[i].reshape(28, 28)
        axes[i].imshow(w, cmap='RdBu', vmin=-w.abs().max(), vmax=w.abs().max())
        axes[i].set_title(f"Digit {i}")
        axes[i].axis('off')
    plt.suptitle(f"Learned SPDE weight field (10×784)  κ = {kappa:.2f}", fontsize=14)
    plt.tight_layout()
    plt.show()

    # Upsample to 2× resolution in the input dimension (784 -> 1568)
    H_fine, W_fine = 10, 28*56   # still 10 classes, but each digit map now 28x56
    kx_f = torch.fft.fftfreq(H_fine) * 2 * torch.pi
    ky_f = torch.fft.fftfreq(W_fine) * 2 * torch.pi
    KX_f, KY_f = torch.meshgrid(kx_f, ky_f, indexing='ij')
    k_sq_f = KX_f**2 + KY_f**2

    # Interpolate noise_hat to the fine grid
    noise_hat_fine = nn.functional.interpolate(
        spde_layer.field.noise_hat[None, None, ...].real, 
        size=(H_fine, W_fine), mode='bilinear', align_corners=False
    ).squeeze().to(torch.complex64)

    kappa_val = torch.exp(spde_layer.field.log_kappa)
    f_hat_fine = noise_hat_fine / (kappa_val**2 + k_sq_f)
    field_fine = torch.fft.ifft2(f_hat_fine).real  # (10, 1568)

    # Show the upsampled digit patterns (each row is a digit, shown as 28x56)
    fig, axes = plt.subplots(2, 5, figsize=(15, 5))
    for i in range(10):
        ax = axes[i//5, i%5]
        w = field_fine[i].reshape(28, 56)
        im = ax.imshow(w, cmap='RdBu', aspect='auto',
                       vmin=-w.abs().max(), vmax=w.abs().max())
        ax.set_title(f"Digit {i}")
        ax.axis('off')
    plt.suptitle(f"Upsampled SPDE weight field (28×56 pixels per digit)", fontsize=14)
    plt.tight_layout()
    plt.show()

    # For comparison: train a normal linear layer
    model_lin = nn.Sequential(nn.Flatten(), nn.Linear(28*28, 10, bias=False))
    opt_lin = torch.optim.Adam(model_lin.parameters(), lr=0.01)
    for epoch in range(epochs):
        for images, labels in train_loader:
            opt_lin.zero_grad()
            loss = criterion(model_lin(images), labels)
            loss.backward()
            opt_lin.step()
    with torch.no_grad():
        lin_weights = model_lin[1].weight.reshape(10, 28, 28)
        fig, axes = plt.subplots(1, 10, figsize=(15, 2))
        for i in range(10):
            w = lin_weights[i]
            axes[i].imshow(w, cmap='RdBu', vmin=-w.abs().max(), vmax=w.abs().max())
            axes[i].set_title(f"Digit {i}")
            axes[i].axis('off')
        plt.suptitle("Standard Linear Layer Weights (no smoothness)", fontsize=14)
        plt.tight_layout()
        plt.show()
