import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
import torchvision
import torchvision.transforms as transforms
from tqdm import tqdm
import numpy as np

# ----------------------------------------------------------------------
# Operator definitions (replacement for addition in parameter updates)
# ----------------------------------------------------------------------
def operator_add(old, delta):
    return old + delta

def operator_mul(old, delta):
    # multiplicative update: new = old * (1 + delta)
    # delta = -lr * grad, so this is old * (1 - lr*grad)
    return old * (1 + delta)

def operator_geo_mean(old, delta, eps=1e-8):
    # geometric mean of old and (old + delta)
    a = torch.abs(old)
    b = torch.abs(old + delta)
    sign = torch.sign(old * (old + delta))
    return sign * torch.sqrt(a * b + eps)

def operator_max(old, delta):
    return torch.max(old, old + delta)

def operator_phase_locked(old, delta):
    # treat as phase on a circle (mod 2π)
    return torch.remainder(old + delta, 2 * torch.pi)

# List of available operators
OPERATORS = {
    'add': operator_add,
    'mul': operator_mul,
    'geo_mean': operator_geo_mean,
    'max': operator_max,
    'phase_locked': operator_phase_locked,
}
OP_NAMES = list(OPERATORS.keys())
NUM_OPS = len(OP_NAMES)

# ----------------------------------------------------------------------
# Policy network (AI-automaton) that selects operator
# ----------------------------------------------------------------------
class OperatorPolicy(nn.Module):
    def __init__(self, state_dim=8, hidden_dim=32):
        super().__init__()
        self.fc1 = nn.Linear(state_dim, hidden_dim)
        self.fc2 = nn.Linear(hidden_dim, NUM_OPS)
        
    def forward(self, state, temperature=1.0, hard=False):
        # state: (batch,) but we have one state per step -> shape (1, state_dim)
        x = F.relu(self.fc1(state))
        logits = self.fc2(x)  # (1, NUM_OPS)
        # Gumbel-Softmax sampling (differentiable)
        probs = F.gumbel_softmax(logits, tau=temperature, hard=hard, dim=-1)
        # probs is one-hot if hard=True, else soft
        return probs, logits

# ----------------------------------------------------------------------
# Main CNN for CIFAR-10
# ----------------------------------------------------------------------
class SimpleCNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(3, 32, 3, padding=1)
        self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
        self.conv3 = nn.Conv2d(64, 128, 3, padding=1)
        self.pool = nn.MaxPool2d(2, 2)
        self.drop = nn.Dropout(0.25)
        self.fc = nn.Linear(128 * 4 * 4, 10)

    def forward(self, x):
        x = self.pool(F.relu(self.conv1(x)))
        x = self.pool(F.relu(self.conv2(x)))
        x = self.pool(F.relu(self.conv3(x)))
        x = self.drop(x)
        x = x.view(x.size(0), -1)
        return self.fc(x)

# ----------------------------------------------------------------------
# Custom optimizer that uses the policy to choose operator per step
# ----------------------------------------------------------------------
class AdaptiveOperatorSGD:
    def __init__(self, model, policy, lr=0.01, temperature=1.0):
        self.model = model
        self.policy = policy
        self.lr = lr
        self.temperature = temperature
        self.prev_loss = None
        self.step_count = 0

    def step(self, loss, closure=None):
        # Build state for the policy BEFORE backward (uses current loss value)
        grad_norm = 0.0
        param_norm = 0.0
        for p in self.model.parameters():
            param_norm += p.norm().item() ** 2
            if p.grad is not None:
                grad_norm += p.grad.norm().item() ** 2
        grad_norm = float(np.sqrt(grad_norm))
        param_norm = float(np.sqrt(param_norm))

        loss_val = loss.item()
        prev_loss_val = self.prev_loss if self.prev_loss is not None else loss_val
        loss_change = loss_val - prev_loss_val
        step_phase = float(np.sin(2 * np.pi * self.step_count / 1000))

        # Clamp state values to avoid NaN/inf in policy input
        loss_val = float(np.clip(loss_val, 0, 100))
        prev_loss_val = float(np.clip(prev_loss_val, 0, 100))
        loss_change = float(np.clip(loss_change, -10, 10))
        grad_norm = float(np.clip(grad_norm, 0, 1e6))
        param_norm = float(np.clip(param_norm, 0, 1e6))

        state = torch.tensor([[
            loss_val,
            prev_loss_val,
            loss_change,
            grad_norm,
            param_norm,
            step_phase,
            min(self.step_count / 1000, 1.0),
            self.lr
        ]], dtype=torch.float32, device=next(self.model.parameters()).device)

        # Get operator selection
        with torch.no_grad():
            probs, _ = self.policy(state, temperature=self.temperature, hard=True)
            op_idx = torch.argmax(probs, dim=1).item()
        chosen_op_name = OP_NAMES[op_idx]
        chosen_op = OPERATORS[chosen_op_name]

        # Apply the chosen operator to each parameter
        with torch.no_grad():
            for p in self.model.parameters():
                if p.grad is None:
                    continue
                delta = -self.lr * p.grad
                # Apply operator
                new_p = chosen_op(p.data, delta)
                # Clamp to avoid extreme values
                new_p = torch.clamp(new_p, -10, 10)
                p.data.copy_(new_p)

        # Update state for next step
        self.prev_loss = loss.item()
        self.step_count += 1
        return chosen_op_name

