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 scipy.interpolate import CubicSpline
from tqdm import tqdm
import matplotlib.pyplot as plt

# --------------------------------------------
# 1. Define the Teacher CNN
# --------------------------------------------
class MNISTCNN(nn.Module):
    """
    A simple CNN with a clear split point:
    - 'lower' (features up to the 128-dim embedding)
    - 'upper' (classifier head: 128 -> 10)
    """
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 32, kernel_size=3, padding=1)
        self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1)
        self.pool = nn.MaxPool2d(2, 2)
        self.dropout1 = nn.Dropout(0.25)
        self.dropout2 = nn.Dropout(0.5)
        self.fc1 = nn.Linear(64 * 7 * 7, 128)   # 7x7 after two 2x2 pools
        self.fc2 = nn.Linear(128, 10)

    def forward(self, x, return_hidden=False):
        # --- Lower half (feature extractor) ---
        x = self.pool(F.relu(self.conv1(x)))
        x = self.pool(F.relu(self.conv2(x)))
        x = self.dropout1(x)
        x = torch.flatten(x, 1)
        h = F.relu(self.fc1(x))      # shape: [batch, 128]
        h = self.dropout2(h)

        if return_hidden:
            return h

        # --- Upper half (classifier) ---
        out = self.fc2(h)
        return out

    # For convenience, split forward into two parts
    def lower(self, x):
        x = self.pool(F.relu(self.conv1(x)))
        x = self.pool(F.relu(self.conv2(x)))
        x = self.dropout1(x)
        x = torch.flatten(x, 1)
        h = F.relu(self.fc1(x))
        h = self.dropout2(h)
        return h

    def upper(self, h):
        return self.fc2(h)


# --------------------------------------------
# 2. Training loop (on full MNIST)
# --------------------------------------------
def train_model(model, train_loader, epochs=5, device='cuda'):
    model.train()
    optimizer = optim.Adam(model.parameters(), lr=0.001)
    criterion = nn.CrossEntropyLoss()

    for epoch in range(epochs):
        total_loss = 0
        correct = 0
        for data, target in tqdm(train_loader, desc=f'Epoch {epoch+1}'):
            data, target = data.to(device), target.to(device)
            optimizer.zero_grad()
            output = model(data)
            loss = criterion(output, target)
            loss.backward()
            optimizer.step()

            total_loss += loss.item()
            pred = output.argmax(dim=1)
            correct += pred.eq(target).sum().item()

        acc = 100. * correct / len(train_loader.dataset)
        print(f'Epoch {epoch+1}: Loss {total_loss/len(train_loader):.4f}, Acc {acc:.2f}%')


# --------------------------------------------
# 3. ODE‑Spline Denoising Inference
# --------------------------------------------
def ode_spline_denoise(model, image, noise_pattern, t_values, alpha0=0.5, decay=2.0, device='cuda'):
    """
    image: (1, 28, 28) tensor on device
    noise_pattern: (1, 28, 28) tensor with same std=1
    t_values: 1D array of time points (e.g., np.linspace(0, 2, 12))
    alpha(t) = alpha0 * exp(-decay * t)
    Returns denoised logits.
    """
    model.eval()
    with torch.no_grad():
        # 3a. Generate perturbed images and collect hidden states
        hidden_states = []
        for t in t_values:
            alpha_t = alpha0 * np.exp(-decay * t)
            perturbed = image + alpha_t * noise_pattern
            h = model.lower(perturbed)          # shape: [1, 128]
            hidden_states.append(h.cpu().numpy().flatten())

        hidden_states = np.array(hidden_states)  # shape: (N_points, 128)

        # 3b. Fit cubic spline for each hidden dimension
        splines = []
        curvature_norms = []
        for dim in range(hidden_states.shape[1]):
            spline = CubicSpline(t_values, hidden_states[:, dim], bc_type='natural')
            splines.append(spline)

            # Evaluate second derivative at each sampled point to find curvature
            d2 = spline(t_values, 2)
            curvature_norms.append(np.linalg.norm(d2))

        # 3c. Find the time point with minimum *average* curvature across dimensions
        # We want to evaluate the spline where the trajectory is smoothest.
        curv_avg = np.mean(np.array(curvature_norms), axis=0)   # shape: (N_points,)
        t_star_idx = np.argmin(curv_avg)
        t_star = t_values[t_star_idx]

        # 3d. Evaluate all splines at t_star to get denoised hidden state
        h_denoised = np.array([splines[dim](t_star) for dim in range(hidden_states.shape[1])])
        h_denoised = torch.tensor(h_denoised, dtype=torch.float32, device=device).unsqueeze(0)  # (1, 128)

        # 3e. Pass through the upper half
        logits = model.upper(h_denoised)
        return logits


