import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
import numpy as np
from copy import deepcopy

# ------------------------------------------------------------
# 1. Define a simple MLP for MNIST
# ------------------------------------------------------------
class MLP(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(28*28, 128)
        self.fc2 = nn.Linear(128, 64)
        self.fc3 = nn.Linear(64, 10)

    def forward(self, x):
        x = x.view(-1, 28*28)
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        return self.fc3(x)

# ------------------------------------------------------------
# 2. Helper functions for parameter vector manipulation
# ------------------------------------------------------------
def get_params(model):
    """Return a flat vector of all model parameters."""
    return torch.cat([p.data.view(-1) for p in model.parameters()])

def set_params(model, vec):
    """Set model parameters from a flat vector."""
    idx = 0
    for p in model.parameters():
        size = p.numel()
        p.data.copy_(vec[idx:idx+size].view(p.shape))
        idx += size

# ------------------------------------------------------------
# 3. Distant‑Target Derivative (SPSA with multiscale)
# ------------------------------------------------------------
def distant_target_gradient(model, loss_fn, x, y, c_values, n_perturbations=1):
    """
    Estimate the gradient of loss w.r.t. model parameters using the
    Distant‑Target Derivative principle.

    Args:
        model: PyTorch model
        loss_fn: loss function (e.g., CrossEntropyLoss)
        x, y: mini‑batch inputs and labels
        c_values: list of step sizes (baselines) to use
        n_perturbations: number of random perturbation vectors per scale

    Returns:
        g_combined: combined gradient vector (same size as parameters)
    """
    # Store current parameters as vector
    theta0 = get_params(model)
    n_params = len(theta0)

    # We'll accumulate gradients for each scale
    g_list = []
    var_list = []   # variance estimate per scale

    for c in c_values:
        g_scale = torch.zeros_like(theta0)
        grad_sq_sum = torch.zeros_like(theta0)  # for variance

        for _ in range(n_perturbations):
            # Random perturbation vector with entries ±1 (SPSA)
            delta = torch.randint(0, 2, (n_params,), dtype=torch.float32) * 2 - 1

            # Evaluate loss at theta + c*delta
            theta_plus = theta0 + c * delta
            set_params(model, theta_plus)
            loss_plus = loss_fn(model(x), y)

            # Evaluate loss at theta - c*delta
            theta_minus = theta0 - c * delta
            set_params(model, theta_minus)
            loss_minus = loss_fn(model(x), y)

            # SPSA gradient estimate: (L_plus - L_minus) / (2*c) * delta^{-1}
            # Since delta_i = ±1, delta^{-1} = delta
            g_est = (loss_plus - loss_minus) / (2 * c) * delta
            g_scale += g_est
            grad_sq_sum += g_est ** 2

        # Average over perturbations
        g_scale /= n_perturbations
        # Variance of gradient estimate (diagonal)
        var_scale = (grad_sq_sum / n_perturbations - g_scale ** 2).mean().item()
        # Avoid zero variance
        var_scale = max(var_scale, 1e-8)

        g_list.append(g_scale)
        var_list.append(var_scale)

    # Restore original parameters
    set_params(model, theta0)

    # Combine scales using inverse‑variance weighting (Axiom 3)
    inv_var = torch.tensor([1.0 / v for v in var_list])
    weights = inv_var / inv_var.sum()
    g_combined = sum(w * g for w, g in zip(weights, g_list))

    return g_combined

# ------------------------------------------------------------
# 4. Training loop with Distant‑Target updates
# ------------------------------------------------------------
def train(model, train_loader, test_loader, epochs=5, lr=0.01, c_values=[0.01, 0.1, 1.0]):
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    model.to(device)
    loss_fn = nn.CrossEntropyLoss()

    # Use a standard SGD optimizer, but we will replace the gradient manually
    optimizer = optim.SGD(model.parameters(), lr=lr)

    for epoch in range(epochs):
        model.train()
        for batch_idx, (data, target) in enumerate(train_loader):
            data, target = data.to(device), target.to(device)

            # Compute gradient using Distant‑Target Derivative
            g = distant_target_gradient(model, loss_fn, data, target, c_values, n_perturbations=2)

            # Zero gradients and set the computed gradient manually
            optimizer.zero_grad()
            # The gradient for each parameter is stored in .grad
            # We need to assign g to the parameter gradients
            idx = 0
            for p in model.parameters():
                size = p.numel()
                p.grad = g[idx:idx+size].view(p.shape).clone()
                idx += size

            # Perform one optimization step
            optimizer.step()

            if batch_idx % 100 == 0:
                loss = loss_fn(model(data), target).item()
                print(f'Train Epoch: {epoch} [{batch_idx * len(data)}/{len(train_loader.dataset)} '
                      f'({100. * batch_idx / len(train_loader):.0f}%)]\tLoss: {loss:.6f}')

        # Evaluate on test set
        model.eval()
        test_loss = 0
        correct = 0
        with torch.no_grad():
            for data, target in test_loader:
                data, target = data.to(device), target.to(device)
                output = model(data)
                test_loss += loss_fn(output, target).item()
                pred = output.argmax(dim=1, keepdim=True)
                correct += pred.eq(target.view_as(pred)).sum().item()

        test_loss /= len(test_loader.dataset)
        accuracy = 100. * correct / len(test_loader.dataset)
        print(f'Test set: Average loss: {test_loss:.4f}, '
              f'Accuracy: {correct}/{len(test_loader.dataset)} ({accuracy:.2f}%)\n')

# ------------------------------------------------------------
# 5. Main: load data, create model, and train
# ------------------------------------------------------------
if __name__ == "__main__":
    # MNIST data
    transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,))])
    train_dataset = datasets.MNIST('../data', train=True, download=True, transform=transform)
    test_dataset = datasets.MNIST('../data', train=False, transform=transform)
    train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
    test_loader = DataLoader(test_dataset, batch_size=1000, shuffle=False)

    model = MLP()
    # Use a few different baselines: coarse, optimal, fine
    # In practice, h* would be tuned per problem.
    c_values = [0.005, 0.02, 0.1]   # example scales

    train(model, train_loader, test_loader, epochs=5, lr=0.01, c_values=c_values)
