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

try:
    from torch.func import functional_call
except ImportError:
    from torch.nn.utils import stateless
    functional_call = stateless.functional_call


# ----------------------------
# Device
# ----------------------------
if torch.cuda.is_available():
    device = torch.device("cuda")
elif torch.backends.mps.is_available():
    device = torch.device("mps")
else:
    device = torch.device("cpu")
print(f"Using device: {device}")


# ----------------------------
# 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) — works for any (C_out, C_in, H, W)
# ----------------------------
class LearnedPattern(nn.Module):
    def __init__(self, kernel_size=3):
        super().__init__()
        padding = kernel_size // 2
        self.conv = nn.Conv3d(
            in_channels=1, out_channels=1,
            kernel_size=(1, kernel_size, kernel_size),
            padding=(0, padding, padding),
            bias=True,
        )
        with torch.no_grad():
            self.conv.weight.zero_()
            self.conv.weight[:, :, 0, padding, padding] = 1.0  # near-identity init
            self.conv.bias.zero_()

    def forward(self, dW):
        c_out, c_in, h, w = dW.shape
        x = dW.view(1, 1, c_out * c_in, h, w)
        out = self.conv(x)
        return out.view(c_out, c_in, h, w)


# ----------------------------
# 3. Custom optimiser using the learned pattern (for base training)
# ----------------------------
class LearnedConvOptimizer:
    def __init__(self, params, pattern, lr=0.01):
        self.params = list(params)
        self.pattern = pattern
        self.lr = lr

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

    @torch.no_grad()
    def step(self):
        for p in self.params:
            if p.grad is None:
                continue
            if p.dim() == 4:
                update = self.pattern(p.grad)
                p.add_(update, alpha=-self.lr)
            else:
                p.add_(p.grad, alpha=-self.lr)


# ----------------------------
# 4. Data
# ----------------------------
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_train_set = Subset(train_set, list(range(0, 10000)))
meta_train_loader = DataLoader(meta_train_set, batch_size=128, shuffle=True)

base_train_loader = DataLoader(train_set, batch_size=64, shuffle=True, num_workers=2)
test_loader = DataLoader(test_set, batch_size=256, shuffle=False, num_workers=2)


# ----------------------------
# 5. Meta-training the learned pattern
# ----------------------------
def fresh_weights():
    """Fresh random init each meta-step so the pattern learns a *general* rule."""
    m = MNISTCNN().to(device)
    return {n: p.detach().clone().requires_grad_(True) for n, p in m.named_parameters()}


def meta_train_pattern(skeleton, pattern, meta_steps=150,
                       inner_steps=4, inner_lr=0.1, outer_lr=1e-3):
    outer_optimizer = optim.Adam(pattern.parameters(), lr=outer_lr)
    criterion = nn.CrossEntropyLoss()
    meta_iter = iter(meta_train_loader)

    def next_batch():
        nonlocal meta_iter
        try:
            return next(meta_iter)
        except StopIteration:
            meta_iter = iter(meta_train_loader)
            return next(meta_iter)

    print("Starting meta-training of the learned pattern...")
    for step in range(meta_steps):
        support_x, support_y = next_batch()
        query_x, query_y = next_batch()
        support_x, support_y = support_x.to(device), support_y.to(device)
        query_x, query_y = query_x.to(device), query_y.to(device)

        weights = fresh_weights()  # diverse starting point every step

        # --- unroll several inner steps so the meta-loss really depends on theta ---
        for _ in range(inner_steps):
            logits = functional_call(skeleton, weights, (support_x,))
            loss = criterion(logits, support_y)
            grads = torch.autograd.grad(
                loss, weights.values(), create_graph=True, allow_unused=True
            )
            new_weights = {}
            for (name, w), g in zip(weights.items(), grads):
                if g is None:
                    new_weights[name] = w
                elif w.dim() == 4:
                    new_weights[name] = w - inner_lr * pattern(g)
                else:
                    new_weights[name] = w - inner_lr * g
            weights = new_weights

        # --- meta-loss on the query batch ---
        logits_q = functional_call(skeleton, weights, (query_x,))
        loss_query = criterion(logits_q, query_y)

        outer_optimizer.zero_grad()
        loss_query.backward()
        torch.nn.utils.clip_grad_norm_(pattern.parameters(), max_norm=1.0)
        outer_optimizer.step()

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

    return pattern


# ----------------------------
# 6. Base training
# ----------------------------
def evaluate(model):
    model.eval()
    correct = 0
    with torch.no_grad():
        for x, y in test_loader:
            x, y = x.to(device), y.to(device)
            correct += (model(x).argmax(1) == y).sum().item()
    return correct / len(test_set)


def train_one(model, optimizer, epochs, criterion, label):
    accs = []
    for epoch in range(epochs):
        model.train()
        total_loss = 0.0
        for x, y in base_train_loader:
            x, y = x.to(device), y.to(device)
            optimizer.zero_grad()
            loss = criterion(model(x), y)
            loss.backward()
            optimizer.step()
            total_loss += loss.item()
        acc = evaluate(model)
        accs.append(acc)
        print(f"Epoch {epoch+1}: {label} loss = {total_loss/len(base_train_loader):.4f}, "
              f"test acc = {acc:.4f}")
    return accs


def train_base_model(pattern, epochs=5, lr=0.01):
    criterion = nn.CrossEntropyLoss()

    print("\n--- Baseline: Adam ---")
    model_adam = MNISTCNN().to(device)
    adam_acc = train_one(model_adam, optim.Adam(model_adam.parameters(), lr=1e-3),
                         epochs, criterion, "Adam")

    print("\n--- Learned Conv Optimizer ---")
    model_learned = MNISTCNN().to(device)
    learned_opt = LearnedConvOptimizer(model_learned.parameters(), pattern, lr=lr)
    learned_acc = train_one(model_learned, learned_opt, epochs, criterion, "Learned")

    return adam_acc, learned_acc


# ----------------------------
# 7. Main
# ----------------------------
if __name__ == "__main__":
    skeleton = MNISTCNN().to(device)            # only a structure for functional_call
    pattern = LearnedPattern(kernel_size=3).to(device)

    pattern = meta_train_pattern(skeleton, pattern, meta_steps=150,
                                 inner_steps=4, inner_lr=0.1, outer_lr=1e-3)

    adam_acc, learned_acc = train_base_model(pattern, epochs=5, lr=0.01)

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