"""
07 配套代码：CIFAR-10 训练技巧对比实验
Baseline vs +Dropout +BatchNorm +数据增强 +学习率调度 +Weight Decay
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader, random_split
from torchvision import datasets, transforms
import matplotlib.pyplot as plt
import os

# ───────────────── 设备 ─────────────────
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(f"Using device: {device}")

# ───────────────── 数据增强配置 ─────────────────
train_transform_baseline = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])

train_transform_tricks = transforms.Compose([
    transforms.RandomCrop(32, padding=4),
    transforms.RandomHorizontalFlip(p=0.5),
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])

test_transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])

# ───────────────── 加载 CIFAR-10 ─────────────────
def get_loaders(train_transform):
    train_full = datasets.CIFAR10(root='./data', train=True, download=True, transform=train_transform)
    test_set   = datasets.CIFAR10(root='./data', train=False, download=True, transform=test_transform)
    
    # 训练集拆出 10% 做验证
    train_size = int(0.9 * len(train_full))
    val_size   = len(train_full) - train_size
    train_set, val_set = random_split(train_full, [train_size, val_size])
    
    train_loader = DataLoader(train_set, batch_size=128, shuffle=True, num_workers=2)
    val_loader   = DataLoader(val_set,   batch_size=128, shuffle=False, num_workers=2)
    test_loader  = DataLoader(test_set,  batch_size=128, shuffle=False, num_workers=2)
    
    return train_loader, val_loader, test_loader

# ───────────────── 模型定义 ─────────────────
class BaselineCNN(nn.Module):
    """简单 CNN，无 BN 无 Dropout"""
    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.fc1 = nn.Linear(128 * 4 * 4, 256)
        self.fc2 = nn.Linear(256, 10)
    
    def forward(self, x):
        x = self.pool(F.relu(self.conv1(x)))   # [B,32,16,16]
        x = self.pool(F.relu(self.conv2(x)))   # [B,64,8,8]
        x = self.pool(F.relu(self.conv3(x)))   # [B,128,4,4]
        x = x.view(x.size(0), -1)
        x = F.relu(self.fc1(x))
        x = self.fc2(x)
        return x

class TricksCNN(nn.Module):
    """加 BN + Dropout + 更深一点"""
    def __init__(self, dropout=0.3):
        super().__init__()
        self.conv1 = nn.Conv2d(3, 32, 3, padding=1)
        self.bn1 = nn.BatchNorm2d(32)
        self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
        self.bn2 = nn.BatchNorm2d(64)
        self.conv3 = nn.Conv2d(64, 128, 3, padding=1)
        self.bn3 = nn.BatchNorm2d(128)
        self.pool = nn.MaxPool2d(2, 2)
        self.fc1 = nn.Linear(128 * 4 * 4, 256)
        self.dropout = nn.Dropout(dropout)
        self.fc2 = nn.Linear(256, 10)
    
    def forward(self, x):
        x = self.pool(F.relu(self.bn1(self.conv1(x))))
        x = self.pool(F.relu(self.bn2(self.conv2(x))))
        x = self.pool(F.relu(self.bn3(self.conv3(x))))
        x = x.view(x.size(0), -1)
        x = F.relu(self.fc1(x))
        x = self.dropout(x)
        x = self.fc2(x)
        return x

# ───────────────── 训练 & 验证 ─────────────────
def train_epoch(model, loader, criterion, optimizer):
    model.train()
    total_loss, correct, total = 0, 0, 0
    for x, y in loader:
        x, y = x.to(device), y.to(device)
        optimizer.zero_grad()
        out = model(x)
        loss = criterion(out, y)
        loss.backward()
        optimizer.step()
        
        total_loss += loss.item() * x.size(0)
        correct += (out.argmax(1) == y).sum().item()
        total += x.size(0)
    return total_loss / total, correct / total

@torch.no_grad()
def evaluate(model, loader, criterion):
    model.eval()
    total_loss, correct, total = 0, 0, 0
    for x, y in loader:
        x, y = x.to(device), y.to(device)
        out = model(x)
        loss = criterion(out, y)
        total_loss += loss.item() * x.size(0)
        correct += (out.argmax(1) == y).sum().item()
        total += x.size(0)
    return total_loss / total, correct / total

# ───────────────── 跑一组实验 ─────────────────
def run_experiment(name, model, train_loader, val_loader, use_scheduler=False, weight_decay=0.0, epochs=50, patience=10):
    print(f"\n{'='*20} {name} {'='*20}")
    model = model.to(device)
    criterion = nn.CrossEntropyLoss()
    optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=weight_decay)
    
    scheduler = None
    if use_scheduler:
        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)
    
    history = {'train_loss': [], 'train_acc': [], 'val_loss': [], 'val_acc': []}
    best_val_loss = float('inf')
    no_improve = 0
    best_state = None
    
    for epoch in range(epochs):
        t_loss, t_acc = train_epoch(model, train_loader, criterion, optimizer)
        v_loss, v_acc = evaluate(model, val_loader, criterion)
        
        if scheduler:
            scheduler.step()
        
        history['train_loss'].append(t_loss)
        history['train_acc'].append(t_acc)
        history['val_loss'].append(v_loss)
        history['val_acc'].append(v_acc)
        
        if v_loss < best_val_loss:
            best_val_loss = v_loss
            best_state = model.state_dict().copy()
            no_improve = 0
        else:
            no_improve += 1
        
        if epoch % 5 == 0 or epoch < 5:
            print(f"Epoch {epoch:02d} | train_loss={t_loss:.4f} train_acc={t_acc:.4f} | val_loss={v_loss:.4f} val_acc={v_acc:.4f}")
        
        if no_improve >= patience:
            print(f"Early stopping at epoch {epoch}")
            break
    
    # 加载最佳权重
    if best_state is not None:
        model.load_state_dict(best_state)
    
    return model, history

# ───────────────── 画对比图 ─────────────────
def plot_comparison(hist_baseline, hist_tricks):
    fig, axes = plt.subplots(1, 2, figsize=(14, 5))
    
    # Loss
    ax = axes[0]
    ax.plot(hist_baseline['train_loss'], 'b-', alpha=0.7, label='Baseline train')
    ax.plot(hist_baseline['val_loss'], 'b--', alpha=0.9, label='Baseline val')
    ax.plot(hist_tricks['train_loss'], 'r-', alpha=0.7, label='+Tricks train')
    ax.plot(hist_tricks['val_loss'], 'r--', alpha=0.9, label='+Tricks val')
    ax.set_xlabel('Epoch')
    ax.set_ylabel('Loss')
    ax.set_title('Loss: Baseline vs +Tricks')
    ax.legend()
    ax.grid(True, alpha=0.3)
    
    # Accuracy
    ax = axes[1]
    ax.plot(hist_baseline['train_acc'], 'b-', alpha=0.7, label='Baseline train')
    ax.plot(hist_baseline['val_acc'], 'b--', alpha=0.9, label='Baseline val')
    ax.plot(hist_tricks['train_acc'], 'r-', alpha=0.7, label='+Tricks train')
    ax.plot(hist_tricks['val_acc'], 'r--', alpha=0.9, label='+Tricks val')
    ax.set_xlabel('Epoch')
    ax.set_ylabel('Accuracy')
    ax.set_title('Accuracy: Baseline vs +Tricks')
    ax.legend()
    ax.grid(True, alpha=0.3)
    
    plt.tight_layout()
    out_path = os.path.join(os.path.dirname(__file__), 'training_tricks_comparison.png')
    plt.savefig(out_path, dpi=150)
    print(f"\n对比图已保存: {out_path}")
    plt.show()

# ───────────────── 主流程 ─────────────────
if __name__ == '__main__':
    # 准备两组数据加载器
    train_loader_base, val_loader_base, test_loader_base = get_loaders(train_transform_baseline)
    train_loader_tricks, val_loader_tricks, test_loader_tricks = get_loaders(train_transform_tricks)
    
    # Baseline 实验
    model_base, hist_base = run_experiment(
        'Baseline', BaselineCNN(), train_loader_base, val_loader_base,
        use_scheduler=False, weight_decay=0.0, epochs=50, patience=10
    )
    
    # +Tricks 实验
    model_tricks, hist_tricks = run_experiment(
        '+Tricks', TricksCNN(dropout=0.3), train_loader_tricks, val_loader_tricks,
        use_scheduler=True, weight_decay=1e-4, epochs=50, patience=10
    )
    
    # 最终测试
    criterion = nn.CrossEntropyLoss()
    _, test_acc_base = evaluate(model_base, test_loader_base, criterion)
    _, test_acc_tricks = evaluate(model_tricks, test_loader_tricks, criterion)
    
    print(f"\n{'='*50}")
    print(f"Baseline   测试准确率: {test_acc_base:.4f} ({test_acc_base*100:.2f}%)")
    print(f"+Tricks    测试准确率: {test_acc_tricks:.4f} ({test_acc_tricks*100:.2f}%)")
    print(f"提升: +{(test_acc_tricks - test_acc_base)*100:.2f}%")
    print(f"{'='*50}")
    
    # 画图
    plot_comparison(hist_base, hist_tricks)
