{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":10338,"databundleVersionId":862042,"sourceType":"competition"},{"sourceId":248096013,"sourceType":"kernelVersion"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# 快速查看可用数据集\nimport os\nprint(os.listdir(\"/kaggle/input\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T10:48:08.194934Z","iopub.execute_input":"2025-07-01T10:48:08.195729Z","iopub.status.idle":"2025-07-01T10:48:08.203153Z","shell.execute_reply.started":"2025-07-01T10:48:08.195691Z","shell.execute_reply":"2025-07-01T10:48:08.202358Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\n# 1. 定义输入数据目录（根据实际需求选择数据集）\nINPUT_DATA = \"/kaggle/input/my-pneumonia-gan\"  # 替换为你要使用的数据集\n\n# 2. 验证路径是否正确\nprint(\"\\n=== 路径验证 ===\")\nprint(f\"输入数据目录: {INPUT_DATA}\")\nprint(\"目录内容:\", os.listdir(INPUT_DATA))\n\n# 3. 验证关键文件是否存在\nessential_files = [\"final_model/model.pth\", \"train_split.csv\"]\nfor file in essential_files:\n    path = f\"{INPUT_DATA}/{file}\"\n    print(f\"{'√' if os.path.exists(path) else '×'} {path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T10:48:09.822280Z","iopub.execute_input":"2025-07-01T10:48:09.822729Z","iopub.status.idle":"2025-07-01T10:48:09.837663Z","shell.execute_reply.started":"2025-07-01T10:48:09.822709Z","shell.execute_reply":"2025-07-01T10:48:09.836924Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.nn.utils import spectral_norm\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport os\nimport numpy as np\nimport pandas as pd\n\n# ===== 1. 定义与训练时完全一致的Generator =====\nclass ResidualBlock(nn.Module):\n    def __init__(self, channels):\n        super().__init__()\n        self.conv = nn.Sequential(\n            spectral_norm(nn.Conv2d(channels, channels, 3, padding=1)),\n            nn.InstanceNorm2d(channels),\n            nn.LeakyReLU(0.2),\n            spectral_norm(nn.Conv2d(channels, channels, 3, padding=1)),\n            nn.InstanceNorm2d(channels)\n        )\n    def forward(self, x):\n        return x + self.conv(x)\n\nclass Generator(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.init_size = 128 // 4  # IMG_SIZE=128\n        self.l1 = nn.Sequential(\n            spectral_norm(nn.Linear(128, 128 * self.init_size ** 2))  # LATENT_DIM=128\n        )\n        self.conv_blocks = nn.Sequential(\n            nn.BatchNorm2d(128),\n            nn.Upsample(scale_factor=2),\n            spectral_norm(nn.Conv2d(128, 128, 3, stride=1, padding=1)),\n            nn.BatchNorm2d(128, 0.8),\n            nn.LeakyReLU(0.2, inplace=True),\n            ResidualBlock(128),\n            nn.Upsample(scale_factor=2),\n            spectral_norm(nn.Conv2d(128, 64, 3, stride=1, padding=1)),\n            nn.BatchNorm2d(64, 0.8),\n            nn.LeakyReLU(0.2, inplace=True),\n            spectral_norm(nn.Conv2d(64, 1, 3, stride=1, padding=1)),  # CHANNELS=1\n            nn.Tanh()\n        )\n    def forward(self, z):\n        out = self.l1(z)\n        out = out.view(out.shape[0], 128, self.init_size, self.init_size)\n        return self.conv_blocks(out)\n\n# ===== 2. 加载训练好的模型 =====\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ngenerator = Generator().to(device)\n\n# 加载检查点（替换为您的实际路径）\ncheckpoint_path = \"/kaggle/input/my-pneumonia-gan/final_model/model.pth\"\ncheckpoint = torch.load(checkpoint_path, map_location=device)\ngenerator.load_state_dict(checkpoint['generator'])  # 严格匹配\ngenerator.eval()\nprint(\"✅ 模型加载成功\")\n\n# ===== 3. 生成增强样本 =====\ndef generate_synthetic_samples(num_samples_normal, num_samples_pneumonia, output_dir=\"/kaggle/working/synthetic_images\"):\n    \"\"\"生成指定数量的正常和肺炎样本\"\"\"\n    os.makedirs(output_dir, exist_ok=True)\n    \n    # 生成正常样本 (label=0)\n    normal_samples = []\n    for i in range(num_samples_normal):\n        z = torch.randn(1, 128).to(device)  # LATENT_DIM=128\n        with torch.no_grad():\n            img = generator(z).cpu().squeeze().numpy()\n            img = (img * 127.5 + 127.5).astype(np.uint8)  # 转换到0-255范围\n            img_path = f\"{output_dir}/synth_normal_{i}.png\"\n            Image.fromarray(img).save(img_path)\n            normal_samples.append({\n                'patientId': f'synth_normal_{i}',\n                'dcm_path': img_path,\n                'label': 0\n            })\n        \n        if i % 100 == 0:\n            print(f\"已生成正常样本 {i+1}/{num_samples_normal}\")\n    \n    # 生成肺炎样本 (label=1)\n    pneumonia_samples = []\n    for i in range(num_samples_pneumonia):\n        z = torch.randn(1, 128).to(device)  # LATENT_DIM=128\n        with torch.no_grad():\n            img = generator(z).cpu().squeeze().numpy()\n            img = (img * 127.5 + 127.5).astype(np.uint8)  # 转换到0-255范围\n            img_path = f\"{output_dir}/synth_pneumonia_{i}.png\"\n            Image.fromarray(img).save(img_path)\n            pneumonia_samples.append({\n                'patientId': f'synth_pneumonia_{i}',\n                'dcm_path': img_path,\n                'label': 1\n            })\n        \n        if i % 100 == 0:\n            print(f\"已生成肺炎样本 {i+1}/{num_samples_pneumonia}\")\n    \n    # 合并并返回元数据\n    all_samples = normal_samples + pneumonia_samples\n    return pd.DataFrame(all_samples)\n\n# ===== 4. 智能计算需要生成的样本数量 =====\ndef calculate_samples_to_generate(original_data_path=\"/kaggle/input/my-pneumonia-gan/train_split.csv\"):\n    \"\"\"根据原始数据的类别分布，计算需要生成的正常和肺炎样本数量\"\"\"\n    try:\n        original_df = pd.read_csv(original_data_path)\n        class_distribution = original_df['label'].value_counts()\n        \n        # 0=正常，1=肺炎\n        normal_count = class_distribution.get(0, 0)\n        pneumonia_count = class_distribution.get(1, 0)\n        \n        print(f\"Original dataset distribution: Normal={normal_count}, Pneumonia={pneumonia_count}\")\n        \n        # 计算需要生成的样本数量，目标是使两类数量平衡\n        if normal_count > pneumonia_count:\n            # Need more pneumonia samples\n            return 0, normal_count - pneumonia_count\n        else:\n            # Need more normal samples\n            return pneumonia_count - normal_count, 0\n    except Exception as e:\n        print(f\"Error calculating sample counts: {e}\")\n        print(\"Using default values: generating 500 normal and 500 pneumonia samples\")\n        return 500, 500\n\n# ===== 5. 主函数 =====\nif __name__ == \"__main__\":\n    # 计算需要生成的样本数量\n    num_normal, num_pneumonia = calculate_samples_to_generate()\n    \n    # 如果不需要生成任何样本，则使用默认值\n    if num_normal == 0 and num_pneumonia == 0:\n        print(\"数据已平衡，无需生成额外样本\")\n        num_normal, num_pneumonia = 500, 500\n    \n    print(f\"计划生成: 正常={num_normal}, 肺炎={num_pneumonia}\")\n    \n    # 生成样本并获取元数据\n    synthetic_metadata = generate_synthetic_samples(num_normal, num_pneumonia)\n    \n    # 保存元数据\n    metadata_path = \"/kaggle/working/synthetic_samples_metadata.csv\"\n    synthetic_metadata.to_csv(metadata_path, index=False)\n    print(f\"✅ 合成样本元数据已保存到 {metadata_path}\")\n    \n    # ===== 6. 验证生成质量 =====\n    print(\"\\n=== 生成样本示例 ===\")\n    sample_files = [f for f in os.listdir(\"/kaggle/working/synthetic_images\") if f.endswith('.png')][:4]\n    fig, axes = plt.subplots(1, 4, figsize=(15, 4))\n    for i, file in enumerate(sample_files):\n        img = Image.open(f\"/kaggle/working/synthetic_images/{file}\")\n        axes[i].imshow(img, cmap='gray')\n        axes[i].set_title('Normal' if 'normal' in file else 'Pneumonia')\n        axes[i].axis('off')\n    plt.tight_layout()\n    plt.savefig('/kaggle/working/synthetic_samples_example.jpg')\n    plt.show()\n    \n    # ===== 7. 分析生成样本的类别分布 =====\n    plt.figure(figsize=(8, 5))\n    class_counts = synthetic_metadata['label'].value_counts()\n    # 确保标签映射正确（0=Normal，1=Pneumonia）\n    class_labels = {0: \"Normal\", 1: \"Pneumonia\"}\n    class_counts.index = [class_labels[idx] for idx in class_counts.index]  # 重命名索引\n\n    class_counts.plot(kind='bar', color=['skyblue', 'salmon'])\n    plt.title('Synthetic Sample Class Distribution')  # 英文标题\n    plt.xticks(rotation=0)  # 刻度旋转（可选，保持水平）\n    plt.xlabel('Class')      # 新增X轴标签\n    plt.ylabel('Number of Samples')  # 英文Y轴标签\n    plt.savefig('/kaggle/working/synthetic_class_distribution.jpg', bbox_inches='tight')  # 优化保存\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T11:34:58.870039Z","iopub.execute_input":"2025-07-01T11:34:58.870291Z","iopub.status.idle":"2025-07-01T11:35:02.219503Z","shell.execute_reply.started":"2025-07-01T11:34:58.870273Z","shell.execute_reply":"2025-07-01T11:35:02.218829Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===== 1. 导入所有依赖 =====\nimport os\nimport random\nimport numpy as np\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport pandas as pd\n\n# ===== 2. 验证生成图像质量 =====\n# 创建输出目录（如果不存在）\nos.makedirs(\"/kaggle/working/synthetic_images\", exist_ok=True)\n\n# 随机检查10张生成图像\nsample_files = random.sample(os.listdir(\"/kaggle/working/synthetic_images\"), min(10, len(os.listdir(\"/kaggle/working/synthetic_images\"))))\n\nplt.figure(figsize=(15, 5))\nfor i, file in enumerate(sample_files):\n    img = Image.open(f\"/kaggle/working/synthetic_images/{file}\")\n    plt.subplot(2, 5, i+1)\n    plt.imshow(img, cmap='gray')\n    plt.axis('off')\nplt.tight_layout()\nplt.show()\n# ===== 3. 构建增强数据集（修正版） =====\ntry:\n    # 加载原始训练集\n    train_df = pd.read_csv(\"/kaggle/input/my-pneumonia-gan/train_split.csv\")\n    \n    # 加载合成样本的元数据（包含正确标签）\n    synth_metadata_path = \"/kaggle/working/synthetic_samples_metadata.csv\"\n    if not os.path.exists(synth_metadata_path):\n        raise FileNotFoundError(f\"合成样本元数据文件不存在: {synth_metadata_path}\")\n    \n    synth_df = pd.read_csv(synth_metadata_path)\n    \n    # 合并数据集（保持元数据中的正确标签）\n    augmented_df = pd.concat([train_df, synth_df], ignore_index=True)\n    augmented_df.to_csv(\"/kaggle/working/augmented_train.csv\", index=False)\n    \n    print(f\"增强后数据集：原始 {len(train_df)} + 合成 {len(synth_df)} = {len(augmented_df)} 张\")\n    print(\"前5条记录：\\n\", augmented_df.head())\n    \n    # 验证类别分布\n    class_distribution = augmented_df['label'].value_counts()\n    print(\"\\n增强后类别分布:\")\n    print(f\"正常样本 (label=0): {class_distribution.get(0, 0)}\")\n    print(f\"肺炎样本 (label=1): {class_distribution.get(1, 0)}\")\n    \nexcept Exception as e:\n    print(f\"发生错误：{str(e)}\")\n    print(\"请检查：\")\n    print(\"1. 原始CSV文件路径是否正确\")\n    print(\"2. 合成样本元数据是否已生成\")\n    print(\"3. 元数据文件路径是否正确\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T11:35:24.728377Z","iopub.execute_input":"2025-07-01T11:35:24.728678Z","iopub.status.idle":"2025-07-01T11:35:25.292404Z","shell.execute_reply.started":"2025-07-01T11:35:24.728657Z","shell.execute_reply":"2025-07-01T11:35:25.291603Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport pandas as pd\n\n# 1. 加载增强数据集\naugmented_df = pd.read_csv(\"/kaggle/working/augmented_train.csv\")\n\n# 2. 计算类别分布（这一步必须在使用class_dist之前！）\nclass_dist = augmented_df['label'].value_counts()\n\n# 3. 显式映射标签含义（避免硬编码）\nlabel_mapping = {0: 'Normal', 1: 'Pneumonia'}\nclass_dist.index = [label_mapping[idx] for idx in class_dist.index]  # 现在class_dist已定义\n\n# 4. 可视化（英文标题）\nplt.figure(figsize=(8, 4))\nplt.pie(\n    class_dist, \n    labels=class_dist.index,  # 使用映射后的标签\n    autopct='%1.1f%%', \n    colors=['lightgreen', 'lightcoral'],\n    textprops={'fontsize': 12}\n)  \nplt.title(\"Class Distribution in Augmented Dataset\", fontsize=14, pad=20)\nplt.savefig('/kaggle/working/class_balance.jpg', dpi=300, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T11:35:40.228027Z","iopub.execute_input":"2025-07-01T11:35:40.228643Z","iopub.status.idle":"2025-07-01T11:35:40.403837Z","shell.execute_reply.started":"2025-07-01T11:35:40.228615Z","shell.execute_reply":"2025-07-01T11:35:40.403065Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 检查当前增强数据集的真实/合成样本比例及类别分布\nreal_samples = augmented_df[~augmented_df['dcm_path'].str.contains('synthetic')]\nsynth_samples = augmented_df[augmented_df['dcm_path'].str.contains('synthetic')]\n\n# 复用标签映射（与可视化一致：0=Normal，1=Pneumonia）\nlabel_mapping = {0: 'Normal', 1: 'Pneumonia'}\n\n# ========== 真实样本统计 ==========\nreal_class_dist = real_samples['label'].value_counts().rename(index=label_mapping)\nreal_pneumonia_ratio = real_samples['label'].mean()  # label=1的比例\n\nprint(f\"真实样本: {len(real_samples)}\")\nfor cls in label_mapping.values():\n    count = real_class_dist.get(cls, 0)\n    print(f\"  - {cls}: {count}\")\nprint(f\"  - 肺炎占比: {real_pneumonia_ratio:.1%}\")\n\n# ========== 合成样本统计 ==========\nsynth_class_dist = synth_samples['label'].value_counts().rename(index=label_mapping)\n\nprint(f\"\\n合成样本: {len(synth_samples)}\")\nfor cls in label_mapping.values():\n    count = synth_class_dist.get(cls, 0)\n    print(f\"  - {cls}: {count}\")\n\n# ========== 整体肺炎占比 ==========\noverall_pneumonia_ratio = augmented_df['label'].mean()\nprint(f\"\\n肺炎总占比: {overall_pneumonia_ratio:.1%}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T11:35:49.925394Z","iopub.execute_input":"2025-07-01T11:35:49.925685Z","iopub.status.idle":"2025-07-01T11:35:49.952235Z","shell.execute_reply.started":"2025-07-01T11:35:49.925666Z","shell.execute_reply":"2025-07-01T11:35:49.951484Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models\nfrom torchvision.models import ResNet18_Weights\nfrom PIL import Image\nimport pydicom  # 添加DICOM处理库\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import accuracy_score, roc_auc_score, confusion_matrix, classification_report\n\n# 设置随机种子以确保结果可复现\ntorch.manual_seed(42)\nnp.random.seed(42)\n\n\n# ====================== 1. 数据集类 ======================\nclass PneumoniaDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = row['dcm_path']\n        \n        # 检查文件是否存在\n        if not os.path.exists(img_path):\n            raise FileNotFoundError(f\"错误：在{img_path}处未找到图像\")\n        \n        # 处理DICOM文件\n        if img_path.lower().endswith('.dcm'):\n            try:\n                dicom = pydicom.dcmread(img_path)\n                pixel_array = dicom.pixel_array\n                \n                # 将像素数组转换为PIL图像\n                if pixel_array.dtype != np.uint8:\n                    # 转换为8位图像（如果需要）\n                    pixel_array = self._convert_to_uint8(pixel_array)\n                    \n                image = Image.fromarray(pixel_array)\n                image = image.convert('RGB')  # 确保为RGB格式\n            except Exception as e:\n                raise RuntimeError(f\"无法处理DICOM文件 {img_path}: {str(e)}\")\n        else:\n            # 处理普通图像（如PNG/JPG）\n            image = Image.open(img_path).convert('RGB')\n\n        label = torch.tensor(row['label'], dtype=torch.long)\n\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label\n    \n    def _convert_to_uint8(self, array):\n        \"\"\"将不同位深度的像素数组转换为uint8格式\"\"\"\n        if np.issubdtype(array.dtype, np.floating):\n            # 浮点型数据（如[-1,1]或[0,1]）\n            array = (array * 255).astype(np.uint8)\n        else:\n            # 整型数据（如16位）\n            array_min, array_max = array.min(), array.max()\n            if array_min == array_max:\n                # 避免除以零\n                array = np.zeros_like(array, dtype=np.uint8)\n            else:\n                array = ((array - array_min) / (array_max - array_min) * 255).astype(np.uint8)\n        return array\n\n\n# ====================== 2. 数据准备 ======================\ndef prepare_data(data_df, batch_size=32, is_augmented=False):\n    transforms_list = [\n        transforms.Resize((224, 224)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n    ]\n    # 仅对增强模型的训练集添加额外增强\n    if is_augmented:\n        transforms_list.insert(1, transforms.RandomHorizontalFlip(p=0.5))\n        transforms_list.insert(2, transforms.RandomRotation(15))\n        transforms_list.insert(3, transforms.ColorJitter(brightness=0.2, contrast=0.2))\n    \n    transform = transforms.Compose(transforms_list)\n    dataset = PneumoniaDataset(data_df, transform=transform)\n    dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True, num_workers=2)\n    return dataloader\n\n# ====================== 3. 验证数据加载器 ======================\ndef validate_data_loader(data_loader, num_samples=5, device=\"cpu\"):\n    print(f\"\\n{'='*50}\")\n    print(f\"验证数据加载器（加载{num_samples}个批次）...\")\n    for i, (images, labels) in enumerate(data_loader):\n        if i >= num_samples:\n            break\n\n        images, labels = images.to(device), labels.to(device)\n\n        print(f\"批次 {i+1}/{num_samples}：\")\n        print(f\"  图像形状: {images.shape}\")\n        print(f\"  标签分布: {np.bincount(labels.cpu().numpy())}\")\n\n        plt.figure(figsize=(4, 4))\n        img = images[0].permute(1, 2, 0).cpu().numpy()\n        img = (img * np.array([0.229, 0.224, 0.225])) + np.array([0.485, 0.456, 0.406])\n        img = np.clip(img, 0, 1)\n        plt.imshow(img)\n        plt.title(f\"label: {labels[0].item()}\")\n        plt.axis('off')\n        plt.tight_layout()\n        plt.show()\n\n    print(f\"数据加载器验证完成{'='*50}\\n\")\n\n\n# ====================== 4. CNN模型 ======================\nclass PneumoniaCNN(nn.Module):\n    def __init__(self, num_classes=2):\n        super(PneumoniaCNN, self).__init__()\n        self.model = models.resnet18(weights=ResNet18_Weights.IMAGENET1K_V1)\n        self.model.fc = nn.Linear(self.model.fc.in_features, num_classes)\n\n    def forward(self, x):\n        return self.model(x)\n\n\n# ====================== 5. 训练函数 ======================\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\ndef train_model(model, train_loader, val_loader, criterion, optimizer, device, epochs=10):\n    best_val_auc = 0.0  # 改用AUC作为保存依据\n    history = {'train_loss': [], 'train_acc': [], 'val_loss': [], 'val_acc': [], 'val_auc': []}\n    \n    # 学习率调度：余弦退火\n    scheduler = CosineAnnealingLR(optimizer, T_max=epochs, eta_min=1e-5)\n    \n    for epoch in range(epochs):\n        model.train()\n        train_loss = 0.0\n        train_correct = 0\n        train_total = 0\n\n        for inputs, labels in train_loader:\n            inputs, labels = inputs.to(device), labels.to(device)\n\n            optimizer.zero_grad()\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n\n            train_loss += loss.item()\n            _, predicted = outputs.max(1)\n            train_total += labels.size(0)\n            train_correct += predicted.eq(labels).sum().item()\n\n        train_acc = 100. * train_correct / train_total\n        history['train_loss'].append(train_loss / len(train_loader))\n        history['train_acc'].append(train_acc)\n\n        model.eval()\n        val_loss = 0.0\n        val_correct = 0\n        val_total = 0\n        all_labels = []\n        all_probs = []\n\n        with torch.no_grad():\n            for inputs, labels in val_loader:\n                inputs, labels = inputs.to(device), labels.to(device)\n\n                outputs = model(inputs)\n                loss = criterion(outputs, labels)\n\n                val_loss += loss.item()\n                _, predicted = outputs.max(1)\n                val_total += labels.size(0)\n                val_correct += predicted.eq(labels).sum().item()\n\n                probs = torch.softmax(outputs, dim=1)[:, 1].cpu().numpy()\n                all_labels.extend(labels.cpu().numpy())\n                all_probs.extend(probs)\n\n        val_acc = 100. * val_correct / val_total\n        val_auc = roc_auc_score(all_labels, all_probs)\n        history['val_loss'].append(val_loss / len(val_loader))\n        history['val_acc'].append(val_acc)\n        history['val_auc'].append(val_auc)\n\n        print(f'轮次 {epoch+1}/{epochs}')\n        print(f'训练损失: {history[\"train_loss\"][-1]:.4f} | 训练准确率: {train_acc:.2f}%')\n        print(f'验证损失: {history[\"val_loss\"][-1]:.4f} | 验证准确率: {val_acc:.2f}% | 验证AUC: {val_auc:.4f}')\n\n        # 保存最佳模型（基于AUC）\n        if val_auc > best_val_auc:\n            best_val_auc = val_auc\n            torch.save(model.state_dict(), '/kaggle/working/best_model_base.pth')\n            print(f'保存最佳模型，AUC: {best_val_auc:.4f}')\n\n        # 更新学习率\n        scheduler.step()\n\n    return history, best_val_auc\n\n# ====================== 6. 评估函数 ======================\ndef evaluate_model(model, test_loader, device):\n    model.eval()\n    all_labels = []\n    all_preds = []\n    all_probs = []\n\n    with torch.no_grad():\n        for inputs, labels in test_loader:\n            inputs, labels = inputs.to(device), labels.to(device)\n\n            outputs = model(inputs)\n            probs = torch.softmax(outputs, dim=1)[:, 1]\n            _, predicted = outputs.max(1)\n\n            all_probs.extend(probs.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n            all_preds.extend(predicted.cpu().numpy())\n\n    acc = accuracy_score(all_labels, all_preds)\n    auc = roc_auc_score(all_labels, all_probs)\n    cm = confusion_matrix(all_labels, all_preds)\n    report = classification_report(all_labels, all_preds, target_names=['正常', '肺炎'])\n\n    print(f'测试准确率: {acc:.4f}')\n    print(f'测试AUC: {auc:.4f}')\n    print(f'混淆矩阵:\\n{cm}')\n    print(f'分类报告:\\n{report}')\n\n    return acc, auc, cm, report\n\n\n# ====================== 7. 主函数（基础模型） ======================\ndef main():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"使用设备: {device}\")\n\n    train_df = pd.read_csv(\"/kaggle/input/my-pneumonia-gan/train_split.csv\")\n\n    train_df, val_df = train_test_split(train_df, test_size=0.2, random_state=42, stratify=train_df['label'])\n    print(f\"原始训练集大小: {len(train_df)} | 验证集大小: {len(val_df)}\")\n\n    train_loader = prepare_data(train_df)\n    val_loader = prepare_data(val_df)\n\n    validate_data_loader(train_loader, num_samples=3, device=device)\n\n    model = PneumoniaCNN().to(device)\n\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.Adam(model.parameters(), lr=0.001)\n\n    print(\"开始训练基础模型...\")\n    history, best_acc = train_model(model, train_loader, val_loader, criterion, optimizer, device, epochs=10)\n\n    plt.figure(figsize=(12, 4))\n    plt.subplot(1, 2, 1)\n    plt.plot(history['train_loss'], label='Train Loss')\n    plt.plot(history['val_loss'], label='Val Loss')\n    plt.legend()\n    plt.title('Loss History')\n\n    plt.subplot(1, 2, 2)\n    plt.plot(history['train_acc'], label='Train Acc')\n    plt.plot(history['val_acc'], label='Val Acc')\n    plt.legend()\n    plt.title('Accuracy History')\n    plt.savefig('/kaggle/working/base_model_training_history.jpg')\n    plt.show()\n\n    test_df = pd.read_csv(\"/kaggle/input/my-pneumonia-gan/val_split.csv\")\n    test_loader = prepare_data(test_df)\n\n    model.load_state_dict(torch.load('/kaggle/working/best_model_base.pth'))\n    print(\"\\n评估基础模型性能:\")\n    acc, auc, cm, report = evaluate_model(model, test_loader, device)\n\n    with open('/kaggle/working/base_model_results.txt', 'w') as f:\n        f.write(f'基础模型测试准确率: {acc:.4f}\\n')\n        f.write(f'基础模型测试AUC: {auc:.4f}\\n')\n        f.write(f'混淆矩阵:\\n{cm}\\n')\n        f.write(f'分类报告:\\n{report}')\n\n    return acc, history\n\n\n# ====================== 8. 使用增强数据训练（可选） ======================\ndef train_with_augmented_data(base_accuracy, base_history):\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"使用设备: {device}\")\n\n    # 检查是否存在增强数据\n    augmented_csv = \"/kaggle/working/augmented_train.csv\"\n    if not os.path.exists(augmented_csv):\n        print(f\"\\n警告: {augmented_csv} 未找到，跳过增强数据训练。\")\n        print(\"请先运行「构建增强数据集」代码，或检查文件路径。\")\n        return None, None\n\n    augmented_df = pd.read_csv(augmented_csv)\n\n    train_df, val_df = train_test_split(augmented_df, test_size=0.2, random_state=42, stratify=augmented_df['label'])\n    print(f\"增强训练集大小: {len(train_df)} | 验证集大小: {len(val_df)}\")\n\n    train_loader = prepare_data(train_df, is_augmented=True)  # 增强训练集\n    val_loader = prepare_data(val_df)  # 验证集不变\n\n    validate_data_loader(train_loader, num_samples=3, device=device)\n\n    model = PneumoniaCNN().to(device)\n\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.Adam(model.parameters(), lr=0.001, weight_decay=1e-4)  # L2正则化\n    print(\"开始训练增强模型...\")\n    \n    # 保存增强模型到不同路径\n    history, best_acc = train_model(model, train_loader, val_loader, criterion, optimizer, device, epochs=10)\n    torch.save(model.state_dict(), '/kaggle/working/best_model_augmented.pth')  # 保存增强模型\n\n    plt.figure(figsize=(12, 4))\n    plt.subplot(1, 2, 1)\n    plt.plot(history['train_loss'], label='Train Loss')\n    plt.plot(history['val_loss'], label='Val Loss')\n    plt.legend()\n    plt.title('Loss History')\n\n    plt.subplot(1, 2, 2)\n    plt.plot(history['train_acc'], label='Train Acc')\n    plt.plot(history['val_acc'], label='Val Acc')\n    plt.legend()\n    plt.title('Accuracy History')\n    plt.savefig('/kaggle/working/augmented_model_training_history.jpg')\n    plt.show()\n\n    test_df = pd.read_csv(\"/kaggle/input/my-pneumonia-gan/val_split.csv\")\n    test_loader = prepare_data(test_df)\n\n    # 加载增强模型\n    model.load_state_dict(torch.load('/kaggle/working/best_model_augmented.pth'))\n    print(\"\\n评估增强模型性能:\")\n    acc, auc, cm, report = evaluate_model(model, test_loader, device)\n\n    with open('/kaggle/working/augmented_model_results.txt', 'w') as f:\n        f.write(f'增强模型测试准确率: {acc:.4f}\\n')\n        f.write(f'增强模型测试AUC: {auc:.4f}\\n')\n        f.write(f'混淆矩阵:\\n{cm}\\n')\n        f.write(f'分类报告:\\n{report}')\n\n    improvement = (acc - base_accuracy) / base_accuracy * 100\n    print(f\"\\n模型准确率提升: {improvement:.2f}%\")\n\n    plt.figure(figsize=(10, 6))\n    plt.bar(['original', 'enhance'], [base_accuracy, acc], color=['lightblue', 'lightgreen'])\n    plt.ylim(0.8, 1.0)\n    plt.title('Model Performance Comparison', fontsize=14)\n    plt.ylabel('Accuracy', fontsize=12)\n    plt.grid(axis='y', linestyle='--', alpha=0.7)\n\n    for i, v in enumerate([base_accuracy, acc]):\n        plt.text(i, v + 0.002, f'{v:.4f}', ha='center', fontsize=12)\n\n    plt.savefig('/kaggle/working/model_comparison.jpg')\n    plt.show()\n\n    return acc, history\n\n\n# ====================== 9. 分析结果 ======================\ndef analyze_results(base_accuracy, base_history, augmented_accuracy, augmented_history):\n    if augmented_accuracy is None:\n        print(\"未训练增强模型，跳过性能比较。\")\n        return\n    \n    improvement = (augmented_accuracy - base_accuracy) / base_accuracy * 100\n    result_summary = f\"\"\"\n    ================= 模型性能比较总结 =================\n    基础模型准确率: {base_accuracy:.4f}\n    增强模型准确率: {augmented_accuracy:.4f}\n    准确率提升: {improvement:.2f}%\n    \n    是否达到5%的提升目标? {'✅ 是' if improvement >= 5 else '❌ 否'}\n    =================================================\n    \"\"\"\n\n    print(result_summary)\n\n    with open('/kaggle/working/results_summary.txt', 'w') as f:\n        f.write(result_summary)\n\n    plt.figure(figsize=(14, 5))\n\n    plt.subplot(1, 2, 1)\n    plt.plot(base_history['val_acc'], label='Original Val Acc')\n    plt.plot(augmented_history['val_acc'], label='Augmented Val Acc')\n    plt.xlabel('Epoch')\n    plt.ylabel('Accuracy')\n    plt.title('Validation Accuracy Comparison')\n    plt.legend()\n    plt.grid(True)\n\n    plt.subplot(1, 2, 2)\n    plt.plot(base_history['val_loss'], label='Original Val Loss')\n    plt.plot(augmented_history['val_loss'], label='Augmented Val Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.title('Validation Loss Comparison')\n    plt.legend()\n    plt.grid(True)\n\n    plt.tight_layout()\n    plt.savefig('/kaggle/working/combined_training_history.jpg')\n    plt.show()\n\n    return result_summary\n\n\n# ====================== 主执行 ======================\nif __name__ == \"__main__\":\n    try:\n        # 确保pydicom库已安装\n        import pydicom\n    except ImportError:\n        print(\"错误：缺少pydicom库。请安装：\")\n        print(\"!pip install pydicom\")\n        exit()\n        \n    try:\n        base_accuracy, base_history = main()\n        \n        # 若基础模型训练成功，尝试增强训练\n        if base_accuracy is not None:\n            augmented_accuracy, augmented_history = train_with_augmented_data(base_accuracy, base_history)\n            analyze_results(base_accuracy, base_history, augmented_accuracy, augmented_history)\n                \n    except Exception as e:\n        print(f\"\\n错误: {str(e)}\")\n        print(\"建议检查：\")\n        print(\"1. 数据集文件是否存在（train_split.csv、val_split.csv、augmented_train.csv）\")\n        print(\"2. 图像路径是否正确（dcm_path字段是否有效）\")\n        print(\"3. 环境依赖是否完整（PyTorch、Torchvision、pydicom等）\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T11:37:52.345018Z","iopub.execute_input":"2025-07-01T11:37:52.345355Z","iopub.status.idle":"2025-07-01T12:20:43.893327Z","shell.execute_reply.started":"2025-07-01T11:37:52.345327Z","shell.execute_reply":"2025-07-01T12:20:43.892246Z"}},"outputs":[],"execution_count":null}]}