{"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"}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":" #安装必要库（Kaggle已预装大部分，只需额外安装pydicom）\n!pip install pydicom --quiet\n\n# === 导入所有需要的库 ===\nimport os\nimport pandas as pd\nimport numpy as np\nimport pydicom\nfrom PIL import Image\nimport cv2\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm  # 进度条工具\n\n# === 设置Kaggle专用路径 ===\n# 输入路径（数据集位置）\nINPUT_DIR = \"/kaggle/input/rsna-pneumonia-detection-challenge\"\n# 输出路径（处理后的文件保存位置）\nOUTPUT_DIR = \"/kaggle/working\"\n\n# === 1. 提取分类标签 ===\nprint(\"步骤1/5: 处理标签...\")\ndf_class = pd.read_csv(os.path.join(INPUT_DIR, \"stage_2_detailed_class_info.csv\"))\n\n# 筛选有效类别\nvalid_classes = ['Normal', 'Lung Opacity']\ndf_valid = df_class[df_class['class'].isin(valid_classes)].copy()\n\n# 创建标签映射\ndf_valid['label'] = df_valid['class'].map({'Normal': 0, 'Lung Opacity': 1})\n\n# 添加DICOM文件路径\ndf_valid['dcm_path'] = df_valid['patientId'].apply(\n    lambda x: os.path.join(INPUT_DIR, \"stage_2_train_images\", f\"{x}.dcm\")\n)\n\n# 检查文件是否存在\ndf_valid['exists'] = df_valid['dcm_path'].apply(os.path.exists)\nprint(f\"有效图像数量: {df_valid['exists'].sum()}/{len(df_valid)}\")\n\n# 保存预处理标签\ndf_valid[['patientId', 'dcm_path', 'label']].to_csv(\n    os.path.join(OUTPUT_DIR, 'preprocessed_labels.csv'), index=False\n)\n\n# === 2. DICOM转图像 + 预处理 ===\nprint(\"\\n步骤2/5: 转换DICOM文件...\")\nos.makedirs(os.path.join(OUTPUT_DIR, \"processed_images\"), exist_ok=True)\n\ndef preprocess_dicom(dcm_path, output_size=256):\n    \"\"\"处理DICOM文件的函数\"\"\"\n    dcm = pydicom.dcmread(dcm_path)\n    img = dcm.pixel_array.astype(float)\n    \n    # 归一化到0-255\n    img = (img - np.min(img)) / (np.max(img) - np.min(img)) * 255\n    img = img.astype(np.uint8)\n    \n    # 转换为3通道\n    if len(img.shape) == 2:\n        img = np.stack([img]*3, axis=-1)\n    \n    # 调整尺寸\n    h, w = img.shape[:2]\n    scale = output_size / max(h, w)\n    new_h, new_w = int(h*scale), int(w*scale)\n    img_resized = cv2.resize(img, (new_w, new_h))\n    \n    # 边缘填充\n    pad_h = (output_size - new_h) // 2\n    pad_w = (output_size - new_w) // 2\n    img_padded = cv2.copyMakeBorder(\n        img_resized, \n        pad_h, output_size - new_h - pad_h, \n        pad_w, output_size - new_w - pad_w, \n        cv2.BORDER_CONSTANT, \n        value=0\n    )\n    return img_padded\n\n# 批量处理所有图像（使用进度条）\nfor idx, row in tqdm(df_valid.iterrows(), total=len(df_valid)):\n    if not row['exists']: \n        continue\n    \n    output_path = os.path.join(OUTPUT_DIR, \"processed_images\", f\"{row['patientId']}.png\")\n    # 如果文件已处理则跳过\n    if not os.path.exists(output_path):\n        try:\n            img = preprocess_dicom(row['dcm_path'])\n            Image.fromarray(img).save(output_path)\n        except Exception as e:\n            print(f\"处理 {row['patientId']} 失败: {str(e)}\")\n\n# === 3. 数据分析 ===\nprint(\"\\n步骤3/5: 数据分析...\")\n# 类别分布\nplt.figure(figsize=(10,5))\ndf_valid['label'].value_counts().plot.pie(\n    autopct='%1.1f%%', \n    labels = ['Normal', 'Pneumonia'],\n    colors=['lightgreen', 'lightcoral']\n)\nplt.title(\"Class Distribution\")\nplt.savefig(os.path.join(OUTPUT_DIR, 'class_distribution.jpg'))\nplt.show()\n\n# === 4. 划分数据集 ===\nprint(\"\\n步骤4/5: 划分数据集...\")\ntrain_df, test_df = train_test_split(\n    df_valid, \n    test_size=0.2, \n    stratify=df_valid['label'],\n    random_state=42\n)\nval_df, test_df = train_test_split(\n    test_df, \n    test_size=0.5, \n    stratify=test_df['label'],\n    random_state=42\n)\n\nprint(f\"训练集: {len(train_df)} 张\")\nprint(f\"验证集: {len(val_df)} 张\")\nprint(f\"测试集: {len(test_df)} 张\")\n\n# 保存划分结果\ntrain_df.to_csv(os.path.join(OUTPUT_DIR, 'train_split.csv'), index=False)\nval_df.to_csv(os.path.join(OUTPUT_DIR, 'val_split.csv'), index=False)\ntest_df.to_csv(os.path.join(OUTPUT_DIR, 'test_split.csv'), index=False)\n\n# === 5. 可视化样本 ===\nprint(\"\\n步骤5/5: 生成样本图像...\")\ndef plot_samples(df, title):\n    fig, axes = plt.subplots(1, 4, figsize=(20, 5))\n    for i in range(4):\n        sample = df.iloc[i]\n        img_path = os.path.join(OUTPUT_DIR, \"processed_images\", f\"{sample['patientId']}.png\")\n        img = Image.open(img_path)\n        axes[i].imshow(img, cmap='gray')\n        axes[i].set_title(f\"{sample['class']} (ID: {sample['patientId']})\")\n        axes[i].axis('off')\n    plt.suptitle(title, fontsize=16)\n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, f\"{title}.jpg\"))\n\nplot_samples(train_df[train_df['label']==0], \"Example of Normal Sample\")\nplot_samples(train_df[train_df['label']==1], \"Example of pneumonia sample\")\n\nprint(\"\\n=== 预处理完成！ ===\")\nprint(f\"处理后的图像保存在: {os.path.join(OUTPUT_DIR, 'processed_images')}\")\nprint(f\"标签文件保存在: {OUTPUT_DIR}\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-06-30T07:03:42.432982Z","iopub.execute_input":"2025-06-30T07:03:42.433665Z","iopub.status.idle":"2025-06-30T07:14:50.744942Z","shell.execute_reply.started":"2025-06-30T07:03:42.433639Z","shell.execute_reply":"2025-06-30T07:14:50.744155Z"}},"outputs":[],"execution_count":null},{"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 transforms, utils\nfrom torch.nn.utils import spectral_norm\nimport os\nfrom PIL import Image\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\n\n# ========== 超参数配置 ==========\nEPOCHS = 50\nBATCH_SIZE = 32\nLR = 1e-4\nLATENT_DIM = 128\nN_CRITIC = 3\nLAMBDA_GP = 100\nIMG_SIZE = 128\nCHANNELS = 1\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# ========== 模型架构 ==========\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    \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 = IMG_SIZE // 4\n        self.l1 = nn.Sequential(\n            spectral_norm(nn.Linear(LATENT_DIM, 128 * self.init_size ** 2))\n        )\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, CHANNELS, 3, stride=1, padding=1)),\n            nn.Tanh()\n        )\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        img = self.conv_blocks(out)\n        return img\n\nclass Discriminator(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.model = nn.Sequential(\n            spectral_norm(nn.Conv2d(CHANNELS, 64, 4, stride=2, padding=1)),\n            nn.LeakyReLU(0.2, inplace=True),\n            spectral_norm(nn.Conv2d(64, 128, 4, stride=2, padding=1)),\n            nn.InstanceNorm2d(128),\n            nn.LeakyReLU(0.2, inplace=True),\n            spectral_norm(nn.Conv2d(128, 256, 4, stride=2, padding=1)),\n            nn.InstanceNorm2d(256),\n            nn.LeakyReLU(0.2, inplace=True),\n            spectral_norm(nn.Conv2d(256, 512, 4, stride=2, padding=1)),\n            nn.InstanceNorm2d(512),\n            nn.LeakyReLU(0.2, inplace=True),\n            spectral_norm(nn.Conv2d(512, 1, 4, stride=1, padding=0))\n        )\n\n    def forward(self, img):\n        return self.model(img)\n\n# ========== 关键函数 ==========\ndef compute_gradient_penalty(D, real_samples, fake_samples, lambda_gp):\n    alpha = torch.rand((real_samples.size(0), 1, 1, 1), device=DEVICE)\n    interpolates = (alpha * real_samples + ((1 - alpha) * fake_samples)).requires_grad_(True)\n    d_interpolates = D(interpolates)\n    \n    gradients = torch.autograd.grad(\n        outputs=d_interpolates,\n        inputs=interpolates,\n        grad_outputs=torch.ones_like(d_interpolates),\n        create_graph=True,\n        retain_graph=True,\n        only_inputs=True\n    )[0]\n    \n    gradients = gradients.view(gradients.size(0), -1)\n    gradient_penalty = ((gradients.norm(2, dim=1) - 1) ** 2).mean() * lambda_gp\n    return gradient_penalty\n\ndef save_image(tensor, filename, nrow=8, padding=2):\n    grid = utils.make_grid(tensor, nrow=nrow, padding=padding, normalize=True)\n    ndarr = grid.mul(255).add_(0.5).clamp_(0, 255).permute(1, 2, 0).to('cpu', torch.uint8).numpy()\n    im = Image.fromarray(ndarr)\n    im.save(filename)\n\n# ========== 数据准备 ==========\ndef prepare_data():\n    data_dir = \"/kaggle/working/processed_images\"\n    os.makedirs(data_dir, exist_ok=True)\n    os.makedirs(\"/kaggle/working/samples\", exist_ok=True)\n    os.makedirs(\"/kaggle/working/checkpoints\", exist_ok=True)\n    \n    # 示例：从Kaggle数据集解压（根据实际情况修改）\n    if len(os.listdir(data_dir)) == 0:\n        print(\"数据目录为空，请确保已添加数据\")\n        print(\"示例解压命令（取消注释使用）：\")\n        print(\"# !unzip /kaggle/input/your-dataset.zip -d /kaggle/working/processed_images\")\n        return False\n    return True\n\n# ========== 训练流程 ==========\ndef train():\n    if not prepare_data():\n        return\n\n    transform = transforms.Compose([\n        transforms.Resize(IMG_SIZE),\n        transforms.ToTensor(),\n        transforms.Normalize([0.5], [0.5])\n    ])\n    \n    class XRayDataset(torch.utils.data.Dataset):\n        def __init__(self, img_dir, transform=None):\n            self.img_dir = img_dir\n            self.transform = transform\n            self.img_list = [f for f in os.listdir(img_dir) \n                           if f.lower().endswith(('.png', '.jpg', '.jpeg', '.tif'))]\n            if not self.img_list:\n                raise ValueError(f\"目录 {img_dir} 中未找到有效图像\")\n            \n        def __len__(self):\n            return len(self.img_list)\n            \n        def __getitem__(self, idx):\n            img_path = os.path.join(self.img_dir, self.img_list[idx])\n            try:\n                with Image.open(img_path) as img:\n                    image = img.convert('L')\n                    if self.transform:\n                        image = self.transform(image)\n                    return image\n            except Exception as e:\n                print(f\"加载 {img_path} 失败: {str(e)}\")\n                return torch.zeros(1, IMG_SIZE, IMG_SIZE)\n    \n    try:\n        dataset = XRayDataset(\"/kaggle/working/processed_images\", transform)\n        print(f\"成功加载 {len(dataset)} 张图像\")\n    except Exception as e:\n        print(f\"数据加载失败: {str(e)}\")\n        return\n\n    dataloader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True, \n                          num_workers=2, pin_memory=True)\n    \n    generator = Generator().to(DEVICE)\n    discriminator = Discriminator().to(DEVICE)\n    \n    optimizer_G = optim.Adam(generator.parameters(), lr=LR, betas=(0.5, 0.9))\n    optimizer_D = optim.Adam(discriminator.parameters(), lr=3e-5, betas=(0.5, 0.9))\n    \n    for epoch in range(EPOCHS):\n        generator.train()\n        discriminator.train()\n        \n        progress_bar = tqdm(dataloader, desc=f\"Epoch {epoch+1}/{EPOCHS}\")\n        for i, real_imgs in enumerate(progress_bar):\n            real_imgs = real_imgs.to(DEVICE)\n            \n            # 训练判别器\n            optimizer_D.zero_grad()\n            z = torch.randn(real_imgs.size(0), LATENT_DIM).to(DEVICE)\n            fake_imgs = generator(z).detach()\n            \n            gp = compute_gradient_penalty(discriminator, real_imgs.data, fake_imgs.data, LAMBDA_GP)\n            d_loss = -torch.mean(discriminator(real_imgs)) + torch.mean(discriminator(fake_imgs)) + gp\n            d_loss.backward()\n            optimizer_D.step()\n            \n            # 训练生成器\n            if i % N_CRITIC == 0:\n                optimizer_G.zero_grad()\n                gen_imgs = generator(z)\n                g_loss = -torch.mean(discriminator(gen_imgs))\n                g_loss.backward()\n                optimizer_G.step()\n            \n            progress_bar.set_postfix({\n                \"D_loss\": f\"{d_loss.item():.4f}\",\n                \"G_loss\": f\"{g_loss.item():.4f}\",\n                \"GP\": f\"{gp.item():.4f}\",\n                \"λ_GP\": LAMBDA_GP\n            })\n\n        # 保存检查点\n        if epoch % 2 == 0:\n            torch.save({\n                'generator': generator.state_dict(),\n                'discriminator': discriminator.state_dict(),\n                'epoch': epoch\n            }, f\"/kaggle/working/checkpoints/epoch_{epoch}.pth\")\n            \n            # 生成示例图像\n            with torch.no_grad():\n                test_z = torch.randn(16, LATENT_DIM).to(DEVICE)\n                gen_imgs = generator(test_z)\n                save_image(gen_imgs, f\"/kaggle/working/samples/epoch_{epoch}.png\")\n                \n                # 实时显示最新生成图像\n                plt.figure(figsize=(10,10))\n                plt.imshow(utils.make_grid(gen_imgs, nrow=4, padding=2, normalize=True).permute(1,2,0).cpu())\n                plt.axis('off')\n                plt.title(f\"Epoch {epoch} Generated Samples\")\n                plt.show()\n\nif __name__ == \"__main__\":\n    try:\n        train()\n    except Exception as e:\n        print(f\"训练中断: {str(e)}\")\n        if 'generator' in locals():\n            torch.save(generator.state_dict(), \"/kaggle/working/emergency_generator.pth\")\n            print(\"已保存紧急备份模型\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T07:14:50.746092Z","iopub.execute_input":"2025-06-30T07:14:50.746316Z","iopub.status.idle":"2025-06-30T09:04:41.740654Z","shell.execute_reply.started":"2025-06-30T07:14:50.746297Z","shell.execute_reply":"2025-06-30T09:04:41.739896Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls /kaggle/working/checkpoints/  \n!ls /kaggle/working/  # 查看是否有generator.pth","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T09:10:49.586793Z","iopub.execute_input":"2025-06-30T09:10:49.587159Z","iopub.status.idle":"2025-06-30T09:10:49.888328Z","shell.execute_reply.started":"2025-06-30T09:10:49.587130Z","shell.execute_reply":"2025-06-30T09:10:49.887413Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 创建最终模型目录\n!mkdir -p /kaggle/working/final_model\n\n# 复制最新检查点并重命名\n!cp /kaggle/working/checkpoints/epoch_48.pth /kaggle/working/final_model/model.pth\n\n# 生成元数据文件\nwith open('/kaggle/working/final_model/dataset-metadata.json', 'w') as f:\n    f.write(\"\"\"{\n    \"title\": \"RSNA_Pneumonia_GAN_Final_Model\",\n    \"id\": \"alanawangmunuk/pneumonia-gan-final\",\n    \"licenses\": [{\n        \"name\": \"CC0-1.0\"\n    }]\n}\"\"\")\n\n# 验证\n!ls -lh /kaggle/working/final_model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T09:14:35.113566Z","iopub.execute_input":"2025-06-30T09:14:35.114399Z","iopub.status.idle":"2025-06-30T09:14:35.633112Z","shell.execute_reply.started":"2025-06-30T09:14:35.114360Z","shell.execute_reply":"2025-06-30T09:14:35.632417Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!stat /kaggle/working/final_model/model.pth","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T09:16:23.182885Z","iopub.execute_input":"2025-06-30T09:16:23.183465Z","iopub.status.idle":"2025-06-30T09:16:23.337651Z","shell.execute_reply.started":"2025-06-30T09:16:23.183441Z","shell.execute_reply":"2025-06-30T09:16:23.336938Z"}},"outputs":[],"execution_count":null}]}