# ----------------------------------------------------------------------
# Training loop with adaptive operator selection
# ----------------------------------------------------------------------
def train_adaptive(model, policy, device, train_loader, epochs=20, lr=0.01):
    model.train()
    optimizer = AdaptiveOperatorSGD(model, policy, lr=lr, temperature=1.0)

    for epoch in range(1, epochs+1):
        total_loss = 0
        correct = 0
        epoch_op_counts = {op: 0 for op in OP_NAMES}
        num_batches = 0

        for data, target in tqdm(train_loader, desc=f'Epoch {epoch}'):
            data, target = data.to(device), target.to(device)

            # Standard pattern: zero_grad -> forward -> loss -> backward -> step
            optimizer.model.zero_grad()
            output = model(data)
            loss = F.cross_entropy(output, target)
            loss.backward()

            # Adaptive step (chooses operator, updates model using pre-computed gradients)
            chosen_op = optimizer.step(loss)
            epoch_op_counts[chosen_op] += 1
            num_batches += 1

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

        avg_loss = total_loss / num_batches
        accuracy = 100. * correct / len(train_loader.dataset)
        print(f'Epoch {epoch}: Loss {avg_loss:.4f}, Acc {accuracy:.2f}%')
        print(f'Operator usage: {epoch_op_counts}')

        optimizer.temperature = max(0.5, optimizer.temperature * 0.95)

    return model

def test(model, device, test_loader):
    model.eval()
    correct = 0
    with torch.no_grad():
        for data, target in test_loader:
            data, target = data.to(device), target.to(device)
            output = model(data)
            pred = output.argmax(dim=1)
            correct += pred.eq(target).sum().item()
    acc = 100. * correct / len(test_loader.dataset)
    print(f'Test accuracy: {acc:.2f}%')
    return acc

# ----------------------------------------------------------------------
# Main
# ----------------------------------------------------------------------
def main():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"Using {device}")
    
    # Data
    transform_train = transforms.Compose([
        transforms.RandomCrop(32, padding=4),
        transforms.RandomHorizontalFlip(),
        transforms.ToTensor(),
        transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)),
    ])
    transform_test = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)),
    ])
    trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform_train)
    testset = torchvision.datasets.CIFAR10(root='./data', train=False, download=True, transform=transform_test)
    train_loader = torch.utils.data.DataLoader(trainset, batch_size=128, shuffle=True, num_workers=2)
    test_loader = torch.utils.data.DataLoader(testset, batch_size=100, shuffle=False, num_workers=2)
    
    # Model and policy
    model = SimpleCNN().to(device)
    policy = OperatorPolicy(state_dim=8, hidden_dim=32).to(device)
    
    print("Training with AI-automaton selecting update operators...")
    train_adaptive(model, policy, device, train_loader, epochs=20, lr=0.01)
    test(model, device, test_loader)

if __name__ == '__main__':
    main()