{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"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":6799,"databundleVersionId":4225553,"sourceType":"competition"}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader\nfrom torchvision import datasets, transforms, models\nimport matplotlib.pyplot as plt\nimport numpy as np\nfrom tqdm import tqdm\nimport os\nimport xml.etree.ElementTree as ET\nfrom pathlib import Path\nfrom PIL import Image\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# ====================== 1. 配置参数 =======================\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nBATCH_SIZE = 8\nEPOCHS = 10\nLEARNING_RATE = 1e-4\nMOMENTUM = 0.9\nWEIGHT_DECAY = 1e-4\nMODEL_NAMES = [\"alexnet\", \"vgg16\", \"resnet50\"]\n\n\ndef parse_xml_annotation(xml_path):\n    \"\"\"解析XML标注文件，获取类别ID\"\"\"\n    try:\n        tree = ET.parse(xml_path)\n        root = tree.getroot()\n        class_id = root.find(\"object\").find(\"name\").text\n        return class_id\n    except Exception as e:\n        print(f\"⚠️ 解析XML失败：{xml_path}，错误：{e}\")\n        return None\n\n\ndef get_image_ids_from_imageset(imageset_path, split=\"train\"):\n    \"\"\"从txt文件读取图像ID - 修复Kaggle格式解析（第二列是序号，不是标记）\"\"\"\n    image_ids = []\n    if not imageset_path.exists():\n        print(f\"❌ 图像ID文件不存在：{imageset_path}\")\n        return image_ids\n    \n    print(f\"📖 读取{split}集图像ID文件：{imageset_path}\")\n    with open(imageset_path, \"r\") as f:\n        lines = f.readlines()\n        print(f\"  - 共{len(lines)}行数据\")\n        \n        for idx, line in enumerate(lines[:5]):  # 打印前5行示例\n            print(f\"  - 示例行{idx+1}：{line.strip()}\")\n        \n        for line in lines:\n            parts = line.strip().split()\n            if len(parts) == 0:\n                continue\n            \n            # 关键修复：Kaggle格式中，第一列是图像ID，第二列是序号（忽略）\n            img_id_with_path = parts[0]\n            image_ids.append(img_id_with_path)\n    \n    print(f\"✅ 成功读取{split}集图像ID数量：{len(image_ids)}\")\n    return image_ids\n\n\n# ====================== 2. 数据预处理 =======================\ndef get_data_transforms():\n    train_transform = transforms.Compose([\n        transforms.RandomResizedCrop(224),\n        transforms.RandomHorizontalFlip(p=0.5),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n    ])\n\n    val_transform = transforms.Compose([\n        transforms.Resize(256),\n        transforms.CenterCrop(224),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n    ])\n    return train_transform, val_transform\n\n\nclass ILSVRC_Kaggle_Dataset(torch.utils.data.Dataset):\n    def __init__(self, image_ids_with_path, data_root, anno_root, split=\"train\", transform=None):\n        self.data_root = Path(data_root)\n        self.anno_root = Path(anno_root)\n        self.split = split\n        self.transform = transform\n        self.valid_image_ids = []  # 存储(图像路径, 标注路径, 类别ID)\n        self.class_to_idx = {}\n\n        # 验证根路径是否存在\n        print(f\"\\n📁 {split}集路径验证：\")\n        print(f\"  - 图像根路径：{self.data_root} {'✅ 存在' if self.data_root.exists() else '❌ 不存在'}\")\n        print(f\"  - 标注根路径：{self.anno_root} {'✅ 存在' if self.anno_root.exists() else '❌ 不存在'}\")\n        \n        if not self.data_root.exists() or not self.anno_root.exists():\n            raise FileNotFoundError(f\"{split}集图像或标注根路径不存在\")\n\n        # 收集所有类别（从标注文件中提取）\n        print(f\"\\n🔍 正在扫描{split}集标注文件，提取类别信息...\")\n        all_class_ids = set()\n        xml_files = list(self.anno_root.glob(\"**/*.xml\"))\n        print(f\"  - 标注文件夹下共有{len(xml_files)}个标注文件\")\n        \n        for xml_file in xml_files[:10]:  # 打印前10个标注文件示例\n            class_id = parse_xml_annotation(xml_file)\n            if class_id:\n                all_class_ids.add(class_id)\n        print(f\"  - 已发现类别数：{len(all_class_ids)}（示例：{list(all_class_ids)[:5]}）\")\n\n        # 构建类别映射\n        self.class_to_idx = {cls: idx for idx, cls in enumerate(sorted(all_class_ids))}\n        self.num_classes = len(self.class_to_idx)\n        print(f\"  - 类别映射构建完成：{self.num_classes}个类别\")\n\n        # 遍历所有图像ID，验证文件存在性\n        print(f\"\\n🔍 正在验证{split}集样本（前10个）：\")\n        valid_count = 0\n        for idx, img_id_with_path in enumerate(image_ids_with_path):\n            try:\n                if self.split == \"train\":\n                    # 训练集图像路径：data_root / 类别文件夹 / 图像ID.JPEG\n                    # 例如：n01440764/n01440764_10026 → data_root/n01440764/n01440764_10026.JPEG\n                    img_path = self.data_root / img_id_with_path\n                    \n                    # 尝试不同的图像后缀\n                    img_suffixes = [\".JPEG\", \".jpg\", \".png\", \".jpeg\"]\n                    found_img_path = None\n                    for suffix in img_suffixes:\n                        test_path = img_path.with_suffix(suffix)\n                        if test_path.exists():\n                            found_img_path = test_path\n                            break\n                    \n                    if not found_img_path:\n                        if idx < 10:\n                            print(f\"  - ❌ 图像不存在：{img_path}（尝试了多种后缀）\")\n                        continue\n                    \n                    # 训练集标注路径：anno_root / 图像ID.xml（关键修复：没有类别子文件夹！）\n                    xml_filename = f\"{os.path.basename(img_id_with_path)}.xml\"\n                    xml_path = self.anno_root / xml_filename\n                    \n                    if not xml_path.exists():\n                        # 备用路径：尝试带类别文件夹的标注路径\n                        xml_path = self.anno_root / img_id_with_path.replace(\"/\", \".\") + \".xml\"\n                        if not xml_path.exists():\n                            if idx < 10:\n                                print(f\"  - ❌ 标注不存在：{xml_path}\")\n                            continue\n\n                else:  # val\n                    # 验证集图像路径：data_root / 图像ID.JPEG\n                    img_id = img_id_with_path\n                    img_path = self.data_root / img_id\n                    \n                    # 尝试不同的图像后缀\n                    img_suffixes = [\".JPEG\", \".jpg\", \".png\", \".jpeg\"]\n                    found_img_path = None\n                    for suffix in img_suffixes:\n                        test_path = img_path.with_suffix(suffix)\n                        if test_path.exists():\n                            found_img_path = test_path\n                            break\n                    \n                    if not found_img_path:\n                        if idx < 10:\n                            print(f\"  - ❌ 图像不存在：{img_path}（尝试了多种后缀）\")\n                        continue\n                    \n                    # 验证集标注路径：anno_root / 图像ID.xml\n                    xml_path = self.anno_root / f\"{img_id}.xml\"\n                    if not xml_path.exists():\n                        if idx < 10:\n                            print(f\"  - ❌ 标注不存在：{xml_path}\")\n                        continue\n\n                # 解析类别ID\n                class_id = parse_xml_annotation(xml_path)\n                if class_id is None or class_id not in self.class_to_idx:\n                    if idx < 10:\n                        print(f\"  - ⚠️ 类别无效：{class_id}\")\n                    continue\n\n                # 存储有效样本信息\n                self.valid_image_ids.append((found_img_path, xml_path, class_id))\n                valid_count += 1\n                \n                # 打印前10个有效样本\n                if idx < 10 and valid_count <= 10:\n                    print(f\"  - ✅ 有效样本：{found_img_path.name} | 类别：{class_id}\")\n\n            except Exception as e:\n                if idx < 10:\n                    print(f\"  - ⚠️ 处理失败：{img_id_with_path}，错误：{str(e)[:50]}\")\n                continue\n\n        # 检查有效样本数\n        if not self.valid_image_ids:\n            print(f\"\\n❌ {split}集无有效样本！详细原因：\")\n            print(f\"  - 输入图像ID数量：{len(image_ids_with_path)}\")\n            print(f\"  - 图像根路径下的文件/文件夹数量：{len(list(self.data_root.glob('*')))}\")\n            print(f\"  - 标注根路径下的XML文件数量：{len(list(self.anno_root.glob('*.xml')))}\")\n            raise ValueError(f\"{split}集无有效样本，请检查路径和数据集格式\")\n        \n        print(f\"\\n✅ {split}集加载完成：\")\n        print(f\"  - 有效样本数：{len(self.valid_image_ids)}\")\n        print(f\"  - 类别数：{self.num_classes}\")\n        print(f\"  - 类别示例：{list(self.class_to_idx.keys())[:5]}\")\n\n    def __len__(self):\n        return len(self.valid_image_ids)\n\n    def __getitem__(self, idx):\n        img_path, xml_path, class_id = self.valid_image_ids[idx]\n\n        # 读取图像\n        try:\n            image = Image.open(img_path).convert(\"RGB\")\n        except Exception as e:\n            raise FileNotFoundError(f\"无法读取图像：{img_path}，错误：{e}\")\n\n        # 获取标签\n        label = self.class_to_idx[class_id]\n\n        if self.transform:\n            image = self.transform(image)\n        return image, label\n\n\ndef load_kaggle_ilsvrc_dataset(train_transform, val_transform):\n    # Kaggle ImageNet数据集的标准路径\n    KAGGLE_DATA_ROOT = Path(\"/kaggle/input/imagenet-object-localization-challenge/ILSVRC\")\n    print(f\"🌐 数据集根路径：{KAGGLE_DATA_ROOT} {'✅ 存在' if KAGGLE_DATA_ROOT.exists() else '❌ 不存在'}\")\n    \n    if not KAGGLE_DATA_ROOT.exists():\n        raise FileNotFoundError(\"Kaggle数据集根路径不存在！请检查数据集是否正确挂载\")\n\n    # 1. 定义所有必要路径\n    train_imageset_path = KAGGLE_DATA_ROOT / \"ImageSets\" / \"CLS-LOC\" / \"train_cls.txt\"\n    val_imageset_path = KAGGLE_DATA_ROOT / \"ImageSets\" / \"CLS-LOC\" / \"val.txt\"\n    \n    train_data_root = KAGGLE_DATA_ROOT / \"Data\" / \"CLS-LOC\" / \"train\"\n    val_data_root = KAGGLE_DATA_ROOT / \"Data\" / \"CLS-LOC\" / \"val\"\n    train_anno_root = KAGGLE_DATA_ROOT / \"Annotations\" / \"CLS-LOC\" / \"train\"\n    val_anno_root = KAGGLE_DATA_ROOT / \"Annotations\" / \"CLS-LOC\" / \"val\"\n\n    # 打印所有路径供调试\n    print(\"\\n📋 所有关键路径：\")\n    paths_to_check = [\n        (\"train_cls.txt\", train_imageset_path),\n        (\"val.txt\", val_imageset_path),\n        (\"train图像根路径\", train_data_root),\n        (\"val图像根路径\", val_data_root),\n        (\"train标注根路径\", train_anno_root),\n        (\"val标注根路径\", val_anno_root)\n    ]\n    for name, path in paths_to_check:\n        print(f\"  - {name}：{path} {'✅ 存在' if path.exists() else '❌ 不存在'}\")\n\n    # 2. 读取图像ID列表（修复了解析逻辑）\n    print(\"\\n\" + \"=\"*50)\n    print(\"📥 正在读取图像ID列表...\")\n    print(\"=\"*50)\n    train_image_ids = get_image_ids_from_imageset(train_imageset_path, split=\"train\")\n    val_image_ids = get_image_ids_from_imageset(val_imageset_path, split=\"val\")\n\n    # 3. 初始化Dataset（修复了标注路径拼接）\n    print(\"\\n\" + \"=\"*50)\n    print(\"📦 正在加载训练集...\")\n    print(\"=\"*50)\n    train_dataset = ILSVRC_Kaggle_Dataset(\n        image_ids_with_path=train_image_ids[:1000],  # 测试用：只取前1000个样本（可选）\n        data_root=train_data_root,\n        anno_root=train_anno_root,\n        split=\"train\",\n        transform=train_transform\n    )\n\n    print(\"\\n\" + \"=\"*50)\n    print(\"📦 正在加载验证集...\")\n    print(\"=\"*50)\n    val_dataset = ILSVRC_Kaggle_Dataset(\n        image_ids_with_path=val_image_ids[:100],  # 测试用：只取前100个样本（可选）\n        data_root=val_data_root,\n        anno_root=val_anno_root,\n        split=\"val\",\n        transform=val_transform\n    )\n\n    # 4. 初始化DataLoader\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=BATCH_SIZE,\n        shuffle=True,\n        num_workers=0 if DEVICE.type == \"cpu\" else 2,  # CPU模式下禁用多线程\n        pin_memory=True if DEVICE.type != \"cpu\" else False\n    )\n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=BATCH_SIZE,\n        shuffle=False,\n        num_workers=0 if DEVICE.type == \"cpu\" else 2,\n        pin_memory=True if DEVICE.type != \"cpu\" else False\n    )\n\n    print(f\"\\n🚀 数据加载完成：\")\n    print(f\"  - 训练集批次数量：{len(train_loader)}\")\n    print(f\"  - 验证集批次数量：{len(val_loader)}\")\n    print(f\"  - 总类别数：{train_dataset.num_classes}\")\n    \n    return train_loader, val_loader, train_dataset.num_classes\n\n\n# ====================== 3. 模型加载/训练/评估（保持不变）======================\ndef load_pretrained_model(model_name, num_classes):\n    print(f\"\\n📦 加载模型：{model_name}\")\n    if model_name == \"alexnet\":\n        try:\n            model = models.alexnet(weights=models.AlexNet_Weights.DEFAULT)\n        except:\n            model = models.alexnet(pretrained=True)\n        num_ftrs = model.classifier[6].in_features\n        model.classifier[6] = nn.Linear(num_ftrs, num_classes)\n    elif model_name == \"vgg16\":\n        try:\n            model = models.vgg16(weights=models.VGG16_Weights.DEFAULT)\n        except:\n            model = models.vgg16(pretrained=True)\n        num_ftrs = model.classifier[6].in_features\n        model.classifier[6] = nn.Linear(num_ftrs, num_classes)\n    elif model_name == \"resnet50\":\n        try:\n            model = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n        except:\n            model = models.resnet50(pretrained=True)\n        num_ftrs = model.fc.in_features\n        model.fc = nn.Linear(num_ftrs, num_classes)\n    else:\n        raise ValueError(f\"❌ 不支持的模型：{model_name}\")\n\n    for param in model.parameters():\n        param.requires_grad = False\n    if model_name in [\"alexnet\", \"vgg16\"]:\n        for param in model.classifier[-1].parameters():\n            param.requires_grad = True\n    elif model_name == \"resnet50\":\n        for param in model.fc.parameters():\n            param.requires_grad = True\n    model = model.to(DEVICE)\n    print(f\"✅ 模型加载完成（设备：{DEVICE}）\")\n    return model\n\n\ndef train_one_epoch(model, train_loader, criterion, optimizer, epoch):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS} | Training\")\n    for inputs, labels in pbar:\n        inputs, labels = inputs.to(DEVICE), labels.to(DEVICE)\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item() * inputs.size(0)\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        pbar.set_postfix({\"Loss\": f\"{running_loss/total:.4f}\", \"Top-1 Acc\": f\"{correct/total:.4f}\"})\n    return running_loss/total, correct/total\n\n\ndef evaluate_model(model, val_loader, criterion):\n    model.eval()\n    running_loss = 0.0\n    top1_correct = 0\n    top5_correct = 0\n    total = 0\n    with torch.no_grad():\n        pbar = tqdm(val_loader, desc=\"Evaluating\")\n        for inputs, labels in pbar:\n            inputs, labels = inputs.to(DEVICE), labels.to(DEVICE)\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n            running_loss += loss.item() * inputs.size(0)\n            total += labels.size(0)\n            _, top1_pred = outputs.max(1)\n            top1_correct += top1_pred.eq(labels).sum().item()\n            top5_pred = outputs.topk(5, dim=1)[1]\n            top5_correct += top5_pred.eq(labels.view(-1, 1).expand_as(top5_pred)).sum().item()\n            pbar.set_postfix({\n                \"Val Loss\": f\"{running_loss/total:.4f}\",\n                \"Top-1 Acc\": f\"{top1_correct/total:.4f}\",\n                \"Top-5 Acc\": f\"{top5_correct/total:.4f}\"\n            })\n    return running_loss/total, top1_correct/total, top5_correct/total\n\n\ndef main():\n    print(\"=\"*60)\n    print(\"🎯 ImageNet图像分类训练开始\")\n    print(\"=\"*60)\n    print(f\"📋 训练配置：\")\n    print(f\"  - 设备：{DEVICE}\")\n    print(f\"  - 批次大小：{BATCH_SIZE}\")\n    print(f\"  - 训练轮数：{EPOCHS}\")\n    print(f\"  - 学习率：{LEARNING_RATE}\")\n    print(f\"  - 模型列表：{MODEL_NAMES}\")\n    print(\"=\"*60)\n\n    # 数据预处理\n    train_transform, val_transform = get_data_transforms()\n    \n    # 加载数据（关键修复）\n    try:\n        train_loader, val_loader, num_classes = load_kaggle_ilsvrc_dataset(train_transform, val_transform)\n    except Exception as e:\n        print(f\"\\n❌ 数据加载失败：{e}\")\n        return\n\n    results = {\"model\": [], \"top1_acc\": [], \"top5_acc\": [], \"best_val_loss\": []}\n\n    for model_name in MODEL_NAMES:\n        print(f\"\\n{'='*60}\\n📌 开始训练模型：{model_name}\\n{'='*60}\")\n        model = load_pretrained_model(model_name, num_classes)\n        criterion = nn.CrossEntropyLoss().to(DEVICE)\n        optimizer = optim.SGD(\n            filter(lambda p: p.requires_grad, model.parameters()),\n            lr=LEARNING_RATE, momentum=MOMENTUM, weight_decay=WEIGHT_DECAY\n        )\n        scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)\n        best_top1_acc = 0.0\n        best_val_loss = float(\"inf\")\n\n        for epoch in range(EPOCHS):\n            train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, epoch)\n            val_loss, top1_acc, top5_acc = evaluate_model(model, val_loader, criterion)\n            scheduler.step()\n\n            if top1_acc > best_top1_acc:\n                best_top1_acc = top1_acc\n                best_val_loss = val_loss\n                torch.save(model.state_dict(), f\"{model_name}_best_weights.pth\")\n                print(f\"💾 保存最佳模型 -> Top-1 Acc: {best_top1_acc:.4f}\")\n\n            print(f\"📊 Epoch {epoch+1} 总结：\")\n            print(f\"  - 训练损失：{train_loss:.4f} | 训练准确率：{train_acc:.4f}\")\n            print(f\"  - 验证损失：{val_loss:.4f} | Top-1准确率：{top1_acc:.4f} | Top-5准确率：{top5_acc:.4f}\")\n\n        results[\"model\"].append(model_name)\n        results[\"top1_acc\"].append(best_top1_acc)\n        results[\"top5_acc\"].append(top5_acc)\n        results[\"best_val_loss\"].append(best_val_loss)\n\n    print(f\"\\n{'='*60}\\n🏆 所有模型训练完成，结果对比：\\n{'='*60}\")\n    for i in range(len(results[\"model\"])):\n        print(f\"{results['model'][i]}:\")\n        print(f\"  - Top-1准确率：{results['top1_acc'][i]:.4f}\")\n        print(f\"  - Top-5准确率：{results['top5_acc'][i]:.4f}\")\n        print(f\"  - 最佳验证损失：{results['best_val_loss'][i]:.4f}\\n\")\n\n    # 可视化\n    plt.figure(figsize=(10, 6))\n    bars = plt.bar(results[\"model\"], results[\"top1_acc\"], color=['#1f77b4', '#ff7f0e', '#2ca02c'])\n    plt.xlabel(\"经典CNN模型\", fontsize=12)\n    plt.ylabel(\"Top-1准确率\", fontsize=12)\n    plt.title(\"ImageNet图像分类 - Top-1准确率对比\", fontsize=14, fontweight='bold')\n    plt.ylim(0, 1.0)\n    for bar, acc in zip(bars, results[\"top1_acc\"]):\n        plt.text(bar.get_x() + bar.get_width()/2., bar.get_height() + 0.02,\n                 f\"{acc:.4f}\", ha='center', va='bottom', fontsize=11)\n    plt.grid(axis='y', alpha=0.3)\n    plt.savefig(\"top1_accuracy_comparison.png\", dpi=300, bbox_inches='tight')\n    plt.show()\n\n    plt.figure(figsize=(10, 6))\n    bars = plt.bar(results[\"model\"], results[\"top5_acc\"], color=['#d62728', '#9467bd', '#8c564b'])\n    plt.xlabel(\"经典CNN模型\", fontsize=12)\n    plt.ylabel(\"Top-5准确率\", fontsize=12)\n    plt.title(\"ImageNet图像分类 - Top-5准确率对比\", fontsize=14, fontweight='bold')\n    plt.ylim(0, 1.0)\n    for bar, acc in zip(bars, results[\"top5_acc\"]):\n        plt.text(bar.get_x() + bar.get_width()/2., bar.get_height() + 0.02,\n                 f\"{acc:.4f}\", ha='center', va='bottom', fontsize=11)\n    plt.grid(axis='y', alpha=0.3)\n    plt.savefig(\"top5_accuracy_comparison.png\", dpi=300, bbox_inches='tight')\n    plt.show()\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-28T07:50:32.065789Z","iopub.execute_input":"2025-11-28T07:50:32.066037Z","execution_failed":"2025-11-28T07:51:10.179Z"}},"outputs":[],"execution_count":null}]}