"""
CIFAR-10 Kaggle比赛Baseline
包含：5折交叉验证、TTA、Label Smoothing、CosineAnnealingLR、混合精度训练
直接跑就能出提交结果，大概95%+准确率
"""
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, Dataset, Subset
from torchvision import datasets, transforms, models
from sklearn.model_selection import KFold
import numpy as np
import pandas as pd
from tqdm import tqdm
import random

# 固定随机种子，保证可复现
def seed_everything(seed=42):
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)
    torch.backends.cudnn.deterministic = True
seed_everything()

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(f'使用设备: {device}')

# ==================== 超参数 ====================
BATCH_SIZE = 128
EPOCHS = 20
LR = 1e-3
NUM_CLASSES = 10
K_FOLDS = 5
WEIGHT_DECAY = 1e-4

# ==================== 数据增强 ====================
train_transform = 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)),  # CIFAR-10统计值
])
test_transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)),
])

# 加载数据集
train_dataset = datasets.CIFAR10(root='./data', train=True, download=True, transform=train_transform)
test_dataset = datasets.CIFAR10(root='./data', train=False, download=True, transform=test_transform)

# ==================== 训练/验证函数 ====================
def train_one_epoch(model, loader, criterion, optimizer, scaler):
    model.train()
    total_loss = 0
    correct = 0
    total = 0
    for imgs, labels in tqdm(loader, desc='训练'):
        imgs, labels = imgs.to(device), labels.to(device)
        optimizer.zero_grad()
        # 混合精度训练加速
        with torch.cuda.amp.autocast():
            outputs = model(imgs)
            loss = criterion(outputs, labels)
        scaler.scale(loss).backward()
        scaler.step(optimizer)
        scaler.update()
        total_loss += loss.item() * imgs.size(0)
        _, predicted = outputs.max(1)
        total += labels.size(0)
        correct += predicted.eq(labels).sum().item()
    return total_loss / total, correct / total

def val_one_epoch(model, loader, criterion):
    model.eval()
    total_loss = 0
    correct = 0
    total = 0
    with torch.no_grad():
        for imgs, labels in tqdm(loader, desc='验证'):
            imgs, labels = imgs.to(device), labels.to(device)
            outputs = model(imgs)
            loss = criterion(outputs, labels)
            total_loss += loss.item() * imgs.size(0)
            _, predicted = outputs.max(1)
            total += labels.size(0)
            correct += predicted.eq(labels).sum().item()
    return total_loss / total, correct / total

# TTA预测
def tta_predict(model, img):
    model.eval()
    preds = []
    with torch.no_grad():
        # 原图
        preds.append(torch.softmax(model(img), dim=1))
        # 水平翻转
        preds.append(torch.softmax(model(torch.flip(img, dims=[-1])), dim=1))
    return torch.stack(preds).mean(dim=0)

# ==================== K折训练 ====================
kf = KFold(n_splits=K_FOLDS, shuffle=True, random_state=42)
all_test_preds = []  # 存所有折的测试预测结果
best_accs = []

for fold, (train_idx, val_idx) in enumerate(kf.split(train_dataset)):
    print(f'\n===== 第 {fold+1}/{K_FOLDS} 折训练 =====')
    # 划分训练/验证集
    train_subset = Subset(train_dataset, train_idx)
    val_subset = Subset(train_dataset, val_idx)
    # 验证集用测试transform（不做增强）
    val_subset.dataset.transform = test_transform
    train_loader = DataLoader(train_subset, batch_size=BATCH_SIZE, shuffle=True, num_workers=4, pin_memory=True)
    val_loader = DataLoader(val_subset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4, pin_memory=True)
    test_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4, pin_memory=True)
    # 模型：ResNet18，改最后一层输出10类
    model = models.resnet18(pretrained=False)
    model.conv1 = nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1, bias=False)  # 适配32x32输入
    model.maxpool = nn.Identity()
    model.fc = nn.Linear(model.fc.in_features, NUM_CLASSES)
    model = model.to(device)
    # 带Label Smoothing的损失函数
    criterion = nn.CrossEntropyLoss(label_smoothing=0.1)
    optimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)
    # 余弦退火学习率
    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)
    scaler = torch.cuda.amp.GradScaler()
    best_acc = 0
    for epoch in range(EPOCHS):
        train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, scaler)
        val_loss, val_acc = val_one_epoch(model, val_loader, criterion)
        scheduler.step()
        print(f'Epoch {epoch+1}/{EPOCHS} | 训练Loss: {train_loss:.4f}  Acc: {train_acc:.4f} | 验证Loss: {val_loss:.4f} Acc: {val_acc:.4f}')
        # 保存最佳模型
        if val_acc > best_acc:
            best_acc = val_acc
            torch.save(model.state_dict(), f'best_fold{fold}.pth')
            print(f'验证集准确率提升，已保存最佳模型: {best_acc:.4f}')
    best_accs.append(best_acc)
    # 加载最佳模型做TTA预测
    model.load_state_dict(torch.load(f'best_fold{fold}.pth'))
    fold_preds = []
    for imgs, _ in tqdm(test_loader, desc='TTA测试'):
        imgs = imgs.to(device)
        pred = tta_predict(model, imgs)
        fold_preds.append(pred.cpu().numpy())
    fold_preds = np.concatenate(fold_preds, axis=0)
    all_test_preds.append(fold_preds)
print(f'\n===== K折训练完成 =====')
print(f'各折最佳准确率: {[f"{acc:.4f}" for acc in best_accs]}')
print(f'平均验证准确率: {np.mean(best_accs):.4f}')

# ==================== 集成所有折的预测结果 ====================
final_preds = np.mean(all_test_preds, axis=0)  # K折结果平均
final_labels = np.argmax(final_preds, axis=1)
# 生成Kaggle提交文件
submission = pd.DataFrame({
    'id': np.arange(len(final_labels)),
    'label': final_labels
})
submission.to_csv('submission.csv', index=False)
print('提交文件 submission.csv 已生成！可以直接上传Kaggle~')
