import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torch.utils.data import DataLoader, Subset
from torchvision import datasets, transforms
import numpy as np
import time
import matplotlib.pyplot as plt
from copy import deepcopy
try:
    from torch.func import functional_call
except ImportError:
    # Fallback for older torch versions if needed, though 2.9.1 should have it
    from torch.nn.utils import stateless
    functional_call = stateless.functional_call


# ----------------------------
# 1. Base CNN model for MNIST
# ----------------------------
class MNISTCNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 16, kernel_size=5, padding=2)
        self.conv2 = nn.Conv2d(16, 32, kernel_size=5, padding=2)
        self.fc1 = nn.Linear(32*7*7, 128)
        self.fc2 = nn.Linear(128, 10)

    def forward(self, x):
        x = F.relu(self.conv1(x))
        x = F.max_pool2d(x, 2)
        x = F.relu(self.conv2(x))
        x = F.max_pool2d(x, 2)
        x = x.view(x.size(0), -1)
        x = F.relu(self.fc1(x))
        x = self.fc2(x)
        return x

# ----------------------------
# 2. Learned Conv Pattern (3D Conv)
# ----------------------------
class LearnedPattern(nn.Module):
    """
    Applies a learnable (1, k, k) 3D convolution to a 4D gradient tensor
    of shape (C_out, C_in, H, W), producing an output of the same shape.
    """
    def __init__(self, in_channels, kernel_size=3):
        super().__init__()
        padding = kernel_size // 2
        # We treat the input as a 5D tensor: (1, 1, C_out, C_in, H, W)
        # Conv3d with kernel (1, k, k) keeps the first two spatial dims unchanged.
        self.conv = nn.Conv3d(
            in_channels=1, out_channels=1,
            kernel_size=(1, kernel_size, kernel_size),
            padding=(0, padding, padding),
            bias=True
        )
        # Initialise to near-identity so that initially the update ≈ dW
        with torch.no_grad():
            self.conv.weight.data.zero_()
            self.conv.weight.data[:, :, 0, padding, padding] = 1.0
            self.conv.bias.data.zero_()

    def forward(self, dW):
        # dW: (C_out, C_in, H, W)
        orig_shape = dW.shape
        # Reshape to (1, 1, C_out * C_in, H, W) for Conv3d
        x = dW.view(1, 1, -1, orig_shape[2], orig_shape[3])
        out = self.conv(x)
        # Reshape back to (C_out, C_in, H, W)
        return out.view(orig_shape)

# ----------------------------
# 3. Custom Optimiser using the learned pattern
# ----------------------------
class LearnedConvOptimizer:
    """
    A custom optimiser that updates parameters using:
        ΔW = learned_pattern(dW)
        W_new = W - lr * ΔW
    Only works for Conv2d weight tensors (4D). Biases & linear layers are updated
    with standard SGD (momentum=0) for simplicity.
    """
    def __init__(self, params, pattern, lr=0.01):
        self.params = list(params)
        self.pattern = pattern  # a LearnedPattern instance
        self.lr = lr

    def zero_grad(self):
        for p in self.params:
            if p.grad is not None:
                p.grad.zero_()

    def step(self):
        for p in self.params:
            if p.grad is None:
                continue
            # Only apply learned pattern to Conv2d weight tensors (4D)
            if p.dim() == 4:
                dW = p.grad
                update = self.pattern(dW)   # shape same as p
                p.data.add_(update, alpha=-self.lr)
            else:
                # Standard SGD for biases, linear weights, etc.
                p.data.add_(p.grad, alpha=-self.lr)

# ----------------------------
# 4. Helper to set weights of a cloned model
# ----------------------------
def set_model_weights(model, flat_weights):
    """Set model parameters from a list of tensors (used for meta-unrolling)."""
    idx = 0
    for p in model.parameters():
        p.data.copy_(flat_weights[idx])
        idx += 1

def get_flat_weights(model):
    """Return a list of parameter tensors (cloned)."""
    return [p.clone().requires_grad_(True) for p in model.parameters()]

