{"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":"none","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nimport glob\nimport random\nimport cv2\nfrom tqdm import tqdm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import Adam\nfrom sklearn.model_selection import train_test_split\nfrom torchvision import transforms\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom sklearn.metrics import precision_recall_curve, average_precision_score\nfrom PIL import Image\n\n# 检测并设置可用的硬件加速器\ndef setup_device():\n    \"\"\"\n    设置最佳可用设备：TPU > GPU > CPU\n    返回设备和适当的PyTorch设备\n    \"\"\"\n    # 检查是否在Kaggle TPU环境\n    try:\n        import torch_xla.core.xla_model as xm\n        print(\"TPU 可用，将使用 TPU 加速训练\")\n        device = xm.xla_device()\n        return \"tpu\", device\n    except ImportError:\n        pass\n    \n    # 检查 GPU 是否可用\n    if torch.cuda.is_available():\n        device_count = torch.cuda.device_count()\n        print(f\"找到 {device_count} 个 GPU 设备\")\n        for i in range(device_count):\n            gpu_name = torch.cuda.get_device_name(i)\n            print(f\"GPU {i}: {gpu_name}\")\n        \n        device = torch.device(\"cuda:0\")\n        # 设置性能优化\n        torch.backends.cudnn.benchmark = True\n        print(f\"将使用 GPU: {torch.cuda.get_device_name(0)}\")\n        return \"gpu\", device\n    \n    # 如果没有可用的加速器，则使用CPU\n    print(\"未找到GPU或TPU，将使用CPU\")\n    return \"cpu\", torch.device(\"cpu\")\n\n# 设置随机种子确保可复现性\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed()\n\n# 定义数据路径\nBASE_PATH = Path(\"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025\")\nTRAIN_PATH = BASE_PATH / \"train\"\nTEST_PATH = BASE_PATH / \"test\"\nSAMPLE_SUBMISSION = BASE_PATH / \"sample_submission.csv\"\nTRAIN_LABELS = BASE_PATH / \"train_labels.csv\"\n\n# 设置设备\naccelerator_type, device = setup_device()\nprint(f\"使用加速器类型: {accelerator_type}\")\nprint(f\"使用设备: {device}\")\n\n# 读取训练标签和检查实际列名\ntry:\n    train_labels_df = pd.read_csv(TRAIN_LABELS)\n    print(f\"标签数据形状: {train_labels_df.shape}\")\n    print(\"标签文件列名:\")\n    print(train_labels_df.columns.tolist())\n    print(\"标签数据前5行:\")\n    print(train_labels_df.head())\nexcept Exception as e:\n    print(f\"读取标签文件时出错: {e}\")\n\n# 读取样本提交文件\ntry:\n    sample_submission_df = pd.read_csv(SAMPLE_SUBMISSION)\n    print(f\"样本提交模板形状: {sample_submission_df.shape}\")\n    print(\"样本提交文件列名:\")\n    print(sample_submission_df.columns.tolist())\n    print(\"样本提交前5行:\")\n    print(sample_submission_df.head())\nexcept Exception as e:\n    print(f\"读取样本提交文件时出错: {e}\")\n\n# 获取所有断层扫描文件夹和切片\ndef get_data_paths():\n    # 获取训练数据断层扫描文件夹\n    train_tomogram_folders = [f for f in TRAIN_PATH.iterdir() if f.is_dir()]\n    train_tomogram_folders.sort()  # 确保一致的顺序\n    \n    # 获取测试数据断层扫描文件夹\n    test_tomogram_folders = [f for f in TEST_PATH.iterdir() if f.is_dir()]\n    test_tomogram_folders.sort()\n    \n    # 获取所有切片路径\n    train_slices = []\n    for folder in train_tomogram_folders:\n        slices = list(folder.glob(\"*.jpg\"))\n        slices.sort()  # 确保切片按顺序\n        train_slices.extend(slices)\n    \n    test_slices = []\n    for folder in test_tomogram_folders:\n        slices = list(folder.glob(\"*.jpg\"))\n        slices.sort()\n        test_slices.extend(slices)\n    \n    return train_tomogram_folders, test_tomogram_folders, train_slices, test_slices\n\ntrain_tomogram_folders, test_tomogram_folders, train_slices, test_slices = get_data_paths()\nprint(f\"训练断层扫描数量: {len(train_tomogram_folders)}\")\nprint(f\"测试断层扫描数量: {len(test_tomogram_folders)}\")\nprint(f\"训练切片总数: {len(train_slices)}\")\nprint(f\"测试切片总数: {len(test_slices)}\")\n\n# 查看一个训练文件夹的名称作为示例\nif train_tomogram_folders:\n    print(f\"示例训练断层扫描文件夹名称: {train_tomogram_folders[0].name}\")\n    # 查看该文件夹下的切片文件名\n    slices = list(train_tomogram_folders[0].glob(\"*.jpg\"))\n    if slices:\n        print(f\"示例切片文件名: {slices[0].name}\")\n\n# 从文件路径提取断层扫描ID和切片索引\ndef extract_tomo_slice_info(file_path):\n    # 从路径中提取断层扫描ID\n    tomo_id = file_path.parent.name\n    # 从文件名中提取切片索引\n    slice_idx = int(file_path.stem.split('_')[1])\n    return tomo_id, slice_idx\n\n# 预处理标签数据 - 根据实际列名调整\ndef preprocess_labels(labels_df):\n    # 检查和调整列名\n    # 假设标签文件有tomo_name/id, Motor axis列和一些其他列，我们手动决定如何处理\n    tomo_col = 'tomo_id'  # tomo_id 列\n    slice_col = 'row_id'  # 假设使用 row_id 作为切片索引列\n    x_col = 'Motor axis 0'\n    y_col = 'Motor axis 1'\n    \n    # 创建字典来存储切片的标签数据\n    labels_dict = {}\n    \n    # 遍历标签数据\n    for _, row in labels_df.iterrows():\n        try:\n            tomo_id = str(row[tomo_col])  # 获取tomo_id\n            slice_idx = int(row[slice_col])  # 使用row_id作为slice索引\n            x, y = float(row[x_col]), float(row[y_col])\n            \n            key = (tomo_id, slice_idx)\n            if key not in labels_dict:\n                labels_dict[key] = []\n            \n            # 将(x, y)坐标添加到对应的(tomo_id, slice_idx)的键中\n            labels_dict[key].append((x, y))\n        \n        except Exception as e:\n            print(f\"处理行时出错: {e}\")\n    \n    return labels_dict\n\n# 首先检查一下数据\nif train_labels_df is not None:\n    labels_dict = preprocess_labels(train_labels_df)\n    print(f\"有标签的切片数量: {len(labels_dict)}\")\n    # 显示前几个标签示例\n    if labels_dict:\n        count = 0\n        for key, points in labels_dict.items():\n            print(f\"切片 {key}: {len(points)} 个标记点\")\n            count += 1\n            if count >= 3:\n                break\n\n# 创建热图标签\ndef create_heatmap(img_shape, points, sigma=10):\n    \"\"\"\n    为给定的点创建高斯热图\n    \n    参数:\n    - img_shape: 原始图像形状 (高, 宽)\n    - points: 坐标点列表 [(x1, y1), (x2, y2), ...]\n    - sigma: 高斯核的标准差\n    \n    返回:\n    - 热图\n    \"\"\"\n    heatmap = np.zeros(img_shape, dtype=np.float32)\n    \n    # 如果没有点，返回全零热图\n    if len(points) == 0:\n        return heatmap\n    \n    for x, y in points:\n        # 确保坐标在图像范围内\n        if 0 <= x < img_shape[1] and 0 <= y < img_shape[0]:\n            # 生成以点为中心的高斯衰减\n            x_grid, y_grid = np.meshgrid(\n                np.arange(img_shape[1]), \n                np.arange(img_shape[0])\n            )\n            \n            # 计算高斯值\n            gaussian = np.exp(-((x_grid - x)**2 + (y_grid - y)**2) / (2 * sigma**2))\n            \n            # 更新热图，取最大值以避免覆盖\n            heatmap = np.maximum(heatmap, gaussian)\n    \n    return heatmap\n\n# 定义自定义数据集\nclass BacterialMotorDataset(Dataset):\n    def __init__(self, image_paths, labels_dict=None, transform=None, is_test=False):\n        self.image_paths = image_paths\n        self.labels_dict = labels_dict\n        self.transform = transform\n        self.is_test = is_test\n        \n    def __len__(self):\n        return len(self.image_paths)\n    \n    def __getitem__(self, idx):\n        img_path = self.image_paths[idx]\n        \n        # 读取图像并转为灰度\n        img = cv2.imread(str(img_path), cv2.IMREAD_GRAYSCALE)\n        \n        # Convert numpy ndarray to PIL Image before transformations\n        img = Image.fromarray(img)  # Convert to PIL Image\n\n        img_shape = img.size  # Use PIL's size (width, height)\n        \n        # 提取断层扫描ID和切片索引\n        tomo_id, slice_idx = extract_tomo_slice_info(img_path)\n        \n        # 应用变换\n        if self.transform:\n            img = self.transform(img)\n        \n        # 如果是测试集，只返回图像和元信息\n        if self.is_test:\n            return {\n                'image': img,\n                'tomo_id': tomo_id,\n                'slice_idx': slice_idx,\n                'image_path': str(img_path)\n            }\n        \n        # 获取标签点，如果没有则为空列表\n        key = (tomo_id, slice_idx)\n        points = self.labels_dict.get(key, [])\n        \n        # 创建热图标签\n        heatmap = create_heatmap(img_shape, points)\n        \n        # 转换为张量\n        heatmap = torch.tensor(heatmap, dtype=torch.float32).unsqueeze(0)  # 添加通道维度\n        \n        return {\n            'image': img,\n            'heatmap': heatmap,\n            'points': points,\n            'tomo_id': tomo_id,\n            'slice_idx': slice_idx,\n            'image_path': str(img_path)\n        }\n\n# 图像转换包括Resize以确保统一大小\ntrain_transform = transforms.Compose([\n    transforms.Resize((256, 256)),  # Resize all images to the same shape\n    transforms.ToTensor(),\n])\n\nval_transform = transforms.Compose([\n    transforms.Resize((256, 256)),  # Resize all images to the same shape\n    transforms.ToTensor(),\n])\n\n# Define a custom collate function to handle padding of images with different sizes\ndef custom_collate_fn(batch):\n    images = [item['image'] for item in batch]\n    heatmaps = [item['heatmap'] for item in batch]\n\n    # Find the maximum height and width in the batch\n    max_height = max(img.shape[1] for img in images)\n    max_width = max(img.shape[2] for img in images)\n\n    # Resize or pad the images and heatmaps to the max height and width\n    resized_images = []\n    resized_heatmaps = []\n    \n    for img, heatmap in zip(images, heatmaps):\n        # Resize or pad the images and heatmaps to the max height and width\n        padded_img = F.pad(img, (0, max_width - img.shape[2], 0, max_height - img.shape[1]), \"constant\", 0)\n        padded_heatmap = F.pad(heatmap, (0, max_width - heatmap.shape[2], 0, max_height - heatmap.shape[1]), \"constant\", 0)\n        \n        resized_images.append(padded_img)\n        resized_heatmaps.append(padded_heatmap)\n\n    # Stack the images and heatmaps tensors\n    images = torch.stack(resized_images, dim=0)\n    heatmaps = torch.stack(resized_heatmaps, dim=0)\n\n    return {'image': images, 'heatmap': heatmaps}\n\n# Define U-Net Model\nclass DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(DoubleConv, self).__init__()\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.double_conv(x)\n\nclass UNet(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1, init_features=32):\n        super(UNet, self).__init__()\n        \n        features = init_features\n        self.encoder1 = DoubleConv(in_channels, features)\n        self.pool1 = nn.MaxPool2d(kernel_size=2, stride=2)\n        \n        self.encoder2 = DoubleConv(features, features * 2)\n        self.pool2 = nn.MaxPool2d(kernel_size=2, stride=2)\n        \n        self.encoder3 = DoubleConv(features * 2, features * 4)\n        self.pool3 = nn.MaxPool2d(kernel_size=2, stride=2)\n        \n        self.encoder4 = DoubleConv(features * 4, features * 8)\n        self.pool4 = nn.MaxPool2d(kernel_size=2, stride=2)\n        \n        self.bottleneck = DoubleConv(features * 8, features * 16)\n        \n        self.upconv4 = nn.ConvTranspose2d(features * 16, features * 8, kernel_size=2, stride=2)\n        self.decoder4 = DoubleConv(features * 16, features * 8)\n        \n        self.upconv3 = nn.ConvTranspose2d(features * 8, features * 4, kernel_size=2, stride=2)\n        self.decoder3 = DoubleConv(features * 8, features * 4)\n        \n        self.upconv2 = nn.ConvTranspose2d(features * 4, features * 2, kernel_size=2, stride=2)\n        self.decoder2 = DoubleConv(features * 4, features * 2)\n        \n        self.upconv1 = nn.ConvTranspose2d(features * 2, features, kernel_size=2, stride=2)\n        self.decoder1 = DoubleConv(features * 2, features)\n        \n        self.conv = nn.Conv2d(features, out_channels, kernel_size=1)\n        \n    def forward(self, x):\n        enc1 = self.encoder1(x)\n        enc2 = self.encoder2(self.pool1(enc1))\n        enc3 = self.encoder3(self.pool2(enc2))\n        enc4 = self.encoder4(self.pool3(enc3))\n        \n        bottleneck = self.bottleneck(self.pool4(enc4))\n        \n        dec4 = self.upconv4(bottleneck)\n        dec4 = torch.cat((dec4, enc4), dim=1)\n        dec4 = self.decoder4(dec4)\n        \n        dec3 = self.upconv3(dec4)\n        dec3 = torch.cat((dec3, enc3), dim=1)\n        dec3 = self.decoder3(dec3)\n        \n        dec2 = self.upconv2(dec3)\n        dec2 = torch.cat((dec2, enc2), dim=1)\n        dec2 = self.decoder2(dec2)\n        \n        dec1 = self.upconv1(dec2)\n        dec1 = torch.cat((dec1, enc1), dim=1)\n        dec1 = self.decoder1(dec1)\n        \n        output = self.conv(dec1)\n        output = torch.sigmoid(output)  # 使用sigmoid确保输出在[0,1]范围\n        \n        return output\n\n# 定义训练函数 - 针对不同加速器\ndef train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, num_epochs=10, accelerator_type='cpu'):\n    best_val_loss = float('inf')\n    history = {'train_loss': [], 'val_loss': []}\n    \n    # 如果使用TPU\n    if accelerator_type == 'tpu':\n        import torch_xla.core.xla_model as xm\n        import torch_xla.distributed.parallel_loader as pl\n    \n    for epoch in range(num_epochs):\n        # 训练阶段\n        model.train()\n        train_loss = 0.0\n        \n        # 根据加速器类型选择合适的数据加载器\n        if accelerator_type == 'tpu':\n            train_device_loader = pl.ParallelLoader(train_loader, [device]).per_device_loader(device)\n            loader_to_use = train_device_loader\n        else:\n            loader_to_use = train_loader\n        \n        for batch in tqdm(loader_to_use, desc=f'Epoch {epoch+1}/{num_epochs} [Train]'):\n            images = batch['image'].to(device)\n            heatmaps = batch['heatmap'].to(device)\n            \n            # 前向传播\n            outputs = model(images)\n            loss = criterion(outputs, heatmaps)\n            \n            # 反向传播和优化\n            optimizer.zero_grad()\n            loss.backward()\n            \n            if accelerator_type == 'tpu':\n                xm.optimizer_step(optimizer, barrier=True)\n            else:\n                optimizer.step()\n            \n            train_loss += loss.item() * images.size(0)\n        \n        train_loss /= len(train_loader.dataset)\n        \n        # 验证阶段\n        model.eval()\n        val_loss = 0.0\n        \n        # 根据加速器类型选择合适的数据加载器\n        if accelerator_type == 'tpu':\n            val_device_loader = pl.ParallelLoader(val_loader, [device]).per_device_loader(device)\n            val_loader_to_use = val_device_loader\n        else:\n            val_loader_to_use = val_loader\n        \n        with torch.no_grad():\n            for batch in tqdm(val_loader_to_use, desc=f'Epoch {epoch+1}/{num_epochs} [Val]'):\n                images = batch['image'].to(device)\n                heatmaps = batch['heatmap'].to(device)\n                \n                outputs = model(images)\n                loss = criterion(outputs, heatmaps)\n                \n                val_loss += loss.item() * images.size(0)\n            \n            val_loss /= len(val_loader.dataset)\n        \n        # 更新学习率\n        if accelerator_type == 'tpu':\n            # 在TPU上，我们需要同步所有核心上的损失\n            val_loss_for_scheduler = xm.mesh_reduce('val_loss', val_loss, lambda x: sum(x) / len(x))\n            scheduler.step(val_loss_for_scheduler)\n        else:\n            scheduler.step(val_loss)\n        \n        # 保存最佳模型\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            # 在TPU上，只在主进程上保存\n            if accelerator_type == 'tpu':\n                if xm.is_master_ordinal():\n                    xm.save(model.state_dict(), 'best_model.pth')\n                    print(f'模型已保存: best_model.pth')\n            else:\n                torch.save(model.state_dict(), 'best_model.pth')\n                print(f'模型已保存: best_model.pth')\n        \n        # 记录损失\n        history['train_loss'].append(train_loss)\n        history['val_loss'].append(val_loss)\n        \n        print(f'Epoch {epoch+1}/{num_epochs}, Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}')\n    \n    return history\n\n# 定义预测函数\ndef predict(model, test_loader, submission_columns, accelerator_type='cpu'):\n    model.eval()\n    predictions = []\n    \n    if accelerator_type == 'tpu':\n        import torch_xla.core.xla_model as xm\n        import torch_xla.distributed.parallel_loader as pl\n        test_device_loader = pl.ParallelLoader(test_loader, [device]).per_device_loader(device)\n        loader_to_use = test_device_loader\n    else:\n        loader_to_use = test_loader\n    \n    with torch.no_grad():\n        for batch in tqdm(loader_to_use, desc='Predicting'):\n            images = batch['image'].to(device)\n            tomo_ids = batch['tomo_id']\n            slice_idxs = batch['slice_idx']\n            \n            # 获取预测热图\n            predicted_heatmaps = model(images)\n            \n            # 从热图中提取最可能的点\n            batch_size = images.size(0)\n            \n            for i in range(batch_size):\n                tomo_id = tomo_ids[i]\n                slice_idx = slice_idxs[i]\n                heatmap = predicted_heatmaps[i, 0].cpu().numpy()\n                \n                # 找到热图中的局部最大值 (可能的电机位置)\n                # 这里使用一个简单的阈值和最大值查找\n                threshold = 0.5  # 可调\n                heatmap_binary = (heatmap > threshold).astype(np.uint8)\n                \n                # 使用OpenCV查找连通区域\n                num_labels, labels, stats, centroids = cv2.connectedComponentsWithStats(heatmap_binary, connectivity=8)\n                \n                # 跳过背景 (第一个组件通常是背景)\n                for j in range(1, num_labels):\n                    x, y = centroids[j]\n                    score = heatmap[int(y), int(x)]  # 获取置信度分数\n                    \n                    predictions.append({\n                        'tomo_id': tomo_id,\n                        'slice_idx': slice_idx,\n                        'x': x,\n                        'y': y,\n                        'confidence': score\n                    })\n    \n    # 创建DataFrame并格式化\n    predictions_df = pd.DataFrame(predictions)\n    \n    # 处理提交格式\n    if len(predictions_df) > 0:\n        # 将预测结果映射到提交格式\n        submission_df = pd.DataFrame(columns=submission_columns)\n        \n        # 这里需要根据实际比赛的提交要求进行调整\n        # 假设需要的列是 'tomo_id', 'row_id', 'Motor axis 0', 'Motor axis 1'\n        submission_df['tomo_id'] = predictions_df['tomo_id']\n        submission_df['row_id'] = predictions_df['slice_idx']\n        submission_df['Motor axis 0'] = predictions_df['x']\n        submission_df['Motor axis 1'] = predictions_df['y']\n        \n        return submission_df\n    else:\n        print(\"警告: 没有找到预测点\")\n        return pd.DataFrame(columns=submission_columns)\n\n# 主函数\ndef main():\n    # 拆分训练和验证集\n    train_tomo_folders, val_tomo_folders = train_test_split(\n        train_tomogram_folders, test_size=0.2, random_state=42\n    )\n    \n    # 获取相应的切片路径\n    train_slices_list = []\n    for folder in train_tomo_folders:\n        train_slices_list.extend(list(folder.glob(\"*.jpg\")))\n    \n    val_slices_list = []\n    for folder in val_tomo_folders:\n        val_slices_list.extend(list(folder.glob(\"*.jpg\")))\n    \n    # 图像转换\n    train_transform = transforms.Compose([\n        transforms.Resize((256, 256)),  # Resize all images to the same shape\n        transforms.ToTensor(),\n    ])\n    \n    val_transform = transforms.Compose([\n        transforms.Resize((256, 256)),  # Resize all images to the same shape\n        transforms.ToTensor(),\n    ])\n    \n    # 创建数据集和数据加载器\n    train_dataset = BacterialMotorDataset(\n        train_slices_list, labels_dict, transform=train_transform\n    )\n    \n    val_dataset = BacterialMotorDataset(\n        val_slices_list, labels_dict, transform=val_transform\n    )\n    \n    test_dataset = BacterialMotorDataset(\n        test_slices, transform=val_transform, is_test=True\n    )\n    \n    # 设置适当的工作线程数\n    # 对于TPU和多GPU，减少工作线程数以避免瓶颈\n    if accelerator_type == 'tpu':\n        num_workers = 4\n    elif accelerator_type == 'gpu':\n        num_workers = 4\n    else:\n        num_workers = 2\n    \n    # 创建数据加载器\n    train_loader = DataLoader(\n        train_dataset, batch_size=8, shuffle=True, num_workers=num_workers, collate_fn=custom_collate_fn\n    )\n    \n    val_loader = DataLoader(\n        val_dataset, batch_size=8, shuffle=False, num_workers=num_workers, collate_fn=custom_collate_fn\n    )\n    \n    test_loader = DataLoader(\n        test_dataset, batch_size=8, shuffle=False, num_workers=num_workers, collate_fn=custom_collate_fn\n    )\n    \n    # 初始化模型\n    model = UNet(in_channels=1, out_channels=1, init_features=32).to(device)\n    \n    # 在GPU上使用混合精度训练以加速\n    if accelerator_type == 'gpu':\n        try:\n            from torch.cuda.amp import GradScaler, autocast\n            use_amp = True\n            scaler = GradScaler()\n            print(\"启用混合精度训练\")\n        except ImportError:\n            use_amp = False\n            scaler = None\n            print(\"无法启用混合精度训练，使用全精度\")\n    else:\n        use_amp = False\n        scaler = None\n    \n    # 定义损失函数和优化器\n    criterion = nn.MSELoss()\n    optimizer = Adam(model.parameters(), lr=1e-3)\n    scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=3, verbose=True)\n    \n    # 训练模型\n    print(\"开始训练模型...\")\n    history = train_model(\n        # model, train_loader, val_loader, criterion, optimizer, scheduler, \n        # num_epochs=15, accelerator_type=accelerator_type\n        model, train_loader, val_loader, criterion, optimizer, scheduler, \n        num_epochs=10, accelerator_type=accelerator_type\n    )\n    \n    # 绘制损失曲线\n    plt.figure(figsize=(10, 5))\n    plt.plot(history['train_loss'], label='Train Loss')\n    plt.plot(history['val_loss'], label='Validation Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.title('Training and Validation Loss')\n    plt.savefig('loss_curve.png')\n    plt.show()\n    \n    # 加载最佳模型\n    if accelerator_type == 'tpu':\n        import torch_xla.core.xla_model as xm\n        model.load_state_dict(torch.load('best_model.pth', map_location=device))\n    else:\n        model.load_state_dict(torch.load('best_model.pth', map_location=device))\n    \n    # 确保提交列名与样本提交文件匹配\n    submission_columns = sample_submission_df.columns.tolist()\n    \n    # 进行预测\n    print(\"开始预测...\")\n    predictions_df = predict(model, test_loader, submission_columns, accelerator_type=accelerator_type)\n    \n    # 保存预测结果\n    predictions_df.to_csv('submission.csv', index=False)\n    print(f'预测已保存，共 {len(predictions_df)} 个预测点')\n    print(predictions_df.head())\n\n# 修改训练循环以支持混合精度训练\ndef train_model_with_amp(model, train_loader, val_loader, criterion, optimizer, scheduler, num_epochs=10):\n    \"\"\"使用混合精度训练的版本\"\"\"\n    from torch.cuda.amp import GradScaler, autocast\n    \n    scaler = GradScaler()\n    best_val_loss = float('inf')\n    history = {'train_loss': [], 'val_loss': []}\n    \n    for epoch in range(num_epochs):\n        # 训练阶段\n        model.train()\n        train_loss = 0.0\n        \n        for batch in tqdm(train_loader, desc=f'Epoch {epoch+1}/{num_epochs} [Train]'):\n            images = batch['image'].to(device)\n            heatmaps = batch['heatmap'].to(device)\n            \n            # 使用autocast进行混合精度训练\n            with autocast():\n                outputs = model(images)\n                loss = criterion(outputs, heatmaps)\n            \n            # 梯度缩放\n            optimizer.zero_grad()\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            \n            train_loss += loss.item() * images.size(0)\n        \n        train_loss /= len(train_loader.dataset)\n        \n        # 验证阶段\n        model.eval()\n        val_loss = 0.0\n        \n        with torch.no_grad():\n            for batch in tqdm(val_loader, desc=f'Epoch {epoch+1}/{num_epochs} [Val]'):\n                images = batch['image'].to(device)\n                heatmaps = batch['heatmap'].to(device)\n                \n                with autocast():\n                    outputs = model(images)\n                    loss = criterion(outputs, heatmaps)\n                \n                val_loss += loss.item() * images.size(0)\n            \n            val_loss /= len(val_loader.dataset)\n        \n        # 更新学习率\n        scheduler.step(val_loss)\n        \n        # 保存最佳模型\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            torch.save(model.state_dict(), 'best_model.pth')\n            print(f'模型已保存: best_model.pth')\n        \n        # 记录损失\n        history['train_loss'].append(train_loss)\n        history['val_loss'].append(val_loss)\n        \n        print(f'Epoch {epoch+1}/{num_epochs}, Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}')\n    \n    return history\n\nif __name__ == \"__main__\":\n    # 检查是否已定义predict函数中所需的submission_columns变量\n    # 如果未定义，则从sample_submission_df中获取\n    try:\n        submission_columns\n    except NameError:\n        submission_columns = sample_submission_df.columns.tolist() if sample_submission_df is not None else None\n        if submission_columns is None:\n            print(\"警告: 无法获取提交列名! 请确保样本提交文件可用.\")\n            submission_columns = ['tomo_id', 'row_id', 'Motor axis 0', 'Motor axis 1']  # 假设的默认列名\n    \n    # 确保使用指定的加速器类型\n    print(f\"使用 {accelerator_type} 加速器进行训练和预测\")\n    \n    if accelerator_type == 'gpu' and torch.cuda.is_available():\n        try:\n            from torch.cuda.amp import GradScaler, autocast\n            print(\"启用GPU混合精度训练\")\n            use_amp = True\n        except ImportError:\n            print(\"无法启用混合精度训练\")\n            use_amp = False\n    else:\n        use_amp = False\n    \n    main()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-05T06:37:27.566351Z","iopub.execute_input":"2025-05-05T06:37:27.567084Z","iopub.status.idle":"2025-05-05T06:38:07.488546Z","shell.execute_reply.started":"2025-05-05T06:37:27.567055Z","shell.execute_reply":"2025-05-05T06:38:07.487542Z"}},"outputs":[],"execution_count":null}]}