# --------------------------------------------
# 4. Evaluation on Test Set
# --------------------------------------------
def evaluate_vanilla(model, test_loader, device='cuda'):
    model.eval()
    correct = 0
    total = 0
    with torch.no_grad():
        for data, target in tqdm(test_loader, desc='Vanilla'):
            data, target = data.to(device), target.to(device)
            output = model(data)
            pred = output.argmax(dim=1)
            correct += pred.eq(target).sum().item()
            total += target.size(0)
    return 100. * correct / total


def evaluate_ode_spline(model, test_loader, noise_patterns, device='cuda'):
    """
    noise_patterns: a pre-generated list of fixed noise tensors (one per test image)
    to ensure determinism.
    """
    model.eval()
    correct = 0
    total = 0

    # ODE parameters
    t_values = np.linspace(0, 2.5, 10)   # 10 points along the ODE trajectory
    alpha0 = 0.8
    decay = 2.0

    with torch.no_grad():
        for idx, (data, target) in enumerate(tqdm(test_loader, desc='ODE-Spline')):
            data, target = data.to(device), target.to(device)

            # Each image gets its own fixed noise pattern (can be precomputed)
            noise = noise_patterns[idx * data.size(0): (idx+1) * data.size(0)].to(device)

            for i in range(data.size(0)):
                logits = ode_spline_denoise(
                    model,
                    data[i].unsqueeze(0),
                    noise[i].unsqueeze(0),
                    t_values,
                    alpha0,
                    decay,
                    device
                )
                pred = logits.argmax(dim=1)
                correct += (pred == target[i]).sum().item()
                total += 1

    return 100. * correct / total


# --------------------------------------------
# 5. Main Script
# --------------------------------------------
if __name__ == "__main__":
    device = 'cuda' if torch.cuda.is_available() else 'cpu'
    print(f"Using device: {device}")

    # 5a. 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, download=True, transform=transform)

    train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
    test_loader = DataLoader(test_dataset, batch_size=128, shuffle=False)

    # 5b. Train the model (or load pre-trained weights)
    model = MNISTCNN().to(device)

    # Uncomment to train from scratch:
    # train_model(model, train_loader, epochs=5, device=device)
    # torch.save(model.state_dict(), 'mnist_cnn_teacher.pt')

    # For quick testing, load pre-trained weights if available, else train.
    try:
        model.load_state_dict(torch.load('mnist_cnn_teacher.pt', map_location=device))
        print("Loaded pre-trained weights.")
    except FileNotFoundError:
        print("Training from scratch...")
        train_model(model, train_loader, epochs=5, device=device)
        torch.save(model.state_dict(), 'mnist_cnn_teacher.pt')

    # 5c. Vanilla evaluation
    vanilla_acc = evaluate_vanilla(model, test_loader, device)
    print(f"\n✅ Vanilla inference accuracy: {vanilla_acc:.2f}%\n")

    # 5d. Precompute a fixed noise pattern for each test sample
    # (We use the same noise pattern across all ODE steps for that sample)
    print("Precomputing noise patterns for ODE-Spline...")
    noise_patterns = []
    for data, _ in test_loader:
        for img in data:
            # Random Gaussian noise with unit variance, same shape as image
            noise = torch.randn_like(img)
            noise_patterns.append(noise)
    noise_patterns = torch.stack(noise_patterns)  # (N_test, 1, 28, 28)

    # 5e. ODE-Spline evaluation
    ode_acc = evaluate_ode_spline(model, test_loader, noise_patterns, device)
    print(f"\n🧊 ODE-Spline denoised inference accuracy: {ode_acc:.2f}%\n")

    # 5f. Compare results
    improvement = ode_acc - vanilla_acc
    print(f"📈 Improvement: {improvement:+.2f} percentage points")