# ----------------------------
# 5. Data loading for MNIST
# ----------------------------
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,))])
train_set = datasets.MNIST('./data', train=True, download=True, transform=transform)
test_set  = datasets.MNIST('./data', train=False, download=True, transform=transform)

# Meta-training uses a small subset as support/query to keep things fast.
meta_train_indices = list(range(0, 10000))      # 10k samples for meta-training
meta_val_indices   = list(range(10000, 12000))  # 2k for meta-validation

meta_train_set = Subset(train_set, meta_train_indices)
meta_val_set   = Subset(train_set, meta_val_indices)

meta_train_loader = DataLoader(meta_train_set, batch_size=128, shuffle=True)
meta_val_loader   = DataLoader(meta_val_set,   batch_size=128, shuffle=False)

# Base training uses the full training set (after meta-training is done)
base_train_loader = DataLoader(train_set, batch_size=64, shuffle=True, num_workers=2)
test_loader       = DataLoader(test_set,  batch_size=64, shuffle=False, num_workers=2)

# ----------------------------
# 6. Meta-Training the Learned Pattern
# ----------------------------
def meta_train_pattern(base_model, pattern, meta_steps=200, inner_lr=0.01, outer_lr=1e-3):
    """
    Meta-train the pattern parameters θ using a support/query split.
    Returns the trained pattern.
    """
    outer_optimizer = optim.Adam(pattern.parameters(), lr=outer_lr)
    criterion = nn.CrossEntropyLoss()
    meta_iter = iter(meta_train_loader)

    print("Starting meta-training of the learned pattern...")
    for step in range(meta_steps):
        try:
            support_x, support_y = next(meta_iter)
            query_x, query_y = next(meta_iter)
        except StopIteration:
            meta_iter = iter(meta_train_loader)
            support_x, support_y = next(meta_iter)
            query_x, query_y = next(meta_iter)


        # 1. Get initial weights as a dict (cloned and with requires_grad)
        weights_dict = {n: p.clone().requires_grad_(True) for n, p in base_model.named_parameters()}
        param_names = list(weights_dict.keys())

        # 2. Compute loss on support batch using weights_dict
        logits_support = functional_call(base_model, weights_dict, (support_x,))
        loss_support = criterion(logits_support, support_y)

        # 3. Compute gradients dW w.r.t. weights_dict
        dW_list = torch.autograd.grad(
            loss_support, weights_dict.values(), 
            create_graph=True, allow_unused=True
        )

        # 4. Apply learned pattern to each gradient
        updates = []
        for name, dw in zip(param_names, dW_list):
            w = weights_dict[name]
            if dw is None:
                updates.append(torch.zeros_like(w))
                continue
            
            if w.dim() == 4:
                # Apply the learned pattern to Conv2d weights
                upd = pattern(dw)
            else:
                # For biases/linear, use plain gradient (SGD)
                upd = dw
            updates.append(upd)

        # 5. Perform inner update: W' = W - inner_lr * update
        new_weights_dict = {
            name: weights_dict[name] - inner_lr * upd
            for name, upd in zip(param_names, updates)
        }

        # 6. Compute meta-loss on query batch using the updated weights
        logits_query = functional_call(base_model, new_weights_dict, (query_x,))
        loss_query = criterion(logits_query, query_y)

        # 7. Compute meta-gradients w.r.t. pattern parameters θ
        outer_optimizer.zero_grad()
        loss_query.backward()
        # Gradient clipping to stabilise meta-training
        torch.nn.utils.clip_grad_norm_(pattern.parameters(), max_norm=1.0)
        outer_optimizer.step()

        if step % 50 == 0:
            print(f"Meta-step {step:4d}, meta-loss: {loss_query.item():.4f}")

    return pattern

