{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":21154,"databundleVersionId":1243559,"sourceType":"competition"}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Flower Classification with EfficientNet (PyTorch + TFRecords)\n\nThis notebook reimplements the original training pipeline in a Kaggle-friendly format. It keeps the project modular, optimizes data loading for TFRecord files, and enables fast experimentation on the `tpu-getting-started` dataset.\n\n","metadata":{"execution":{"iopub.status.busy":"2025-11-26T08:56:56.299140Z","iopub.execute_input":"2025-11-26T08:56:56.299717Z","iopub.status.idle":"2025-11-26T08:56:56.304617Z","shell.execute_reply.started":"2025-11-26T08:56:56.299695Z","shell.execute_reply":"2025-11-26T08:56:56.303677Z"}}},{"cell_type":"markdown","source":"## 安装环境依赖\n在 Kaggle Notebook 初次运行时，需要提前安装 `tfrecord` 与 `timm` 等第三方库，以便后续读取 TFRecord 数据并加载 EfficientNet 模型。\n","metadata":{}},{"cell_type":"code","source":"pip install -q tfrecord timm==0.9.2","metadata":{"vscode":{"languageId":"plaintext"},"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T12:14:03.045970Z","iopub.execute_input":"2025-11-26T12:14:03.046288Z","iopub.status.idle":"2025-11-26T12:15:21.242153Z","shell.execute_reply.started":"2025-11-26T12:14:03.046262Z","shell.execute_reply":"2025-11-26T12:15:21.241425Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 导入核心依赖与工具函数\n该部分导入 PyTorch、Torchvision、timm、tfrecord 等训练所需库，同时加载常用的科学计算与可视化工具，为后续数据管道与模型构建做准备。\n","metadata":{}},{"cell_type":"code","source":"import io\nimport math\nimport os\nimport random\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom PIL import Image\n\nimport torch\nfrom torch import nn\nfrom torch.utils.data import IterableDataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torchvision import transforms\n\nfrom sklearn.metrics import f1_score, precision_score, recall_score, confusion_matrix\n\nimport timm\nimport tfrecord\n\nfrom tqdm.auto import tqdm\n\nplt.style.use('seaborn-v0_8')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T12:15:21.243619Z","iopub.execute_input":"2025-11-26T12:15:21.243862Z","iopub.status.idle":"2025-11-26T12:15:31.675361Z","shell.execute_reply.started":"2025-11-26T12:15:21.243836Z","shell.execute_reply":"2025-11-26T12:15:31.674771Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 配置参数与随机种子设置\n这里集中定义训练所需的超参数、数据路径以及类别名称，并设置随机种子、创建工作目录，确保在 Kaggle 环境中的可复现性。\n","metadata":{}},{"cell_type":"code","source":"class CFG:\n    base_path = Path('/kaggle/input/tpu-getting-started')\n    tfrecord_dir = base_path / 'tfrecords-jpeg-512x512'\n    train_dir = tfrecord_dir / 'train'\n    val_dir = tfrecord_dir / 'val'\n    test_dir = tfrecord_dir / 'test'\n\n    image_size = (384, 384)\n    num_classes = 104\n\n    batch_size = 16\n    val_batch_size = 64\n    test_batch_size = 64\n    epochs = 10\n\n    learning_rate = 3e-4\n    min_lr = 1e-6\n    weight_decay = 1e-5\n    grad_clip = 1.0\n\n    num_workers = 1\n    seed = 2235\n\n    model_name = 'swin_large_patch4_window12_384'\n    drop_rate = 0.214\n    drop_path_rate = 0.118\n    use_amp = True\n\n    output_dir = Path('/kaggle/working')\n    model_path = output_dir / 'best_model.pt'\n    submission_path = output_dir / 'submission.csv'\n\n    train_limit = None  # 调试时可以设置为例如 2048\n    val_limit = None\n    test_limit = None\n\n    classes = [\n        'pink primrose', 'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea', 'wild geranium',\n        'tiger lily', 'moon orchid', 'bird of paradise', 'monkshood', 'globe thistle',\n        'snapdragon', \"colt's foot\", 'king protea', 'spear thistle', 'yellow iris',\n        'globe-flower', 'purple coneflower', 'peruvian lily', 'balloon flower', 'giant white arum lily',\n        'fire lily', 'pincushion flower', 'fritillary', 'red ginger', 'grape hyacinth',\n        'corn poppy', 'prince of wales feathers', 'stemless gentian', 'artichoke', 'sweet william',\n        'carnation', 'garden phlox', 'love in the mist', 'cosmos', 'alpine sea holly',\n        'ruby-lipped cattleya', 'cape flower', 'great masterwort', 'siam tulip', 'lenten rose',\n        'barberton daisy', 'daffodil', 'sword lily', 'poinsettia', 'bolero deep blue',\n        'wallflower', 'marigold', 'buttercup', 'daisy', 'common dandelion',\n        'petunia', 'wild pansy', 'primula', 'sunflower', 'lilac hibiscus',\n        'bishop of llandaff', 'gaura', 'geranium', 'orange dahlia', 'pink-yellow dahlia',\n        'cautleya spicata', 'japanese anemone', 'black-eyed susan', 'silverbush', 'californian poppy',\n        'osteospermum', 'spring crocus', 'iris', 'windflower', 'tree poppy',\n        'gazania', 'azalea', 'water lily', 'rose', 'thorn apple',\n        'morning glory', 'passion flower', 'lotus', 'toad lily', 'anthurium',\n        'frangipani', 'clematis', 'hibiscus', 'columbine', 'desert-rose',\n        'tree mallow', 'magnolia', 'cyclamen ', 'watercress', 'canna lily',\n        'hippeastrum ', 'bee balm', 'pink quill', 'foxglove', 'bougainvillea',\n        'camellia', 'mallow', 'mexican petunia', 'bromelia', 'blanket flower',\n        'trumpet creeper', 'blackberry lily', 'common tulip', 'wild rose'\n    ]\n\n\nCFG.output_dir.mkdir(parents=True, exist_ok=True)\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n\ndef set_seed(seed: int = CFG.seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n\nset_seed()\nprint(f'device: {device}')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T12:15:31.676208Z","iopub.execute_input":"2025-11-26T12:15:31.676469Z","iopub.status.idle":"2025-11-26T12:15:31.776196Z","shell.execute_reply.started":"2025-11-26T12:15:31.676444Z","shell.execute_reply":"2025-11-26T12:15:31.775391Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 多 GPU 使用说明\n若机器上有多块 GPU，可自动检测并启用 `nn.DataParallel`，如下修改将模型复制到所有可用 GPU 上并在训练、验证和推理阶段自动聚合结果。\n","metadata":{}},{"cell_type":"code","source":"# 自动检测 GPU 并设置随机种子\nCFG.output_dir.mkdir(parents=True, exist_ok=True)\nuse_cuda = torch.cuda.is_available()\nnum_gpus = torch.cuda.device_count() if use_cuda else 0\ndevice = torch.device('cuda' if use_cuda else 'cpu')\n\n\ndef set_seed(seed: int = CFG.seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if use_cuda:\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n\n\nset_seed()\nprint(f'device: {device}, available gpus: {num_gpus}')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T12:15:31.777973Z","iopub.execute_input":"2025-11-26T12:15:31.778244Z","iopub.status.idle":"2025-11-26T12:15:33.858798Z","shell.execute_reply.started":"2025-11-26T12:15:31.778224Z","shell.execute_reply":"2025-11-26T12:15:33.857993Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 训练期间的指标计算与可视化工具\n该部分实现滑动平均器、分类指标计算、训练曲线绘制以及混淆矩阵展示函数，便于监控模型表现并进行结果分析。\n","metadata":{}},{"cell_type":"code","source":"class AverageMeter:\n    \"\"\"跟踪并更新滑动平均。\"\"\"\n\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0.0\n        self.avg = 0.0\n        self.sum = 0.0\n        self.count = 0\n\n    def update(self, val: float, n: int = 1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count if self.count else 0.0\n\n\ndef accuracy_from_logits(logits: torch.Tensor, targets: torch.Tensor) -> float:\n    preds = logits.argmax(dim=1)\n    return (preds == targets).float().mean().item()\n\n\ndef compute_classification_metrics(preds: np.ndarray, labels: np.ndarray):\n    accuracy = (preds == labels).mean()\n    f1 = f1_score(labels, preds, labels=np.arange(CFG.num_classes), average='macro')\n    precision = precision_score(labels, preds, labels=np.arange(CFG.num_classes), average='macro')\n    recall = recall_score(labels, preds, labels=np.arange(CFG.num_classes), average='macro')\n    cmat = confusion_matrix(labels, preds, labels=np.arange(CFG.num_classes))\n    return accuracy, f1, precision, recall, cmat\n\n\ndef plot_training_curves(history):\n    fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\n    axes[0].plot(history['train_loss'], label='train')\n    axes[0].plot(history['val_loss'], label='val')\n    axes[0].set_title('Loss')\n    axes[0].set_xlabel('epoch')\n    axes[0].legend()\n\n    axes[1].plot(history['train_acc'], label='train')\n    axes[1].plot(history['val_acc'], label='val')\n    axes[1].set_title('Accuracy')\n    axes[1].set_xlabel('epoch')\n    axes[1].legend()\n\n    axes[2].plot(history['lr'])\n    axes[2].set_title('Learning Rate')\n    axes[2].set_xlabel('epoch')\n\n    plt.show()\n\n\ndef show_confusion_matrix(cmat: np.ndarray, title_suffix: str = ''):\n    fig, ax = plt.subplots(figsize=(14, 12))\n    cax = ax.imshow(cmat, interpolation='nearest', cmap='Reds')\n    fig.colorbar(cax)\n    ax.set_title(f'Confusion Matrix {title_suffix}')\n    ax.set_xlabel('Predicted')\n    ax.set_ylabel('True')\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T12:15:33.860267Z","iopub.execute_input":"2025-11-26T12:15:33.860549Z","iopub.status.idle":"2025-11-26T12:15:33.870099Z","shell.execute_reply.started":"2025-11-26T12:15:33.860525Z","shell.execute_reply":"2025-11-26T12:15:33.869344Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## TFRecord 数据集与 DataLoader 构建\n这里实现对 TFRecord 文件的遍历、随机顺序控制与图像预处理，并封装成可复用的数据加载器，确保在 Kaggle 上按需读取数据且控制显存占用。\n","metadata":{}},{"cell_type":"code","source":"import re\n\n\ndef get_tfrecord_files(directory: Path):\n    files = sorted(str(p) for p in directory.glob('*.tfrec'))\n    if not files:\n        raise FileNotFoundError(f'No TFRecord files found in {directory}')\n    return files\n\n\ndef count_data_items(filenames):\n    pattern = re.compile(r'-([0-9]*)\\.tfrec$')\n    total = 0\n    for fname in filenames:\n        match = pattern.search(fname)\n        if match:\n            total += int(match.group(1))\n    return total\n\ndef get_train_transforms():\n    return transforms.Compose([\n        transforms.RandomResizedCrop(CFG.image_size, scale=(0.7, 1.0), ratio=(0.8, 1.2)),\n        transforms.RandomHorizontalFlip(p=0.618),\n        transforms.RandomVerticalFlip(p=0.314),\n        transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.05),\n        transforms.RandomRotation(degrees=20),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n        transforms.RandomErasing(p=0.2, scale=(0.02, 0.1), ratio=(0.3, 3.3), value='random'),\n    ])\n\n\ndef get_eval_transforms():\n    return transforms.Compose([\n        transforms.Resize(CFG.image_size),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    ])\n\n\nclass FlowerTFRecordDataset(IterableDataset):\n    def __init__(self, filenames, labeled=True, transform=None, seed=CFG.seed, limit=None):\n        self.filenames = list(filenames)\n        self.labeled = labeled\n        self.transform = transform\n        self.seed = seed\n        self.limit = limit\n        self.epoch = 0\n        total_items = count_data_items(self.filenames)\n        self.num_samples = min(total_items, limit) if limit is not None else total_items\n\n    def set_epoch(self, epoch: int):\n        self.epoch = epoch\n\n    def __len__(self):\n        return self.num_samples\n\n    def _shuffle_files(self):\n        rng = random.Random(self.seed + self.epoch)\n        files = self.filenames.copy()\n        rng.shuffle(files)\n        return files\n\n    def _shuffle_records(self, records, *, fname):\n        rng = random.Random(self.seed + self.epoch + hash(fname) % 10_000)\n        records = list(records)\n        rng.shuffle(records)\n        return records\n\n    def __iter__(self):\n        files = self._shuffle_files()\n        worker_info = torch.utils.data.get_worker_info()\n        if worker_info is not None:\n            files = files[worker_info.id :: worker_info.num_workers]\n\n        produced = 0\n        for fname in files:\n            if self.limit is not None and produced >= self.limit:\n                break\n\n            try:\n                records = tfrecord.tfrecord_loader(fname, None, None)\n            except Exception as exc:\n                print(f'Warning: failed to read {fname}: {exc}')\n                continue\n\n            for record in self._shuffle_records(records, fname=fname):\n                if self.limit is not None and produced >= self.limit:\n                    break\n\n                image_bytes = record['image']\n                image = Image.open(io.BytesIO(image_bytes)).convert('RGB')\n                if self.transform is not None:\n                    image = self.transform(image)\n\n                if self.labeled:\n                    label_raw = record['class']\n                    if isinstance(label_raw, np.ndarray):\n                        label = int(label_raw.reshape(-1)[0])\n                    elif isinstance(label_raw, (list, tuple)):\n                        label = int(label_raw[0])\n                    else:\n                        label = int(label_raw)\n                    yield image, label\n                else:\n                    image_id = record.get('id', '')\n                    if isinstance(image_id, bytes):\n                        image_id = image_id.decode()\n                    yield image, image_id\n\n                produced += 1\n\n\ndef build_loader(directory: Path, *, labeled: bool, transform, batch_size: int, limit=None):\n    filenames = get_tfrecord_files(directory)\n    dataset = FlowerTFRecordDataset(filenames, labeled=labeled, transform=transform, limit=limit)\n\n    loader = DataLoader(\n        dataset,\n        batch_size=batch_size,\n        shuffle=False,  # 数据集内部已经处理了随机逻辑\n        num_workers=CFG.num_workers,\n        pin_memory=True,\n        drop_last=False,\n    )\n\n    return loader, dataset\n\n\ndef get_loaders():\n    train_loader, _ = build_loader(\n        CFG.train_dir,\n        labeled=True,\n        transform=get_train_transforms(),\n        batch_size=CFG.batch_size,\n        limit=CFG.train_limit,\n    )\n\n    val_loader, _ = build_loader(\n        CFG.val_dir,\n        labeled=True,\n        transform=get_eval_transforms(),\n        batch_size=CFG.val_batch_size,\n        limit=CFG.val_limit,\n    )\n\n    return train_loader, val_loader\n\n\ndef get_test_loader():\n    test_loader, _ = build_loader(\n        CFG.test_dir,\n        labeled=False,\n        transform=get_eval_transforms(),\n        batch_size=CFG.test_batch_size,\n        limit=CFG.test_limit,\n    )\n    return test_loader\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T12:15:33.870938Z","iopub.execute_input":"2025-11-26T12:15:33.871155Z","iopub.status.idle":"2025-11-26T12:15:33.888660Z","shell.execute_reply.started":"2025-11-26T12:15:33.871139Z","shell.execute_reply":"2025-11-26T12:15:33.887900Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 模型、优化器与学习率调度器\n在此处创建使用 `timm` 的 EfficientNet 骨干网络，并配置 `AdamW` 优化器与余弦退火调度器，为后续训练循环提供基础组件。\n","metadata":{}},{"cell_type":"code","source":"def build_model():\n    model = timm.create_model(\n        CFG.model_name,\n        pretrained=True,\n        num_classes=CFG.num_classes,\n        drop_rate=CFG.drop_rate,\n        drop_path_rate=CFG.drop_path_rate,\n    )\n    return model\n\n\ndef create_optimizer(model):\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=CFG.learning_rate,\n        weight_decay=CFG.weight_decay,\n    )\n    return optimizer\n\n\ndef create_scheduler(optimizer):\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer,\n        T_max=CFG.epochs,\n        eta_min=CFG.min_lr,\n    )\n    return scheduler\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T12:15:33.889283Z","iopub.execute_input":"2025-11-26T12:15:33.889537Z","iopub.status.idle":"2025-11-26T12:15:33.905341Z","shell.execute_reply.started":"2025-11-26T12:15:33.889495Z","shell.execute_reply":"2025-11-26T12:15:33.904652Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 训练与验证主循环\n该段代码实现单个 epoch 的前向/反向过程、AMP 混合精度训练以及整体训练流程的历史记录与最佳模型保存逻辑。\n","metadata":{}},{"cell_type":"code","source":"def run_epoch(loader, model, criterion, optimizer=None, scaler=None, epoch=0, training=True):\n    if training and hasattr(loader.dataset, 'set_epoch'):\n        loader.dataset.set_epoch(epoch)\n\n    if training:\n        model.train()\n    else:\n        model.eval()\n\n    losses = AverageMeter()\n    accuracies = AverageMeter()\n\n    total_steps = math.ceil(len(loader.dataset) / loader.batch_size)\n\n    pbar = tqdm(loader, total=total_steps, leave=False)\n    pbar.set_description(f\"Epoch {epoch + 1} [{'train' if training else 'val'}]\")\n\n    for images, targets in pbar:\n        images = images.to(device, non_blocking=True)\n        targets = targets.to(device, non_blocking=True)\n\n        with torch.set_grad_enabled(training):\n            with autocast(enabled=scaler is not None and scaler.is_enabled()):\n                outputs = model(images)\n                loss = criterion(outputs, targets)\n\n        batch_size = targets.size(0)\n        if training and optimizer is not None:\n            optimizer.zero_grad(set_to_none=True)\n            if scaler is not None and scaler.is_enabled():\n                scaler.scale(loss).backward()\n                torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.grad_clip)\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                loss.backward()\n                torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.grad_clip)\n                optimizer.step()\n\n        acc = accuracy_from_logits(outputs.detach(), targets)\n        losses.update(loss.item(), batch_size)\n        accuracies.update(acc, batch_size)\n\n        pbar.set_postfix({\n            'loss': f'{losses.avg:.4f}',\n            'acc': f'{accuracies.avg:.4f}',\n            'lr': optimizer.param_groups[0]['lr'] if optimizer is not None else 0.0,\n        })\n\n    return losses.avg, accuracies.avg\n\n\ndef train_model():\n    train_loader, val_loader = get_loaders()\n\n    model = build_model().to(device)\n    if use_cuda and num_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f'Using DataParallel on {num_gpus} GPUs')\n    criterion = nn.CrossEntropyLoss()\n    optimizer = create_optimizer(model)\n    scheduler = create_scheduler(optimizer)\n    scaler = GradScaler(enabled=CFG.use_amp and device.type == 'cuda')\n\n    history = {\n        'train_loss': [],\n        'train_acc': [],\n        'val_loss': [],\n        'val_acc': [],\n        'lr': [],\n    }\n\n    best_val_acc = 0.0\n\n    for epoch in range(CFG.epochs):\n        train_loss, train_acc = run_epoch(\n            train_loader,\n            model,\n            criterion,\n            optimizer=optimizer,\n            scaler=scaler,\n            epoch=epoch,\n            training=True,\n        )\n\n        val_loss, val_acc = run_epoch(\n            val_loader,\n            model,\n            criterion,\n            optimizer=None,\n            scaler=None,\n            epoch=epoch,\n            training=False,\n        )\n\n        scheduler.step()\n        current_lr = optimizer.param_groups[0]['lr']\n\n        history['train_loss'].append(train_loss)\n        history['train_acc'].append(train_acc)\n        history['val_loss'].append(val_loss)\n        history['val_acc'].append(val_acc)\n        history['lr'].append(current_lr)\n\n        print(f\"Epoch {epoch + 1}/{CFG.epochs}: train_loss={train_loss:.4f}, train_acc={train_acc:.4f}, val_loss={val_loss:.4f}, val_acc={val_acc:.4f}, lr={current_lr:.2e}\")\n\n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n            model_to_save = model.module if isinstance(model, nn.DataParallel) else model\n            torch.save({'model_state_dict': model_to_save.state_dict(), 'val_acc': best_val_acc}, CFG.model_path)\n            print(f\"  ✓ New best model saved (val_acc={best_val_acc:.4f})\")\n\n    return model, history, best_val_acc\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T12:15:33.906089Z","iopub.execute_input":"2025-11-26T12:15:33.906382Z","iopub.status.idle":"2025-11-26T12:15:33.922696Z","shell.execute_reply.started":"2025-11-26T12:15:33.906362Z","shell.execute_reply":"2025-11-26T12:15:33.921972Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 验证评估与提交文件生成\n该部分提供模型验证评估、计算指标以及在测试集上生成预测并导出提交 CSV 的工具函数。\n","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef evaluate_model(model, loader):\n    model.eval()\n    all_preds = []\n    all_labels = []\n\n    pbar = tqdm(loader, leave=False)\n    pbar.set_description('Evaluating')\n\n    for images, labels in pbar:\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        outputs = model(images)\n        preds = outputs.argmax(dim=1)\n\n        all_preds.append(preds.cpu().numpy())\n        all_labels.append(labels.cpu().numpy())\n\n    preds = np.concatenate(all_preds)\n    labels = np.concatenate(all_labels)\n\n    accuracy, f1, precision, recall, cmat = compute_classification_metrics(preds, labels)\n    return preds, labels, accuracy, f1, precision, recall, cmat\n\n\n@torch.no_grad()\ndef generate_submission(model, loader):\n    model.eval()\n    all_preds = []\n    all_ids = []\n\n    pbar = tqdm(loader, leave=False)\n    pbar.set_description('Predicting test')\n\n    for images, image_ids in pbar:\n        images = images.to(device, non_blocking=True)\n        outputs = model(images)\n        preds = outputs.argmax(dim=1).cpu().numpy()\n        all_preds.extend(preds.tolist())\n        all_ids.extend(image_ids)\n\n    submission = pd.DataFrame({'id': all_ids, 'label': all_preds})\n    submission.to_csv(CFG.submission_path, index=False)\n    print(f'Submission saved to {CFG.submission_path}')\n    return submission\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T12:15:33.923479Z","iopub.execute_input":"2025-11-26T12:15:33.923736Z","iopub.status.idle":"2025-11-26T12:15:33.937426Z","shell.execute_reply.started":"2025-11-26T12:15:33.923715Z","shell.execute_reply":"2025-11-26T12:15:33.936915Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 启动训练并绘制学习曲线\n执行该单元即可运行完整训练流程，输出损失与准确率、学习率的变化，并记录最佳验证精度。\n","metadata":{}},{"cell_type":"code","source":"%%time\n\nmodel, history, best_val_acc = train_model()\nplot_training_curves(history)\nprint(f'Best validation accuracy: {best_val_acc:.4f}')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T12:15:33.939212Z","iopub.execute_input":"2025-11-26T12:15:33.939642Z","iopub.status.idle":"2025-11-26T12:30:17.735606Z","shell.execute_reply.started":"2025-11-26T12:15:33.939619Z","shell.execute_reply":"2025-11-26T12:30:17.734842Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 加载最佳模型并查看验证集表现\n该单元会读取训练过程中保存的最佳权重，再次在验证集上计算综合指标并展示归一化混淆矩阵。\n","metadata":{}},{"cell_type":"code","source":"checkpoint = torch.load(CFG.model_path, map_location=device)\nbase_model = build_model()\nbase_model.load_state_dict(checkpoint['model_state_dict'])\nif use_cuda and num_gpus > 1:\n    model = nn.DataParallel(base_model).to(device)\nelse:\n    model = base_model.to(device)\nprint(f\"Loaded best checkpoint with val_acc={checkpoint['val_acc']:.4f}\")\n\n_, val_loader = get_loaders()\npreds, labels, acc, f1, precision, recall, cmat = evaluate_model(model, val_loader)\nprint(f'Validation accuracy: {acc:.4f}, F1: {f1:.4f}, precision: {precision:.4f}, recall: {recall:.4f}')\n\ncmat_norm = (cmat.T / cmat.sum(axis=1, keepdims=True)).T\nshow_confusion_matrix(cmat_norm, title_suffix='(normalized)')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T12:30:17.736650Z","iopub.execute_input":"2025-11-26T12:30:17.736894Z","iopub.status.idle":"2025-11-26T12:33:00.772828Z","shell.execute_reply.started":"2025-11-26T12:30:17.736875Z","shell.execute_reply":"2025-11-26T12:33:00.771999Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 生成测试集预测与提交文件\n运行该单元会在测试集上推理，输出提交所需的 `submission.csv` 文件，并预览前几行结果。\n","metadata":{}},{"cell_type":"code","source":"test_loader = get_test_loader()\nsubmission = generate_submission(model, test_loader)\nsubmission.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T12:33:00.774181Z","iopub.execute_input":"2025-11-26T12:33:00.774461Z","iopub.status.idle":"2025-11-26T12:38:26.408644Z","shell.execute_reply.started":"2025-11-26T12:33:00.774429Z","shell.execute_reply":"2025-11-26T12:38:26.407700Z"}},"outputs":[],"execution_count":null}]}