"""
快速版 CIFAR-10 对比实验 — epoch 少、batch 大，CPU 上也能跑完
"""

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
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import os
import time

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

# 数据
train_t_base = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,0.5,0.5),(0.5,0.5,0.5))])
train_t_tri  = 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_t = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,0.5,0.5),(0.5,0.5,0.5))])

def get_loaders(train_t, bs=256):
    full = datasets.CIFAR10(root='./data', train=True, download=False, transform=train_t)
    test = datasets.CIFAR10(root='./data', train=False, download=False, transform=test_t)
    tr, va = random_split(full, [int(0.9*len(full)), len(full)-int(0.9*len(full))])
    return DataLoader(tr, bs, shuffle=True, num_workers=0), DataLoader(va, bs, shuffle=False, num_workers=0), DataLoader(test, bs, shuffle=False, num_workers=0)

# 模型
class BaselineCNN(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.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))); x=self.pool(F.relu(self.conv2(x))); x=self.pool(F.relu(self.conv3(x)))
        x=x.view(x.size(0),-1); x=F.relu(self.fc1(x)); return self.fc2(x)

class TricksCNN(nn.Module):
    def __init__(self, drop=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.drop=nn.Dropout(drop); 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.drop(x); return self.fc2(x)

@torch.no_grad()
def evaluate(model, loader, crit):
    model.eval(); tl, c, t = 0, 0, 0
    for x,y in loader:
        x,y=x.to(device),y.to(device); o=model(x); l=crit(o,y)
        tl+=l.item()*x.size(0); c+=(o.argmax(1)==y).sum().item(); t+=x.size(0)
    return tl/t, c/t

def train_epoch(model, loader, crit, opt):
    model.train(); tl, c, t = 0, 0, 0
    for x,y in loader:
        x,y=x.to(device),y.to(device); opt.zero_grad(); o=model(x); l=crit(o,y); l.backward(); opt.step()
        tl+=l.item()*x.size(0); c+=(o.argmax(1)==y).sum().item(); t+=x.size(0)
    return tl/t, c/t

def run(name, model, tr_ld, va_ld, sched=False, wd=0.0, ep=15, pat=5):
    print(f"\n{'='*15} {name} {'='*15}")
    model=model.to(device); crit=nn.CrossEntropyLoss(); opt=torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=wd)
    sc = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=ep) if sched else None
    hist={'train_loss':[],'train_acc':[],'val_loss':[],'val_acc':[]}
    best, noimp, bst = float('inf'), 0, None
    for e in range(ep):
        t0=time.time()
        tl,ta=train_epoch(model,tr_ld,crit,opt); vl,va=evaluate(model,va_ld,crit)
        if sc: sc.step()
        for k,v in zip(hist.keys(),[tl,ta,vl,va]): hist[k].append(v)
        if vl<best: best, noimp, bst = vl, 0, model.state_dict().copy()
        else: noimp+=1
        print(f"Ep{e:02d} {time.time()-t0:.1f}s | train={tl:.4f}/{ta:.4f} val={vl:.4f}/{va:.4f}")
        if noimp>=pat:
            print(f"Early stop @ {e}"); break
    if bst: model.load_state_dict(bst)
    return model, hist

def plot(hb, ht):
    fig,axs=plt.subplots(1,2,figsize=(12,4.5))
    ax=axs[0]; ax.plot(hb['train_loss'],'b-',alpha=0.6,label='Base train'); ax.plot(hb['val_loss'],'b--',alpha=0.9,label='Base val')
    ax.plot(ht['train_loss'],'r-',alpha=0.6,label='Tricks train'); ax.plot(ht['val_loss'],'r--',alpha=0.9,label='Tricks val')
    ax.set_title('Loss'); ax.legend(); ax.grid(True,alpha=0.3)
    ax=axs[1]; ax.plot(hb['train_acc'],'b-',alpha=0.6); ax.plot(hb['val_acc'],'b--')
    ax.plot(ht['train_acc'],'r-',alpha=0.6); ax.plot(ht['val_acc'],'r--')
    ax.set_title('Accuracy'); ax.legend(['Base train','Base val','Tricks train','Tricks val']); ax.grid(True,alpha=0.3)
    plt.tight_layout()
    p=os.path.join(os.path.dirname(__file__),'training_tricks_comparison.png')
    plt.savefig(p,dpi=150); print(f"\nSaved: {p}"); plt.close()

if __name__=='__main__':
    bs=256; ep=15; pat=5
    trb,vab,teb=get_loaders(train_t_base,bs); trt,vat,tet=get_loaders(train_t_tri,bs)
    mb,hb=run('Baseline',BaselineCNN(),trb,vab,sched=False,wd=0.0,ep=ep,pat=pat)
    mt,ht=run('+Tricks',TricksCNN(0.3),trt,vat,sched=True,wd=1e-4,ep=ep,pat=pat)
    _,tb=evaluate(mb,teb,nn.CrossEntropyLoss()); _,tt=evaluate(mt,tet,nn.CrossEntropyLoss())
    print(f"\n{'='*40}\nBaseline test: {tb:.4f} ({tb*100:.2f}%)\n+Tricks test:  {tt:.4f} ({tt*100:.2f}%)\nImprovement: +{(tt-tb)*100:.2f}%\n{'='*40}")
    plot(hb,ht)