# ----------------------------
# 7. Base Training (using the meta-trained pattern)
# ----------------------------
def train_base_model(base_model, pattern, epochs=5, lr=0.01):
    """
    Train the base model using the custom LearnedConvOptimizer.
    Also train a standard Adam baseline for comparison.
    """
    # --- Baseline: Adam ---
    print("\n--- Baseline: Adam ---")
    model_adam = MNISTCNN()
    adam_opt = optim.Adam(model_adam.parameters(), lr=1e-3)
    criterion = nn.CrossEntropyLoss()
    adam_accuracies = []
    for epoch in range(epochs):
        model_adam.train()
        total_loss = 0
        for x, y in base_train_loader:
            adam_opt.zero_grad()
            loss = criterion(model_adam(x), y)
            loss.backward()
            adam_opt.step()
            total_loss += loss.item()
        # Evaluate on test set
        model_adam.eval()
        correct = 0
        with torch.no_grad():
            for x, y in test_loader:
                pred = model_adam(x).argmax(dim=1)
                correct += (pred == y).sum().item()
        acc = correct / len(test_set)
        adam_accuracies.append(acc)
        print(f"Epoch {epoch+1}: Adam loss = {total_loss/len(base_train_loader):.4f}, test acc = {acc:.4f}")

    # --- Learned Conv Optimizer ---
    print("\n--- Learned Conv Optimizer ---")
    model_learned = MNISTCNN()
    # Use the meta-trained pattern (shared across layers? We'll use the same pattern instance,
    # but note that our pattern expects (C_in, C_out, H, W) which varies per layer.
    # To handle multiple layers, we need a separate pattern per layer or a single pattern
    # that adapts to different input sizes. Here we use the same pattern for both conv layers.
    # For simplicity, we re-initialise the pattern for each layer? 
    # In practice we would have one pattern per layer, but to keep it simple we'll use
    # one pattern for conv1 and another for conv2, sharing the same meta-trained θ?
    # Actually our LearnedPattern works with any (C_out, C_in, H, W) because Conv3d
    # is size-agnostic (the kernel is fixed, but it slides over spatial dimensions).
    # The only issue is that the number of input channels (C_in) must match. 
    # Conv1: (16, 1, 5, 5), Conv2: (32, 16, 5, 5). They have different C_in.
    # We can either use separate pattern instances or use a single one that accepts any C_in.
    # The current implementation uses Conv3d with in_channels=1, out_channels=1, which works for any C_in,C_out
    # because it treats C_out and C_in as spatial dimensions of the 3D volume. So ONE pattern works for ALL conv layers!
    # Great.
    learned_opt = LearnedConvOptimizer(model_learned.parameters(), pattern, lr=lr)
    learned_accuracies = []
    for epoch in range(epochs):
        model_learned.train()
        total_loss = 0
        for x, y in base_train_loader:
            learned_opt.zero_grad()
            loss = criterion(model_learned(x), y)
            loss.backward()
            learned_opt.step()
            total_loss += loss.item()
        # Evaluate
        model_learned.eval()
        correct = 0
        with torch.no_grad():
            for x, y in test_loader:
                pred = model_learned(x).argmax(dim=1)
                correct += (pred == y).sum().item()
        acc = correct / len(test_set)
        learned_accuracies.append(acc)
        print(f"Epoch {epoch+1}: Learned loss = {total_loss/len(base_train_loader):.4f}, test acc = {acc:.4f}")

    return adam_accuracies, learned_accuracies

# ----------------------------
# 8. Main execution
# ----------------------------
if __name__ == "__main__":
    # Instantiate base model (only used for its structure during meta-training)
    base_model = MNISTCNN()

    # Create the learned pattern (kernel size 3)
    pattern = LearnedPattern(in_channels=1, kernel_size=3)

    # ---- Meta-train the pattern ----
    pattern = meta_train_pattern(base_model, pattern, meta_steps=200, inner_lr=0.01, outer_lr=1e-3)

    # ---- Train base model with the learned pattern ----
    adam_acc, learned_acc = train_base_model(base_model, pattern, epochs=5, lr=0.01)

    # Print summary
    print("\n=== Final Test Accuracies ===")
    print(f"Adam:        {adam_acc[-1]:.4f}")
    print(f"LearnedConv: {learned_acc[-1]:.4f}")
