import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
import numpy as np
from itertools import chain


# ============================================
# 4-Layer MLP with Momentum (No Conv)
# ============================================

class MomentumMLP(nn.Module):
    """
    4 layers: 2 momentum layers (trained on replay) + 2 fresh layers
    """
    def __init__(self, input_size=3072, hidden_size=256, num_classes=10):
        super().__init__()
        
        # 4 layers total
        self.layer1 = nn.Linear(input_size, hidden_size)  # Momentum layer 1
        self.layer2 = nn.Linear(hidden_size, hidden_size)  # Momentum layer 2
        self.layer3 = nn.Linear(hidden_size, hidden_size)  # Fresh layer 1
        self.layer4 = nn.Linear(hidden_size, num_classes)  # Fresh layer 2
        
        self.relu = nn.ReLU()
        
    def forward(self, x, use_momentum_layers=True):
        x = x.view(x.size(0), -1)  # Flatten
        
        if use_momentum_layers:
            # Momentum path: layers 1-2 with accumulated knowledge
            x = self.relu(self.layer1(x))
            x = self.relu(self.layer2(x))
            # Fresh path: layers 3-4
            x = self.relu(self.layer3(x))
        else:
            # Fresh only path (when no replay available)
            x = self.relu(self.layer1(x))
            x = self.relu(self.layer2(x))
            x = self.relu(self.layer3(x))
            
        return self.layer4(x)


class ReplayBuffer:
    """Simple buffer for storing previous samples"""
    def __init__(self, capacity=5000):
        self.capacity = capacity
        self.buffer = []
        self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    
    def add(self, samples, labels):
        for s, l in zip(samples, labels):
            if len(self.buffer) < self.capacity:
                self.buffer.append((s.clone(), l.clone()))
    
    def sample(self, batch_size):
        if len(self.buffer) == 0:
            return None, None
        indices = np.random.choice(len(self.buffer), min(batch_size, len(self.buffer)), replace=False)
        samples = torch.stack([self.buffer[i][0] for i in indices]).to(self.device)
        labels = torch.stack([self.buffer[i][1] for i in indices]).to(self.device)
        return samples, labels


def train_momentum_model():
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    print(f"Using device: {device}")
    
    # CIFAR-10 transforms (no conv, so flatten to 3072)
    transform = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
    ])
    
    train_dataset = datasets.CIFAR10(root='../data', train=True, transform=transform, download=True)
    test_dataset = datasets.CIFAR10(root='../data', train=False, transform=transform, download=True)
    
    train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True, num_workers=2)
    test_loader = DataLoader(test_dataset, batch_size=256, shuffle=False, num_workers=2)
    
    # === Model with Momentum ===
    model = MomentumMLP(input_size=3072, hidden_size=512, num_classes=10).to(device)
    
    # Two optimizers: momentum optimizer (slower lr, less updates)
    # and fresh optimizer (faster learning on new data)
    momentum_optimizer = optim.Adam(model.layer1.parameters(), lr=0.0005)  # Conservative
    #fresh_optimizer = optim.Adam([model.layer2.parameters(),model.layer3.parameters(),model.layer4.parameters()], lr=0.001)
    # Correct way to combine multiple parameter generators
    params_to_update = chain(model.layer2.parameters(), 
                             model.layer3.parameters(), 
                             model.layer4.parameters())

    fresh_optimizer = optim.Adam(params_to_update, lr=0.001)
    
    replay_buffer = ReplayBuffer(capacity=10000)
    criterion = nn.CrossEntropyLoss()
    
    print("\n" + "="*50)
    print("Training 4-Layer Momentum MLP on CIFAR-10")
    print("="*50)
    
    epochs = 20
    for epoch in range(epochs):
        model.train()
        train_loss = 0
        correct = 0
        total = 0
        
        for batch_idx, (inputs, targets) in enumerate(train_loader):
            inputs, targets = inputs.to(device), targets.to(device)
            
            # 1. Train fresh layers on current batch (main training)
            fresh_optimizer.zero_grad()
            outputs = model(inputs, use_momentum_layers=True)
            loss_fresh = criterion(outputs, targets)
            loss_fresh.backward()
            fresh_optimizer.step()
            #print(loss_fresh.item())            
            # 2. Sample from replay buffer and train momentum layers
            replay_inputs, replay_targets = replay_buffer.sample(batch_size=64)
            if replay_inputs is not None:
                # Train layer1 on replay (accumulated knowledge)
                momentum_optimizer.zero_grad()
                x = inputs.view(inputs.size(0), -1)
                #x = torch.cat([inputs.view(inputs.size(0), -1), 
                #               replay_inputs], dim=0)
                
                # Forward through momentum layers only
                h = torch.relu(model.layer1(x))
                h = torch.relu(model.layer2(h))
                
                # Combine with fresh features
                loss_momentum = criterion(model.layer4(
                    torch.relu(model.layer3(h[-inputs.size(0):]))
                ), targets)
                loss_momentum.backward()
                momentum_optimizer.step()
                #print(loss_momentum.item())
            # 3. Add current batch to replay buffer
            replay_buffer.add(inputs.detach().cpu(), targets.detach().cpu())
            
            # Stats
            train_loss += loss_fresh.item()
            _, predicted = outputs.max(1)
            total += targets.size(0)
            correct += predicted.eq(targets).sum().item()
        
        # Evaluate
        model.eval()
        test_loss = 0
        test_correct = 0
        test_total = 0
        
        with torch.no_grad():
            for inputs, targets in test_loader:
                inputs, targets = inputs.to(device), targets.to(device)
                outputs = model(inputs, use_momentum_layers=True)
                loss = criterion(outputs, targets)
                test_loss += loss.item()
                _, predicted = outputs.max(1)
                test_total += targets.size(0)
                test_correct += predicted.eq(targets).sum().item()
        
        print(f"Epoch {epoch+1:2d}/{epochs} | "
              f"Train Acc: {100.*correct/total:.2f}% | "
              f"Test Acc: {100.*test_correct/test_total:.2f}% | "
              f"Replay: {len(replay_buffer.buffer)}")
    
    return model


# === Run ===
if __name__ == "__main__":
    model = train_momentum_model()
