{"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":407317,"sourceType":"datasetVersion","datasetId":181273}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Step 1: 完全卸载冲突的 TensorFlow 和旧版 tensorboard\n!pip uninstall -y tensorflow tensorflow-cpu tensorflow-gpu tensorboard 2>/dev/null || true\n\n# Step 2: 安装 PyTorch 兼容的 tensorboard\n!pip install tensorboard\n\n# Step 3: 强制降级 protobuf 到兼容版本（关键！）\n!pip install \"protobuf<4.24\" --force-reinstall\n\n# Step 4: 清除缓存\nimport os\nos.system(\"rm -rf /root/.cache/pip\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-08T08:23:08.680817Z","iopub.execute_input":"2025-11-08T08:23:08.681184Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# 添加CUDA警告抑制\nimport os\n# 指定仅使用第一个GPU，避免多GPU检测冲突，同时移除可能与PyTorch冲突的TensorFlow配置\nos.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0\"\n# 导入必要的库\nimport numpy as np # 用于数组运算和数值处理\nimport torch # PyTorch深度学习框架核心库\nimport torch.nn as nn # 神经网络层与模块\nimport torch.nn.functional as F # 神经网络功能函数（激活函数、卷积操作等）\nfrom torch.utils.data import Dataset, DataLoader # 数据集管理和批量加载工具\nfrom torch.utils.tensorboard import SummaryWriter # 训练过程可视化工具\nimport matplotlib.pyplot as plt # 绘图工具，用于生成训练曲线和结果可视化\nimport seaborn as sns # 增强版绘图工具，优化图表样式\nfrom sklearn.model_selection import train_test_split # 数据集拆分工具\nfrom sklearn.metrics import confusion_matrix # 计算混淆矩阵，用于评估分类性能\nimport cv2 # OpenCV库，处理图像读写和基本变换\nimport albumentations as A # 高效数据增强库，支持多种图像变换\nfrom albumentations.pytorch import ToTensorV2 # 将图像转换为PyTorch张量的工具\nimport torch.optim as optim # 优化器模块，实现各种参数更新算法\nfrom tqdm import tqdm # 进度条工具，直观显示训练和评估进度\nimport threading # 用于掩码读取超时控制\nimport time # 计时工具，用于统计训练耗时和超时控制\n# 修复混合精度训练的导入问题（兼容不同PyTorch版本）\ntry:\n    from torch.amp import GradScaler, autocast # PyTorch 1.10+的混合精度接口\nexcept ImportError:\n    from torch.cuda.amp import GradScaler, autocast # 旧版本接口\nfrom scipy.spatial.distance import directed_hausdorff # 用于计算Hausdorff距离（暂未使用）\n# 使用IPython的魔术命令清理目录\n!rm -rf /kaggle/working/*\n# ==============================\n# 1. 基础配置与设备优化\n# 作用：设置随机种子确保实验可复现性，自动选择计算设备（GPU/CPU）\n# ==============================\ndef set_seed(seed=42):\n    \"\"\"设置随机种子，保证实验结果在不同运行中可复现\"\"\"\n    torch.manual_seed(seed) # 设置PyTorch随机种子\n    np.random.seed(seed) # 设置NumPy随机种子\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed) # 设置当前GPU随机种子\n        torch.cuda.manual_seed_all(seed) # 设置所有GPU随机种子\n    # 确保CUDA卷积算法固定，避免因算法选择导致的结果差异\n    torch.backends.cudnn.deterministic = True\n    # 禁用自动选择最快卷积算法（可能影响复现性）\n    torch.backends.cudnn.benchmark = False\n# 初始化随机种子，保证实验可复现\nset_seed()\ndef get_device():\n    \"\"\"自动检测并返回可用的计算设备（优先使用GPU以加速训练）\"\"\"\n    if torch.cuda.is_available():\n        device = torch.device(\"cuda\") # 使用GPU\n        print(f\"使用GPU设备: {torch.cuda.get_device_name(0)}\")\n        print(f\"可用GPU数量: {torch.cuda.device_count()}\")\n        print(f\"GPU内存: {torch.cuda.get_device_properties(0).total_memory / 1024 ** 3:.2f} GB\")\n    else:\n        device = torch.device(\"cpu\") # 无GPU时使用CPU\n        print(\"警告: 未检测到GPU，训练将非常缓慢\")\n    return device\n# 获取计算设备并标记是否使用CUDA\ndevice = get_device()\nuse_cuda = device.type == 'cuda' # 布尔值：是否使用GPU加速\n# ==============================\n# 2. 模型组件（含注意力权重保存）\n# 作用：定义轻量级多尺度注意力UNet的各个模块，兼顾性能与效率\n# ==============================\nclass DepthwiseSeparableConv(nn.Module):\n    \"\"\"深度可分离卷积：将标准卷积拆分为深度卷积和逐点卷积，减少参数和计算量\n    适用于资源有限场景，在保持精度的同时降低计算成本\n    \"\"\"\n    def __init__(self, in_channels, out_channels, kernel_size=3, stride=1, dilation=1):\n        super().__init__()\n        # 计算填充大小，确保输出尺寸与输入一致（当stride=1时）\n        padding = (kernel_size - 1) // 2 * dilation\n        # 深度卷积：每个输入通道单独卷积（groups=in_channels）\n        self.depthwise = nn.Conv2d(\n            in_channels, in_channels,\n            kernel_size=kernel_size,\n            stride=stride,\n            padding=padding,\n            dilation=dilation, # 膨胀率（用于扩大感受野）\n            groups=in_channels, # 分组卷积：每组对应一个输入通道\n            bias=False # 后续有BN层，可省略偏置\n        )\n        self.bn_depth = nn.BatchNorm2d(in_channels) # 深度卷积后的批归一化\n        # 逐点卷积：1x1卷积，融合不同通道特征\n        self.pointwise = nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)\n        self.bn_point = nn.BatchNorm2d(out_channels) # 逐点卷积后的批归一化\n        self.relu = nn.ReLU6(inplace=True) # ReLU6激活函数（更适合移动端部署）\n    def forward(self, x):\n        \"\"\"前向传播：深度卷积 -> 批归一化 -> 激活 -> 逐点卷积 -> 批归一化 -> 激活\"\"\"\n        x = self.depthwise(x)\n        x = self.bn_depth(x)\n        x = self.relu(x)\n        x = self.pointwise(x)\n        x = self.bn_point(x)\n        return self.relu(x)\nclass LightConvBlock(nn.Module):\n    \"\"\"轻量级卷积块：由两个深度可分离卷积组成，带残差连接（类似ResNet结构）\n    增强特征传播能力，缓解深层网络的梯度消失问题\n    \"\"\"\n    def __init__(self, in_channels, out_channels, stride=1):\n        super().__init__()\n        # 主分支：两个深度可分离卷积\n        self.conv = nn.Sequential(\n            DepthwiseSeparableConv(in_channels, out_channels, stride=stride),\n            DepthwiseSeparableConv(out_channels, out_channels)\n        )\n        # 残差分支：当输入输出通道数不同或步长不为1时，用1x1卷积调整维度\n        self.shortcut = nn.Sequential()\n        if stride != 1 or in_channels != out_channels:\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False),\n                nn.BatchNorm2d(out_channels)\n            )\n        self.relu = nn.ReLU6(inplace=True) # 最终激活函数\n    def forward(self, x):\n        \"\"\"前向传播：主分支输出 + 残差分支输出 -> 激活\"\"\"\n        residual = self.shortcut(x) # 残差\n        x = self.conv(x) # 主分支计算\n        x += residual # 残差连接\n        return self.relu(x)\nclass EnhancedAttentionGate(nn.Module):\n    \"\"\"增强注意力门（含注意力权重保存）：突出编码器与解码器特征的关联区域（如肿瘤区域）\n    让模型更关注有意义的区域，减少背景干扰\n    \"\"\"\n    def __init__(self, low_channels, high_channels, out_channels):\n        super().__init__()\n        # 调整编码器特征（低分辨率）的通道数\n        self.low_adapt = nn.Sequential(\n            nn.Conv2d(low_channels, out_channels, kernel_size=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU6(inplace=True)\n        )\n        # 调整解码器特征（高分辨率）的通道数\n        self.high_adapt = nn.Sequential(\n            nn.Conv2d(high_channels, out_channels, kernel_size=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU6(inplace=True)\n        )\n        # 注意力权重计算：通过卷积生成0-1的权重图\n        self.attention = nn.Sequential(\n            nn.Conv2d(out_channels * 2, out_channels, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU6(inplace=True),\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.Sigmoid() # 输出权重（0-1）\n        )\n        self.attn_weight = None # 保存注意力权重用于可视化\n    def forward(self, low_feat, high_feat):\n        \"\"\"\n        前向传播：\n        1. 将解码器特征上采样到编码器特征尺寸\n        2. 调整两者通道数并融合\n        3. 计算注意力权重，对编码器特征加权\n        \"\"\"\n        # 解码器特征上采样（与编码器特征尺寸一致）\n        high_feat_up = F.interpolate(\n            high_feat, size=low_feat.shape[2:],\n            mode='bilinear', align_corners=False # 双线性插值\n        )\n        # 调整通道数\n        low_adapted = self.low_adapt(low_feat)\n        high_adapted = self.high_adapt(high_feat_up)\n        # 融合特征（相乘+相加，增强关联性）\n        combined = torch.cat([low_adapted * high_adapted, low_adapted + high_adapted], dim=1)\n        # 计算注意力权重并应用到编码器特征\n        self.attn_weight = self.attention(combined) # 保存注意力权重\n        return low_feat * self.attn_weight # 加权后的编码器特征\nclass MultiScaleModule(nn.Module):\n    \"\"\"多尺度模块：通过不同膨胀率的卷积捕捉多尺度特征（适合不同大小的肿瘤）\n    解决肿瘤尺寸差异大的问题，同时扩大感受野\n    \"\"\"\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        assert out_channels % 4 == 0, \"输出通道数必须能被4整除（4个分支）\"\n        branch_channels = out_channels // 4 # 每个分支的通道数\n        # 4个分支：不同膨胀率的深度可分离卷积（感受野依次增大）\n        self.conv1 = DepthwiseSeparableConv(in_channels, branch_channels, kernel_size=7, dilation=1) # 膨胀率1（小感受野）\n        self.conv2 = DepthwiseSeparableConv(in_channels, branch_channels, kernel_size=7, dilation=2) # 膨胀率2\n        self.conv3 = DepthwiseSeparableConv(in_channels, branch_channels, kernel_size=7, dilation=4) # 膨胀率4\n        self.conv4 = DepthwiseSeparableConv(in_channels, branch_channels, kernel_size=7, dilation=8) # 膨胀率8（大感受野）\n        # 尺度注意力：自动调整不同分支的权重\n        self.scale_attention = nn.Sequential(\n            nn.Conv2d(out_channels, out_channels, kernel_size=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.Sigmoid() # 输出各分支的权重\n        )\n        self.bn = nn.BatchNorm2d(out_channels) # 批归一化\n        self.relu = nn.ReLU6(inplace=True) # 激活函数\n    def forward(self, x):\n        \"\"\"前向传播：多分支特征提取 -> 融合 -> 尺度注意力加权\"\"\"\n        x1 = self.conv1(x) # 分支1（小感受野）\n        x2 = self.conv2(x) # 分支2\n        x3 = self.conv3(x) # 分支3\n        x4 = self.conv4(x) # 分支4（大感受野）\n        out = torch.cat([x1, x2, x3, x4], dim=1) # 融合4个分支特征\n        scale_attn = self.scale_attention(out) # 计算尺度权重\n        out = out * scale_attn # 应用权重\n        return self.relu(self.bn(out)) # 批归一化+激活\nclass LightweightMultiScaleAttentionUNet(nn.Module):\n    \"\"\"轻量级多尺度注意力UNet（含中间特征保存）：结合UNet结构、多尺度特征和注意力机制，适合肿瘤分割\n    在保证分割精度的同时，减少模型参数和计算量\n    \"\"\"\n    def __init__(self, img_size=256, in_channels=3, out_channels=1, base_channels=32):\n        super().__init__()\n        self.img_size = img_size # 输入图像尺寸\n        # 编码器（下采样）：逐步减小尺寸，增加通道数\n        self.enc1 = LightConvBlock(in_channels, base_channels) # 输入3通道 -> 32通道\n        self.pool1 = nn.MaxPool2d(2) # 下采样2倍（256->128）\n        self.enc2 = LightConvBlock(base_channels, base_channels * 2) # 32->64通道\n        self.pool2 = nn.MaxPool2d(2) # 128->64\n        self.enc3 = LightConvBlock(base_channels * 2, base_channels * 4) # 64->128通道\n        self.pool3 = nn.MaxPool2d(2) # 64->32\n        self.enc4 = LightConvBlock(base_channels * 4, base_channels * 8) # 128->256通道\n        self.pool4 = nn.MaxPool2d(2) # 32->16\n        # 瓶颈层：多尺度特征提取（捕捉不同大小的肿瘤）\n        self.bottleneck = nn.Sequential(\n            MultiScaleModule(base_channels * 8, base_channels * 8), # 多尺度特征\n            LightConvBlock(base_channels * 8, base_channels * 8) # 特征融合\n        )\n        # 解码器（上采样）：逐步恢复尺寸，结合编码器特征\n        self.up4 = nn.ConvTranspose2d(base_channels * 8, base_channels * 8, kernel_size=2, stride=2) # 上采样2倍（16->32）\n        self.dec4 = LightConvBlock(base_channels * 8 * 2, base_channels * 8) # 融合后256*2->256通道\n        self.att4 = EnhancedAttentionGate(base_channels * 8, base_channels * 8, base_channels * 8) # 注意力门（编码器与解码器特征关联）\n        self.up3 = nn.ConvTranspose2d(base_channels * 8, base_channels * 4, kernel_size=2, stride=2) # 32->64\n        self.dec3 = LightConvBlock(base_channels * 4 * 2, base_channels * 4) # 128*2->128通道\n        self.att3 = EnhancedAttentionGate(base_channels * 4, base_channels * 4, base_channels * 4)\n        self.up2 = nn.ConvTranspose2d(base_channels * 4, base_channels * 2, kernel_size=2, stride=2) # 64->128\n        self.dec2 = LightConvBlock(base_channels * 2 * 2, base_channels * 2) # 64*2->64通道\n        self.att2 = EnhancedAttentionGate(base_channels * 2, base_channels * 2, base_channels * 2)\n        self.up1 = nn.ConvTranspose2d(base_channels * 2, base_channels, kernel_size=2, stride=2) # 128->256\n        self.dec1 = LightConvBlock(base_channels * 2, base_channels) # 32*2->32通道\n        self.att1 = EnhancedAttentionGate(base_channels, base_channels, base_channels)\n        # 精细化输出层：进一步优化分割结果\n        self.refine = nn.Sequential(\n            DepthwiseSeparableConv(base_channels, base_channels),\n            nn.Conv2d(base_channels, base_channels, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(base_channels),\n            nn.ReLU6(inplace=True)\n        )\n        self.final_conv = nn.Conv2d(base_channels, out_channels, 1) # 1x1卷积输出最终结果（1通道：肿瘤/背景）\n        # 保存中间特征（用于可视化）\n        self.mid_features = {\n            'enc4': None, # 编码器最深层特征\n            'bottleneck': None # 多尺度模块输出特征\n        }\n    def forward(self, x):\n        \"\"\"前向传播：编码器 -> 瓶颈层 -> 解码器（结合注意力） -> 输出\"\"\"\n        # 编码器：提取特征并下采样\n        e1 = self.enc1(x) # 256x256, 32通道\n        p1 = self.pool1(e1) # 128x128, 32通道\n        e2 = self.enc2(p1) # 128x128, 64通道\n        p2 = self.pool2(e2) # 64x64, 64通道\n        e3 = self.enc3(p2) # 64x64, 128通道\n        p3 = self.pool3(e3) # 32x32, 128通道\n        e4 = self.enc4(p3) # 32x32, 256通道\n        self.mid_features['enc4'] = e4 # 保存编码器特征\n        p4 = self.pool4(e4) # 16x16, 256通道\n        # 瓶颈层：多尺度特征提取\n        bottleneck = self.bottleneck(p4) # 16x16, 256通道\n        self.mid_features['bottleneck'] = bottleneck # 保存多尺度特征\n        # 解码器：上采样并融合编码器特征（带注意力）\n        d4 = self.up4(bottleneck) # 32x32, 256通道\n        e4_att = self.att4(e4, d4) # 对编码器e4特征加权（注意力）\n        d4 = torch.cat([e4_att, d4], dim=1) # 融合注意力特征和解码器特征（256+256=512通道）\n        d4 = self.dec4(d4) # 32x32, 256通道\n        d3 = self.up3(d4) # 64x64, 128通道\n        e3_att = self.att3(e3, d3) # 对编码器e3特征加权\n        d3 = torch.cat([e3_att, d3], dim=1) # 128+128=256通道\n        d3 = self.dec3(d3) # 64x64, 128通道\n        d2 = self.up2(d3) # 128x128, 64通道\n        e2_att = self.att2(e2, d2) # 对编码器e2特征加权\n        d2 = torch.cat([e2_att, d2], dim=1) # 64+64=128通道\n        d2 = self.dec2(d2) # 128x128, 64通道\n        d1 = self.up1(d2) # 256x256, 32通道\n        e1_att = self.att1(e1, d1) # 对编码器e1特征加权\n        d1 = torch.cat([e1_att, d1], dim=1) # 32+32=64通道\n        d1 = self.dec1(d1) # 256x256, 32通道\n        # 精细化处理并输出\n        d1 = self.refine(d1) # 优化特征\n        out = self.final_conv(d1) # 256x256, 1通道（肿瘤概率）\n        return out\n# ==============================\n# 3. 数据处理\n# 作用：加载、预处理数据，实现数据增强和平衡采样（解决类别不平衡）\n# ==============================\nclass TumorBalancedSampler(torch.utils.data.Sampler):\n    \"\"\"\n    肿瘤平衡采样器（鲁棒版）：确保每个batch包含指定比例的肿瘤样本，解决类别不平衡问题\n    核心优化：\n    1. 动态索引生成，避免DataLoader阻塞；\n    2. 鲁棒的文件读取（超时控制+损坏文件过滤）；\n    3. 边界条件处理（样本数不足时自动适配）；\n    4. 详细日志记录，方便问题定位；\n    5. 缓存有效性校验，支持数据集更新。\n    \"\"\"\n    def __init__(self, dataset, tumor_ratio=0.6, cache_validity_check=True,\n                 file_read_timeout=5, min_lesion_pixels=5):\n        self.dataset = dataset # LGGDataset实例（需带config属性）\n        self.tumor_ratio = tumor_ratio # batch中肿瘤样本占比（默认60%）\n        self.cache_validity_check = cache_validity_check # 缓存有效性校验\n        self.file_read_timeout = file_read_timeout # 文件读取超时时间（秒）\n        self.min_lesion_pixels = min_lesion_pixels # 判定肿瘤样本的最小像素数\n        self.save_dir = self.dataset.config.get('save_dir', './results') # 从数据集获取保存目录\n        os.makedirs(self.save_dir, exist_ok=True)\n        self.cache_path = os.path.join(self.save_dir, \"tumor_sample_cache.npy\")\n        # 1. 加载或生成肿瘤/非肿瘤样本索引（带有效性校验）\n        self.tumor_indices, self.non_tumor_indices = self._load_or_generate_indices()\n        # 2. 边界条件处理：确保样本数足够生成至少一个batch\n        self._validate_sample_count()\n        # 3. 打印初始化信息（便于调试）\n        print(f\"[平衡采样器] 初始化完成：\"\n              f\"肿瘤样本{len(self.tumor_indices)}个 | \"\n              f\"非肿瘤样本{len(self.non_tumor_indices)}个 | \"\n              f\"batch肿瘤占比{self.tumor_ratio:.2f} | \"\n              f\"缓存路径{self.cache_path}\")\n    def _load_or_generate_indices(self):\n        \"\"\"加载缓存索引（带有效性校验），若无效则重新生成\"\"\"\n        # 检查缓存是否存在且有效\n        if os.path.exists(self.cache_path) and self.cache_validity_check:\n            try:\n                cache = np.load(self.cache_path, allow_pickle=True).item()\n                # 校验缓存字段完整性\n                required_fields = [\"tumor\", \"non_tumor\", \"dataset_len\", \"timestamp\", \"valid_pairs_sample\"]\n                if not all(field in cache for field in required_fields):\n                    raise ValueError(\"缓存字段不完整\")\n                # 校验缓存与当前数据集匹配（长度+样本路径）\n                if (cache[\"dataset_len\"] != len(self.dataset) or\n                        not self._check_cache_paths_match(cache)):\n                    raise ValueError(\"缓存与当前数据集不匹配\")\n                print(f\"[平衡采样器] 加载有效缓存（生成时间：{cache['timestamp']}）\")\n                return cache[\"tumor\"], cache[\"non_tumor\"]\n            except Exception as e:\n                print(f\"[平衡采样器] 缓存无效或损坏：{str(e)[:100]}，将重新生成索引\")\n        # 重新生成索引（带超时控制和损坏文件过滤）\n        print(f\"[平衡采样器] 开始生成样本索引（共{len(self.dataset)}个样本）\")\n        tumor_indices = []\n        non_tumor_indices = []\n        invalid_samples = [] # 记录无效样本（便于用户排查）\n        for idx in tqdm(range(len(self.dataset)), desc=\"[平衡采样器] 标注样本类型\"):\n            try:\n                # 获取样本路径（从LGGDataset的valid_pairs中获取）\n                if not hasattr(self.dataset, 'valid_pairs') or idx >= len(self.dataset.valid_pairs):\n                    raise ValueError(f\"样本{idx}无有效路径（超出valid_pairs范围）\")\n                mask_path = self.dataset.valid_pairs[idx][1]\n                # 校验文件存在性和大小（排除空文件）\n                if not os.path.exists(mask_path):\n                    raise ValueError(f\"掩码文件不存在：{mask_path}\")\n                if os.path.getsize(mask_path) < 1024:\n                    raise ValueError(f\"掩码文件过小（<1KB）：{mask_path}\")\n                # 鲁棒的掩码读取（带超时控制）\n                mask = self._read_mask_with_timeout(mask_path)\n                # 判定肿瘤样本（病灶像素数≥最小阈值）\n                if (mask > 127).sum() >= self.min_lesion_pixels:\n                    tumor_indices.append(idx)\n                else:\n                    non_tumor_indices.append(idx)\n            except Exception as e:\n                invalid_samples.append((idx, str(e)[:80])) # 截取80字符避免日志过长\n                continue # 跳过无效样本，避免影响整体采样\n        # 保存新缓存（带时间戳和数据集信息，便于后续校验）\n        cache = {\n            \"tumor\": tumor_indices,\n            \"non_tumor\": non_tumor_indices,\n            \"dataset_len\": len(self.dataset),\n            \"timestamp\": time.strftime(\"%Y-%m-%d %H:%M:%S\", time.localtime()),\n            \"valid_pairs_sample\": self.dataset.valid_pairs[:5] # 保存前5个路径用于校验\n        }\n        np.save(self.cache_path, cache)\n        # 打印无效样本信息（便于用户排查数据问题）\n        if invalid_samples:\n            print(f\"[平衡采样器] 发现{len(invalid_samples)}个无效样本（已跳过）：\")\n            for idx, err in invalid_samples[:5]: # 只打印前5个，避免日志过长\n                print(f\" - 样本{idx}：{err}\")\n            if len(invalid_samples) > 5:\n                print(f\" - 剩余{len(invalid_samples) - 5}个无效样本详见日志\")\n        return tumor_indices, non_tumor_indices\n    def _read_mask_with_timeout(self, mask_path):\n        \"\"\"带超时控制的掩码读取（避免cv2.imread阻塞）\"\"\"\n        mask = None\n        error = None\n        def _read_mask():\n            nonlocal mask, error\n            try:\n                mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n                if mask is None:\n                    error = ValueError(f\"掩码读取失败（格式错误或损坏）\")\n            except Exception as e:\n                error = RuntimeError(f\"读取过程异常：{str(e)[:50]}\")\n        # 启动线程读取文件，超时则终止\n        thread = threading.Thread(target=_read_mask, daemon=True) # daemon=True确保主线程退出时子线程关闭\n        thread.start()\n        thread.join(timeout=self.file_read_timeout)\n        if thread.is_alive():\n            raise TimeoutError(f\"读取超时（超过{self.file_read_timeout}秒）\")\n        if error is not None:\n            raise error\n        return mask\n    def _check_cache_paths_match(self, cache):\n        \"\"\"校验缓存中的样本路径与当前数据集是否匹配（避免数据集更新后缓存失效）\"\"\"\n        # 取缓存中前5个样本路径与当前数据集对比\n        cache_sample_paths = [pair[1] for pair in cache[\"valid_pairs_sample\"]]\n        current_sample_paths = [self.dataset.valid_pairs[i][1] for i in range(min(5, len(self.dataset.valid_pairs)))]\n        return cache_sample_paths == current_sample_paths\n    def _validate_sample_count(self):\n        \"\"\"校验样本数是否足够，不足时自动调整tumor_ratio或报错\"\"\"\n        batch_size = self.dataset.config.get('batch_size', 4)\n        # 若肿瘤样本为0：报错（无法平衡采样）\n        if len(self.tumor_indices) == 0:\n            raise ValueError(\"数据集中无有效肿瘤样本！请检查：1. 掩码文件是否正常；2. min_lesion_pixels是否过小\")\n        # 若非肿瘤样本为0：强制tumor_ratio=1.0（全用肿瘤样本）\n        if len(self.non_tumor_indices) == 0:\n            print(f\"[平衡采样器] 警告：无有效非肿瘤样本，强制tumor_ratio=1.0\")\n            self.tumor_ratio = 1.0\n            return\n        # 计算最小需要的样本数（确保至少能生成1个batch）\n        min_tumor_needed = max(1, int(batch_size * self.tumor_ratio)) # 至少1个肿瘤样本\n        min_non_tumor_needed = batch_size - min_tumor_needed\n        # 若肿瘤样本不足：降低tumor_ratio\n        if len(self.tumor_indices) < min_tumor_needed:\n            new_tumor_ratio = len(self.tumor_indices) / batch_size\n            print(f\"[平衡采样器] 警告：肿瘤样本不足（需{min_tumor_needed}个，实际{len(self.tumor_indices)}个），\"\n                  f\"自动调整tumor_ratio从{self.tumor_ratio:.2f}到{new_tumor_ratio:.2f}\")\n            self.tumor_ratio = new_tumor_ratio\n        # 若非肿瘤样本不足：用肿瘤样本补全（避免batch过小）\n        if len(self.non_tumor_indices) < min_non_tumor_needed:\n            deficit = min_non_tumor_needed - len(self.non_tumor_indices)\n            print(f\"[平衡采样器] 警告：非肿瘤样本不足（需{min_non_tumor_needed}个，实际{len(self.non_tumor_indices)}个），\"\n                  f\"将用{deficit}个肿瘤样本补全batch\")\n    def __iter__(self):\n        \"\"\"生成采样索引（动态洗牌+无限迭代，避免DataLoader阻塞）\"\"\"\n        batch_size = self.dataset.config.get('batch_size', 4)\n        # 计算每个batch的肿瘤/非肿瘤样本数（确保整数且≥1）\n        tumor_num = max(1, int(batch_size * self.tumor_ratio))\n        non_tumor_num = batch_size - tumor_num\n        # 动态迭代：无限生成索引（DataLoader会根据total_batches自动停止）\n        while True:\n            # 1. 洗牌肿瘤/非肿瘤索引（每次洗牌确保多样性）\n            tumor_shuffled = np.random.permutation(self.tumor_indices)\n            non_tumor_shuffled = np.random.permutation(self.non_tumor_indices)\n            # 2. 循环取数（索引耗尽时重新洗牌）\n            tumor_ptr = 0 # 肿瘤样本指针\n            non_tumor_ptr = 0 # 非肿瘤样本指针\n            while True:\n                # 若肿瘤索引耗尽：重新洗牌\n                if tumor_ptr + tumor_num > len(tumor_shuffled):\n                    tumor_shuffled = np.random.permutation(self.tumor_indices)\n                    tumor_ptr = 0\n                # 若非肿瘤索引耗尽：用肿瘤样本补全，同时重置非肿瘤指针和洗牌\n                if non_tumor_ptr + non_tumor_num > len(non_tumor_shuffled):\n                    # 计算需要补全的肿瘤样本数\n                    needed = non_tumor_num - (len(non_tumor_shuffled) - non_tumor_ptr)\n                    # 取剩余非肿瘤样本 + 补全的肿瘤样本\n                    selected_non_tumor = list(non_tumor_shuffled[non_tumor_ptr:]) + \\\n                                         list(tumor_shuffled[tumor_ptr:tumor_ptr + needed])\n                    # 更新指针（肿瘤指针需跳过补全的样本）\n                    tumor_ptr += needed\n                    non_tumor_ptr = 0\n                    non_tumor_shuffled = np.random.permutation(self.non_tumor_indices)\n                else:\n                    # 正常取非肿瘤样本\n                    selected_non_tumor = non_tumor_shuffled[non_tumor_ptr:non_tumor_ptr + non_tumor_num]\n                    non_tumor_ptr += non_tumor_num\n                # 3. 取当前batch的肿瘤样本\n                selected_tumor = tumor_shuffled[tumor_ptr:tumor_ptr + tumor_num]\n                tumor_ptr += tumor_num\n                # 4. 合并并打乱batch内顺序（避免样本类型集中）\n                batch_indices = np.concatenate([selected_tumor, selected_non_tumor])\n                np.random.shuffle(batch_indices)\n                # 5. 返回当前batch的索引（转为int避免类型错误）\n                yield from batch_indices.astype(int)\n    def __len__(self):\n        \"\"\"采样器长度：与数据集长度一致（确保DataLoader迭代次数正确）\"\"\"\n        return len(self.dataset)\n    def get_sample_distribution(self):\n        \"\"\"获取当前样本分布（用于调试）\"\"\"\n        return {\n            \"total_samples\": len(self.dataset),\n            \"tumor_samples\": len(self.tumor_indices),\n            \"non_tumor_samples\": len(self.non_tumor_indices),\n            \"tumor_ratio\": self.tumor_ratio,\n            \"batch_tumor_num\": max(1, int(self.dataset.config.get('batch_size', 4) * self.tumor_ratio))\n        }\nclass SmallLesionAugmentation:\n    \"\"\"小病灶增强：针对小肿瘤样本的专用数据增强（提高小肿瘤检测能力）\n    小肿瘤样本容易被忽略，通过增强提升模型对其的敏感度\n    \"\"\"\n    def __init__(self, p=0.5):\n        self.p = p # 增强概率（50%）\n        # 非肿瘤区域干扰：模拟噪声或空洞（迫使模型关注小病灶）\n        self.disturb = A.OneOf([ # 随机选择一种干扰方式\n            A.CoarseDropout( # 随机挖空洞\n                max_holes=5, # 最多5个洞\n                max_height=15, max_width=15, # 洞的最大尺寸\n                min_holes=1, min_height=5, min_width=5, # 洞的最小尺寸\n                fill_value=0, # 洞填充为0（黑色）\n                p=0.5\n            ),\n            A.GaussNoise(var_limit=(10.0, 50.0), p=0.5) # 高斯噪声\n        ], p=0.7)\n        self.zoom = A.RandomScale(scale_limit=(-0.3, 0.3), p=0.5) # 随机缩放（±30%）\n        self.rotate = A.RandomRotate90(p=0.5) # 随机旋转90度倍数\n    def __call__(self, image, mask=None):\n        \"\"\"对小病灶样本应用增强：缩放+旋转+非肿瘤区域干扰\"\"\"\n        if np.random.random() < self.p and mask is not None:\n            # 缩放和旋转（同时作用于图像和掩码）\n            augmented = self.zoom(image=image, mask=mask)\n            image, mask = augmented['image'], augmented['mask']\n            augmented = self.rotate(image=image, mask=mask)\n            image, mask = augmented['image'], augmented['mask']\n            # 若病灶较小（像素数<100），在非肿瘤区域添加干扰\n            if mask.sum() < 100:\n                non_tumor_mask = (mask == 0).astype(np.uint8) * 255 # 非肿瘤区域掩码\n                augmented = self.disturb(image=image, mask=non_tumor_mask) # 仅在非肿瘤区域添加干扰\n                image = augmented['image']\n        return image, mask\ndef get_transforms(img_size=256):\n    \"\"\"获取数据转换策略（训练集增强，验证集仅标准化）\n    训练集增强提高模型鲁棒性，验证集保持稳定以准确评估性能\n    \"\"\"\n    # 训练集转换：包含多种数据增强（提高模型鲁棒性）\n    train_transform = A.Compose([\n        A.Resize(height=img_size, width=img_size), # resize到指定尺寸\n        A.HorizontalFlip(p=0.5), # 水平翻转（50%概率）\n        A.VerticalFlip(p=0.3), # 垂直翻转（30%概率）\n        # 使用Affine：平移、缩放、旋转\n        A.Affine(\n            translate_percent={\"x\": (-0.05, 0.05), \"y\": (-0.05, 0.05)}, # 平移±5%\n            scale=(0.85, 1.15), # 缩放85%-115%\n            rotate=(-20, 20), # 旋转±20度\n            p=0.5\n        ),\n        # 弹性形变（模拟组织变形）\n        A.ElasticTransform(\n            alpha=120, # 形变强度\n            sigma=120 * 0.05, # 平滑系数\n            p=0.3 # 30%概率\n        ),\n        # 颜色增强（提高对光照变化的鲁棒性）\n        A.OneOf([\n            A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.5), # 亮度/对比度\n            A.CLAHE(clip_limit=2.0, p=0.5), # 对比度受限自适应直方图均衡\n            A.Equalize(p=0.5) # 直方图均衡\n        ], p=0.5),\n        # 模糊（提高抗噪能力）\n        A.OneOf([\n            A.GaussianBlur(blur_limit=3, p=0.5), # 高斯模糊\n            A.MedianBlur(blur_limit=3, p=0.5) # 中值模糊\n        ], p=0.3),\n        A.Normalize( # 标准化（使用ImageNet均值和标准差）\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225]\n        ),\n        ToTensorV2() # 转换为PyTorch张量（HWC->CHW）\n    ])\n    # 验证集转换：仅resize和标准化（无增强，确保评估稳定）\n    val_transform = A.Compose([\n        A.Resize(height=img_size, width=img_size),\n        A.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225]\n        ),\n        ToTensorV2()\n    ])\n    return train_transform, val_transform\nclass LGGDataset(Dataset):\n    \"\"\"LGG脑肿瘤数据集类（带config属性，支持平衡采样器获取配置）\n    处理数据读取、格式转换和基本验证，确保输入模型的数据合法\n    \"\"\"\n    def __init__(self, image_paths, mask_paths, transform=None, min_lesion_size=5, config=None):\n        self.image_paths = image_paths # 图像路径列表\n        self.mask_paths = mask_paths # 掩码路径列表\n        self.transform = transform # 数据转换函数\n        self.min_lesion_size = min_lesion_size # 最小病灶尺寸（过滤无效样本）\n        self.config = config or {} # 新增：保存配置（如batch_size、save_dir）\n        # 筛选有效数据对（图像和掩码均存在且大小合理）\n        self.valid_pairs = []\n        for img_path, mask_path in zip(image_paths, mask_paths):\n            if (os.path.exists(img_path) and os.path.exists(mask_path) and\n                    os.path.getsize(img_path) > 1024 and os.path.getsize(mask_path) > 1024):\n                self.valid_pairs.append((img_path, mask_path))\n        print(f\"[数据集] 有效数据对: {len(self.valid_pairs)}/{len(image_paths)}\")\n    def __len__(self):\n        \"\"\"返回数据集大小\"\"\"\n        return len(self.valid_pairs)\n    def __getitem__(self, idx):\n        \"\"\"获取单个样本：图像和对应的掩码\"\"\"\n        img_path, mask_path = self.valid_pairs[idx]\n        try:\n            # 读取图像（BGR格式）并转换为RGB\n            image = cv2.imread(img_path)\n            if image is None:\n                raise ValueError(f\"图像读取失败: {img_path}\")\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n            # 读取掩码（灰度图）\n            mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n            if mask is None:\n                raise ValueError(f\"掩码读取失败: {mask_path}\")\n            # 应用数据转换\n            if self.transform:\n                augmented = self.transform(image=image, mask=mask)\n                image = augmented['image']\n                mask = augmented['mask']\n            else:\n                # 无转换时的默认处理\n                image = A.Resize(256, 256)(image=image)['image']\n                mask = A.Resize(256, 256)(image=mask)['image']\n                image = torch.from_numpy(image).permute(2, 0, 1).float() / 255.0 # 归一化到0-1\n                mask = torch.from_numpy(mask).unsqueeze(0).float()\n            # 掩码二值化（>0.5视为肿瘤）\n            mask = (mask > 0.5).float()\n            return image, mask\n        except Exception as e:\n            # 处理异常：返回随机dummy数据避免程序崩溃\n            print(f\"[数据集] 样本{idx}处理警告: {str(e)[:80]}, 使用备用数据\")\n            dummy_image = torch.randn(3, 256, 256)\n            dummy_mask = torch.zeros(1, 256, 256)\n            return dummy_image, dummy_mask\ndef load_lgg_data(data_root):\n    \"\"\"加载LGG脑肿瘤数据集：收集图像和掩码路径，并分割为训练/验证/测试集\n    自动查找数据路径，兼容多种环境（本地/Kaggle）\n    \"\"\"\n    image_paths = [] # 图像路径列表\n    mask_paths = [] # 掩码路径列表\n    # 检查数据根目录是否存在，不存在则尝试备选路径\n    if not os.path.exists(data_root):\n        alternative_paths = [\n            \"/kaggle/input/lgg-mri-segmentation/kaggle_3m\", # Kaggle默认路径\n            \"./data/lgg-mri-segmentation\", # 本地数据路径1\n            \"./kaggle_3m\" # 本地数据路径2\n        ]\n        for path in alternative_paths:\n            if os.path.exists(path):\n                data_root = path\n                print(f\"[数据加载] 使用备选路径: {data_root}\")\n                break\n        else:\n            raise FileNotFoundError(f\"数据目录不存在: {data_root}\")\n    # 遍历目录收集图像和掩码路径（图像为.tif，掩码为_mask.tif）\n    for root, dirs, files in os.walk(data_root):\n        for file in files:\n            # 筛选图像文件（排除掩码文件）\n            if file.endswith('.tif') and not file.endswith('_mask.tif'):\n                img_path = os.path.join(root, file)\n                mask_path = img_path.replace('.tif', '_mask.tif') # 掩码路径（与图像对应）\n                if os.path.exists(mask_path):\n                    image_paths.append(img_path)\n                    mask_paths.append(mask_path)\n    if not image_paths:\n        raise ValueError(\"未找到有效的图像-掩码对（检查文件格式是否为.tif）\")\n    # 分割数据集：先分为训练集（70%）和临时集（30%）\n    train_imgs, temp_imgs, train_masks, temp_masks = train_test_split(\n        image_paths, mask_paths, test_size=0.3, random_state=42, shuffle=True\n    )\n    # 临时集再分为验证集（20%总数据）和测试集（10%总数据）\n    val_imgs, test_imgs, val_masks, test_masks = train_test_split(\n        temp_imgs, temp_masks, test_size=1 / 3, random_state=42, shuffle=True\n    )\n    print(f\"[数据加载] 数据集分割完成: 训练集 {len(train_imgs)} | 验证集 {len(val_imgs)} | 测试集 {len(test_imgs)}\")\n    return train_imgs, val_imgs, test_imgs, train_masks, val_masks, test_masks\n# ==============================\n# 4. 损失函数与评估指标\n# 作用：定义适合肿瘤分割的损失函数，实现多种评估指标计算\n# ==============================\nclass FocalDiceLoss(nn.Module):\n    \"\"\"融合Focal Loss和Dice Loss：解决类别不平衡，提高边界精度\n    Focal Loss聚焦难分样本，Dice Loss优化分割重叠度，结合两者优势\n    \"\"\"\n    def __init__(self, alpha=0.85, gamma=2.0, smooth=1e-6, bg_penalty=0.3, edge_weight=1.2):\n        super().__init__()\n        self.alpha = alpha # Focal Loss的权重参数（平衡正负样本）\n        self.gamma = gamma # Focal Loss的聚焦参数（降低易分类样本权重）\n        self.smooth = smooth # 平滑项（避免除零）\n        self.bg_penalty = bg_penalty # 背景误判惩罚（减少假阳性）\n        self.edge_weight = edge_weight # 肿瘤边缘权重（提高边界精度）\n    def forward(self, pred, target):\n        \"\"\"\n        前向传播：\n        1. 计算肿瘤边缘掩码（用于加权）\n        2. 计算Focal Loss（处理类别不平衡）\n        3. 计算Dice Loss（提高分割重叠度）\n        4. 融合两种损失（各占50%）\n        \"\"\"\n        # 去除通道维度（假设输入为[B,1,H,W]，转换为[B,H,W]）\n        if pred.dim() == 4 and pred.size(1) == 1:\n            pred = pred.squeeze(1)\n        if target.dim() == 4 and target.size(1) == 1:\n            target = target.squeeze(1)\n        def get_tumor_edge(mask):\n            \"\"\"计算肿瘤边缘掩码（使用Sobel算子检测边缘）\"\"\"\n            # Sobel算子（水平和垂直方向）\n            sobel_x = torch.tensor([[-1, 0, 1], [-2, 0, 2], [-1, 0, 1]], device=mask.device).float().unsqueeze(\n                0).unsqueeze(0)\n            sobel_y = torch.tensor([[-1, -2, -1], [0, 0, 0], [1, 2, 1]], device=mask.device).float().unsqueeze(\n                0).unsqueeze(0)\n            # 卷积计算边缘\n            edge_x = F.conv2d(mask.unsqueeze(1), sobel_x, padding=1).abs()\n            edge_y = F.conv2d(mask.unsqueeze(1), sobel_y, padding=1).abs()\n            edge = (edge_x + edge_y) > 0.1 # 阈值筛选边缘\n            return edge.squeeze(1).float() # 边缘掩码（1表示边缘）\n        # 计算肿瘤边缘掩码\n        tumor_edge = get_tumor_edge(target)\n        pred_prob = torch.sigmoid(pred) # 预测概率（0-1）\n        # 1. Focal Loss计算\n        # pt: 对于正样本=pred_prob，负样本=1-pred_prob\n        pt = torch.where(target == 1, pred_prob, 1 - pred_prob)\n        # 基础Focal Loss\n        focal_loss = self.alpha * (1 - pt) ** self.gamma * F.binary_cross_entropy_with_logits(pred, target,\n                                                                                              reduction='none')\n        # 背景误判惩罚：对背景区域的假阳性（预测为肿瘤）加重惩罚\n        bg_false_pos = (target == 0) & (pred_prob > 0.5)\n        bg_edge_false_pos = bg_false_pos & (tumor_edge > 0) # 边缘附近的背景假阳性\n        focal_loss[bg_false_pos] *= (1 + self.bg_penalty) # 普通背景假阳性惩罚\n        focal_loss[bg_edge_false_pos] *= (1 + self.bg_penalty * 0.5) # 边缘背景假阳性惩罚（稍轻）\n        # 肿瘤边缘加权：提高边缘区域的损失权重\n        focal_loss[tumor_edge > 0] *= self.edge_weight\n        focal_loss = focal_loss.mean() # 平均Focal Loss\n        # 2. Dice Loss计算（带边缘加权）\n        pred_flat = pred_prob.contiguous().view(-1) # 展平预测\n        target_flat = target.contiguous().view(-1) # 展平目标\n        edge_flat = tumor_edge.contiguous().view(-1) # 展平边缘掩码\n        # 对边缘区域的预测和目标加权\n        weighted_pred = pred_flat * (1 + (self.edge_weight - 1) * edge_flat)\n        weighted_target = target_flat * (1 + (self.edge_weight - 1) * edge_flat)\n        # Dice公式：2*交集/(预测和+目标和)\n        intersection = (weighted_pred * weighted_target).sum() + self.smooth * 2\n        union = weighted_pred.sum() + weighted_target.sum() + self.smooth\n        dice_loss = 1 - (2. * intersection) / (union) # Dice Loss（1-Dice）\n        # 融合两种损失（各占50%）\n        return 0.5 * focal_loss + 0.5 * dice_loss\ndef calculate_dice_score(pred, target, smooth=1e-6):\n    \"\"\"计算Dice系数（分割任务核心指标，范围0-1，越高越好）\n    衡量预测与目标的重叠程度，对类别不平衡敏感\n    \"\"\"\n    pred_prob = torch.sigmoid(pred) # 预测概率\n    pred_bin = (pred_prob > 0.5).float() # 二值化（0.5为阈值）\n    # 去除通道维度\n    if target.dim() == 4:\n        target = target.squeeze(1)\n    if pred_bin.dim() == 4:\n        pred_bin = pred_bin.squeeze(1)\n    # 计算交集和并集\n    intersection = (pred_bin * target).sum()\n    union = pred_bin.sum() + target.sum()\n    # Dice公式\n    dice = (2. * intersection + smooth) / (union + smooth)\n    return dice.item() # 返回标量值\ndef calculate_iou(pred, target, smooth=1e-6):\n    \"\"\"计算IoU（交并比，范围0-1，越高越好）\n    衡量预测与目标的重叠比例，对边界精度敏感\n    \"\"\"\n    pred_prob = torch.sigmoid(pred)\n    pred_bin = (pred_prob > 0.5).float()\n    # 去除通道维度\n    if target.dim() == 4:\n        target = target.squeeze(1)\n    if pred_bin.dim() == 4:\n        pred_bin = pred_bin.squeeze(1)\n    # 计算交集和并集\n    intersection = (pred_bin * target).sum()\n    union = pred_bin.sum() + target.sum() - intersection # 并集=预测+目标-交集\n    # IoU公式\n    iou = (intersection + smooth) / (union + smooth)\n    return iou.item()\ndef calculate_metrics(pred, target, smooth=1e-6):\n    \"\"\"计算多种评估指标：Dice、IoU、灵敏度、特异度、精确率、F1\n    全面评估模型性能，避免单一指标的局限性\n    \"\"\"\n    pred_prob = torch.sigmoid(pred)\n    pred_bin = (pred_prob > 0.5).float() # 二值化预测\n    # 去除通道维度并展平\n    if target.dim() == 4:\n        target = target.squeeze(1)\n    if pred_bin.dim() == 4:\n        pred_bin = pred_bin.squeeze(1)\n    pred_flat = pred_bin.view(-1).cpu().numpy() # 转换为NumPy数组（CPU）\n    target_flat = target.view(-1).cpu().numpy()\n    # 计算混淆矩阵（TN, FP, FN, TP）\n    tn, fp, fn, tp = confusion_matrix(target_flat, pred_flat, labels=[0, 1]).ravel()\n    # 灵敏度（召回率）：TP/(TP+FN) （正确预测的肿瘤占实际肿瘤的比例）\n    sensitivity = tp / (tp + fn + smooth)\n    # 特异度：TN/(TN+FP) （正确预测的背景占实际背景的比例）\n    specificity = tn / (tn + fp + smooth)\n    # 精确率：TP/(TP+FP) （预测为肿瘤的样本中实际是肿瘤的比例）\n    precision = tp / (tp + fp + smooth)\n    # F1分数：2*精确率*灵敏度/(精确率+灵敏度)\n    f1 = 2 * precision * sensitivity / (precision + sensitivity + smooth)\n    return {\n        \"dice\": calculate_dice_score(pred, target),\n        \"iou\": calculate_iou(pred, target),\n        \"sensitivity\": sensitivity,\n        \"specificity\": specificity,\n        \"precision\": precision,\n        \"f1\": f1,\n        \"tn\": tn,  # 新增：返回混淆矩阵值，便于累积\n        \"fp\": fp,\n        \"fn\": fn,\n        \"tp\": tp\n    }\n# ==============================\n# 5. 训练与评估函数\n# 作用：实现模型训练和验证的核心逻辑\n# ==============================\ndef train_epoch(model, train_loader, criterion, optimizer, device, scaler, total_batches, writer, epoch):\n    \"\"\"训练一个epoch：迭代训练集，更新模型参数\n    实现前向传播、损失计算、反向传播和参数更新的完整流程\n    \"\"\"\n    model.train() # 模型设为训练模式（启用 dropout/batch norm更新）\n    total_loss = 0.0 # 总损失\n    total_dice = 0.0 # 总Dice系数\n    # 进度条（tqdm）：显示训练进度\n    pbar = tqdm(enumerate(train_loader), desc=\"Training\", total=total_batches)\n    for batch_idx, (images, masks) in pbar:\n        # 限制batch数量（避免采样器生成过多样本）\n        if batch_idx >= total_batches:\n            break\n        # 数据移至计算设备\n        images = images.to(device)\n        masks = masks.to(device)\n        optimizer.zero_grad() # 清空梯度\n        # 记录当前学习率（每个epoch记录一次）\n        current_lr = optimizer.param_groups[0]['lr']\n        if batch_idx == 0:\n            writer.add_scalar('Learning Rate', current_lr, epoch)\n        # 仅在CUDA环境下使用混合精度训练（加速训练，减少内存占用）\n        if use_cuda:\n            try:\n                # 自动混合精度（AMP）：前向传播使用FP16，反向传播使用FP32\n                with torch.amp.autocast(device_type='cuda'):\n                    outputs = model(images) # 模型输出\n                    loss = criterion(outputs, masks) # 计算损失\n            except:\n                # 兼容旧版本PyTorch\n                with autocast():\n                    outputs = model(images)\n                    loss = criterion(outputs, masks)\n            # 混合精度反向传播\n            scaler.scale(loss).backward() # 梯度缩放（避免FP16下溢）\n            scaler.step(optimizer) # 优化器更新（根据缩放后的梯度）\n            scaler.update() # 更新缩放器状态\n        else:\n            # CPU环境：禁用混合精度\n            outputs = model(images)\n            loss = criterion(outputs, masks)\n            loss.backward() # 反向传播计算梯度\n            optimizer.step() # 优化器更新参数\n        # 计算当前batch的Dice系数\n        dice = calculate_dice_score(outputs, masks)\n        total_loss += loss.item()\n        total_dice += dice\n        # 记录每个batch的损失和Dice到TensorBoard\n        global_step = epoch * total_batches + batch_idx\n        writer.add_scalar('Train/Batch Loss', loss.item(), global_step)\n        writer.add_scalar('Train/Batch Dice', dice, global_step)\n        # 更新进度条信息\n        pbar.set_postfix({\n            \"Loss\": f\"{loss.item():.4f}\",\n            \"Dice\": f\"{dice:.4f}\",\n            \"LR\": f\"{current_lr:.6f}\",\n            \"Batch\": f\"{batch_idx + 1}/{total_batches}\"\n        })\n    # 计算当前epoch的平均损失和Dice\n    avg_loss = total_loss / total_batches\n    avg_dice = total_dice / total_batches\n    # 记录epoch级指标到TensorBoard\n    writer.add_scalar('Train/Epoch Loss', avg_loss, epoch)\n    writer.add_scalar('Train/Epoch Dice', avg_dice, epoch)\n    return avg_loss, avg_dice\ndef plot_confusion_matrix(conf_matrix, save_path, writer, epoch, phase='Validation'):\n    \"\"\"绘制混淆矩阵热图并保存\"\"\"\n    labels = ['Background', 'Tumor']\n    plt.figure(figsize=(6, 5))\n    sns.heatmap(conf_matrix, annot=True, fmt='g', cmap='Blues', xticklabels=labels, yticklabels=labels)\n    plt.xlabel('Predicted')\n    plt.ylabel('True')\n    plt.title(f'{phase} Confusion Matrix')\n    plt.savefig(save_path, dpi=300, bbox_inches='tight')\n    writer.add_figure(f'{phase}/Confusion Matrix', plt.gcf(), epoch)\n    plt.close()\n\ndef validate_epoch(model, val_loader, criterion, device, writer, epoch, phase='Validation'):\n    \"\"\"验证一个epoch：在验证集/测试集上评估模型性能（不更新参数）\n    禁用梯度计算，提高效率；计算多种指标，全面评估模型\n    \"\"\"\n    model.eval() # 模型设为评估模式（禁用 dropout/batch norm固定）\n    total_loss = 0.0 # 总损失\n    total_dice = 0.0 # 总Dice系数\n    all_metrics = [] # 存储所有batch的指标\n    total_batches = len(val_loader) # 总batch数\n    total_tn, total_fp, total_fn, total_tp = 0, 0, 0, 0  # 累积混淆矩阵值\n    # 进度条\n    pbar = tqdm(enumerate(val_loader), desc=phase, total=total_batches)\n    for batch_idx, (images, masks) in pbar:\n        # 数据移至计算设备\n        images = images.to(device)\n        masks = masks.to(device)\n        # 禁用梯度计算（节省内存，加速计算）\n        with torch.no_grad():\n            outputs = model(images) # 模型输出\n            loss = criterion(outputs, masks) # 计算损失\n        # 计算当前batch的指标\n        dice = calculate_dice_score(outputs, masks)\n        metrics = calculate_metrics(outputs, masks)\n        all_metrics.append(metrics)\n        total_loss += loss.item()\n        total_dice += dice\n        # 累积混淆矩阵值\n        total_tn += metrics['tn']\n        total_fp += metrics['fp']\n        total_fn += metrics['fn']\n        total_tp += metrics['tp']\n        # 更新进度条\n        pbar.set_postfix({\n            \"Loss\": f\"{loss.item():.4f}\",\n            \"Dice\": f\"{dice:.4f}\",\n            \"Batch\": f\"{batch_idx + 1}/{total_batches}\"\n        })\n    # 计算平均指标\n    avg_metrics = {}\n    for key in all_metrics[0].keys():\n        if key in ['tn', 'fp', 'fn', 'tp']:  # 跳过累积值\n            continue\n        avg_metrics[key] = np.mean([m[key] for m in all_metrics]) # 每个指标的平均值\n    avg_loss = total_loss / total_batches\n    avg_dice = total_dice / total_batches\n    # 构建整体混淆矩阵\n    conf_matrix = np.array([[total_tn, total_fp], [total_fn, total_tp]])\n    # 记录指标到TensorBoard\n    writer.add_scalar(f'{phase}/Loss', avg_loss, epoch)\n    writer.add_scalar(f'{phase}/Dice', avg_dice, epoch)\n    writer.add_scalar(f'{phase}/IoU', avg_metrics['iou'], epoch)\n    writer.add_scalar(f'{phase}/Sensitivity', avg_metrics['sensitivity'], epoch)\n    writer.add_scalar(f'{phase}/Specificity', avg_metrics['specificity'], epoch)\n    writer.add_scalar(f'{phase}/Precision', avg_metrics['precision'], epoch)\n    writer.add_scalar(f'{phase}/F1', avg_metrics['f1'], epoch)\n    # 绘制并保存混淆矩阵\n    conf_save_path = os.path.join(config['save_dir'], f'{phase.lower()}_confusion_matrix_epoch_{epoch + 1}.png')\n    plot_confusion_matrix(conf_matrix, conf_save_path, writer, epoch, phase)\n    print(f\"[可视化] {phase}混淆矩阵已保存：{conf_save_path}\")\n    return avg_loss, avg_dice, avg_metrics\n# ==============================\n# 6. 可视化函数（含注意力热图和特征图）\n# 作用：生成训练曲线、指标变化和预测结果可视化\n# ==============================\ndef plot_training_curves(train_losses, val_losses, train_dices, val_dices, save_path):\n    \"\"\"绘制训练/验证损失和Dice曲线\n    直观展示模型训练过程中的收敛情况和过拟合风险\n    \"\"\"\n    plt.figure(figsize=(12, 5))\n    # 左侧：损失曲线\n    plt.subplot(1, 2, 1)\n    plt.plot(train_losses, label='Training Loss', linewidth=2)\n    plt.plot(val_losses, label='Validation Loss', linewidth=2)\n    plt.xlabel('Epoch', fontsize=10)\n    plt.ylabel('Loss', fontsize=10)\n    plt.legend(fontsize=9)\n    plt.title('Training and Validation Loss', fontsize=11)\n    plt.grid(True, alpha=0.3) # 网格线（透明度0.3）\n    # 右侧：Dice曲线\n    plt.subplot(1, 2, 2)\n    plt.plot(train_dices, label='Training Dice', linewidth=2)\n    plt.plot(val_dices, label='Validation Dice', linewidth=2)\n    plt.xlabel('Epoch', fontsize=10)\n    plt.ylabel('Dice Coefficient', fontsize=10)\n    plt.legend(fontsize=9)\n    plt.title('Training and Validation Dice', fontsize=11)\n    plt.grid(True, alpha=0.3)\n    plt.tight_layout() # 自动调整布局\n    plt.savefig(save_path, dpi=300, bbox_inches='tight') # 保存图像（300dpi）\n    plt.close() # 关闭图像（释放内存）\ndef plot_additional_metrics(val_ious, val_sensitivities, val_specificities, save_path):\n    \"\"\"绘制验证集的IoU、灵敏度、特异度曲线\n    补充评估模型在不同维度的性能表现\n    \"\"\"\n    plt.figure(figsize=(15, 5))\n    epochs = range(1, len(val_ious) + 1) # epoch范围\n    # 左侧：IoU曲线\n    plt.subplot(1, 3, 1)\n    plt.plot(epochs, val_ious, label='Validation IoU', color='#2E86AB', linewidth=2)\n    plt.xlabel('Epoch', fontsize=10)\n    plt.ylabel('IoU', fontsize=10)\n    plt.title('Validation IoU Over Epochs', fontsize=11)\n    plt.legend(fontsize=9)\n    plt.grid(True, alpha=0.3)\n    # 中间：灵敏度曲线\n    plt.subplot(1, 3, 2)\n    plt.plot(epochs, val_sensitivities, label='Validation Sensitivity', color='#A23B72', linewidth=2)\n    plt.xlabel('Epoch', fontsize=10)\n    plt.ylabel('Sensitivity', fontsize=10)\n    plt.title('Validation Sensitivity Over Epochs', fontsize=11)\n    plt.legend(fontsize=9)\n    plt.grid(True, alpha=0.3)\n    # 右侧：特异度曲线\n    plt.subplot(1, 3, 3)\n    plt.plot(epochs, val_specificities, label='Validation Specificity', color='#F18F01', linewidth=2)\n    plt.xlabel('Epoch', fontsize=10)\n    plt.ylabel('Specificity', fontsize=10)\n    plt.title('Validation Specificity Over Epochs', fontsize=11)\n    plt.legend(fontsize=9)\n    plt.grid(True, alpha=0.3)\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=300, bbox_inches='tight')\n    plt.close()\ndef generate_local_dice_heatmap(pred, target, grid_size=16, smooth=1e-6):\n    \"\"\"生成局部Dice热力图：将图像分网格，计算每个网格的Dice系数\n    定位模型分割效果好/差的区域，辅助分析模型弱点\n    \"\"\"\n    pred_prob = torch.sigmoid(pred)\n    pred_bin = (pred_prob > 0.5).float() # 二值化预测\n    # 去除通道维度\n    if target.dim() == 4 and target.size(1) == 1:\n        target = target.squeeze(1)\n    if pred_bin.dim() == 4 and pred_bin.size(1) == 1:\n        pred_bin = pred_bin.squeeze(1)\n    batch_size, H, W = pred_bin.shape # batch大小、图像高、宽\n    grid_h = H // grid_size # 每个网格的高度\n    grid_w = W // grid_size # 每个网格的宽度\n    # 初始化局部Dice矩阵（batch_size x grid_size x grid_size）\n    local_dice = torch.zeros((batch_size, grid_size, grid_size), device=pred_bin.device)\n    for b in range(batch_size): # 遍历batch\n        for i in range(grid_size): # 遍历行网格\n            for j in range(grid_size): # 遍历列网格\n                # 计算当前网格的坐标范围\n                h_start = i * grid_h\n                h_end = min((i + 1) * grid_h, H)\n                w_start = j * grid_w\n                w_end = min((j + 1) * grid_w, W)\n                # 提取当前网格的预测和目标\n                pred_grid = pred_bin[b, h_start:h_end, w_start:w_end]\n                target_grid = target[b, h_start:h_end, w_start:w_end]\n                # 计算网格内的Dice\n                intersection = (pred_grid * target_grid).sum()\n                union = pred_grid.sum() + target_grid.sum()\n                local_dice[b, i, j] = (2 * intersection + smooth) / (union + smooth)\n    return local_dice\ndef visualize_predictions(model, val_loader, device, save_path, writer, epoch, num_samples=4, grid_size=16,\n                          phase='Validation'):\n    \"\"\"可视化预测结果：输入图像、真实掩码、预测掩码、叠加图、局部Dice热力图、注意力热图、中间特征图\n    直观展示模型的分割效果，辅助分析错误模式\n    \"\"\"\n    model.eval() # 评估模式\n  \n    # 提取注意力热图和中间特征图的辅助函数\n    def get_attention_heatmap(model, img_size):\n        attn_weight = model.att4.attn_weight # 最深层注意力门权重\n        attn_heatmap = F.interpolate(\n            attn_weight, size=(img_size, img_size),\n            mode='bilinear', align_corners=False\n        )\n        attn_heatmap = attn_heatmap.mean(dim=1, keepdim=True) # 通道平均\n        attn_heatmap = (attn_heatmap - attn_heatmap.min()) / (attn_heatmap.max() - attn_heatmap.min() + 1e-6)\n        return attn_heatmap.squeeze(1).cpu().numpy()\n    def get_mid_feature_map(model, img_size):\n        feat = model.mid_features['bottleneck'] # 多尺度瓶颈层特征\n        feat_map = F.interpolate(\n            feat, size=(img_size, img_size),\n            mode='bilinear', align_corners=False\n        )\n        feat_map = feat_map.mean(dim=1, keepdim=True) # 通道平均\n        feat_map = (feat_map - feat_map.min()) / (feat_map.max() - feat_map.min() + 1e-6)\n        return feat_map.squeeze(1).cpu().numpy()\n    # 创建7列子图\n    fig, axes = plt.subplots(num_samples, 7, figsize=(28, 4 * num_samples))\n    if num_samples == 1:\n        axes = axes.reshape(1, -1)\n    with torch.no_grad():\n        for images, masks in val_loader:\n            images = images.to(device)\n            masks = masks.to(device)\n            outputs = model(images)\n            outputs_prob = torch.sigmoid(outputs)\n            local_dice = generate_local_dice_heatmap(outputs, masks, grid_size=grid_size)\n            attn_heatmaps = get_attention_heatmap(model, config['img_size']) # 注意力热图\n            feat_maps = get_mid_feature_map(model, config['img_size']) # 中间特征图\n            for i in range(min(num_samples, images.size(0))):\n                # 1. 输入图像\n                image = images[i].cpu().permute(1, 2, 0).numpy()\n                mean = np.array([0.485, 0.456, 0.406])\n                std = np.array([0.229, 0.224, 0.225])\n                image = image * std + mean\n                image = np.clip(image, 0, 1)\n                axes[i, 0].imshow(image)\n                axes[i, 0].set_title('Input Image', fontsize=10)\n                axes[i, 0].axis('off')\n                # 2. 真实掩码\n                mask = masks[i].squeeze().cpu().numpy()\n                axes[i, 1].imshow(mask, cmap='gray', vmin=0, vmax=1)\n                axes[i, 1].set_title('Ground Truth', fontsize=10)\n                axes[i, 1].axis('off')\n                # 3. 预测掩码\n                pred = (outputs_prob[i].squeeze().cpu().numpy() > 0.5).astype(np.float32)\n                axes[i, 2].imshow(pred, cmap='gray', vmin=0, vmax=1)\n                axes[i, 2].set_title('Prediction', fontsize=10)\n                axes[i, 2].axis('off')\n                # 4. 叠加图\n                overlay = image.copy()\n                overlay[mask == 1] = np.array([1, 0, 0]) * 0.3 + overlay[mask == 1] * 0.7\n                overlay[pred == 1] = np.array([0, 1, 0]) * 0.3 + overlay[pred == 1] * 0.7\n                axes[i, 3].imshow(overlay)\n                axes[i, 3].set_title('Overlay (Red=GT, Green=Pred)', fontsize=10)\n                axes[i, 3].axis('off')\n                # 5. 局部Dice热力图\n                dice_grid = local_dice[i].cpu().numpy()\n                dice_heatmap = cv2.resize(dice_grid, (image.shape[0], image.shape[1]), interpolation=cv2.INTER_LINEAR)\n                im5 = axes[i, 4].imshow(dice_heatmap, cmap='jet', vmin=0, vmax=1)\n                axes[i, 4].set_title('Local Dice Heatmap', fontsize=10)\n                axes[i, 4].axis('off')\n                if i == num_samples - 1:\n                    cbar5 = plt.colorbar(im5, ax=axes[i, 4], fraction=0.046, pad=0.04)\n                    cbar5.set_label('Dice Coefficient', fontsize=8)\n                # 6. 注意力热图\n                attn_heatmap = attn_heatmaps[i]\n                im6 = axes[i, 5].imshow(attn_heatmap, cmap='viridis', vmin=0, vmax=1)\n                axes[i, 5].set_title('Attention Heatmap (Deepest Gate)', fontsize=10)\n                axes[i, 5].axis('off')\n                if i == num_samples - 1:\n                    cbar6 = plt.colorbar(im6, ax=axes[i, 5], fraction=0.046, pad=0.04)\n                    cbar6.set_label('Attention Weight', fontsize=8)\n                # 7. 中间特征图\n                feat_map = feat_maps[i]\n                im7 = axes[i, 6].imshow(feat_map, cmap='plasma', vmin=0, vmax=1)\n                axes[i, 6].set_title('Mid Feature Map (Bottleneck)', fontsize=10)\n                axes[i, 6].axis('off')\n                if i == num_samples - 1:\n                    cbar7 = plt.colorbar(im7, ax=axes[i, 6], fraction=0.046, pad=0.04)\n                    cbar7.set_label('Feature Response', fontsize=8)\n            break\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=300, bbox_inches='tight')\n    writer.add_figure(f'{phase}/Predictions_With_Attn_Feat', fig, epoch)\n    plt.close(fig)\n# ==============================\n# 7. 主函数\n# 作用：整合所有模块，实现端到端的训练流程\n# ==============================\ndef main():\n    global config # 全局配置变量（供采样器、数据集等模块调用）\n    # 根据设备类型自动调整核心参数（GPU/CPU性能差异适配）\n    batch_size = 16 if use_cuda else 4 # GPU批大小16（兼顾速度与内存），CPU批大小4（避免内存溢出）\n    learning_rate = 1e-4 if use_cuda else 5e-5 # GPU学习率稍大（收敛更快），CPU学习率较小（避免震荡）\n    base_channels = 48 if use_cuda else 16 # GPU用更多基础通道（模型更复杂），CPU用更少通道（降低计算量）\n    # 核心配置参数（可根据需求调整）\n    config = {\n        'img_size': 256, # 输入图像尺寸（固定为256x256，适配模型结构）\n        'batch_size': batch_size,\n        'num_epochs': 30, # 总训练轮次（根据数据量调整，30轮适合LGG数据集）\n        'learning_rate': learning_rate,\n        'base_channels': base_channels,\n        'data_root': \"/kaggle/working/data\", # 数据根目录（本地可改为'./data'）\n        'save_dir': \"/kaggle/working/results\", # 结果保存目录（模型、图表、日志）\n        'log_dir': \"/kaggle/working/logs\", # TensorBoard日志目录\n        'target_dice': 0.90, # 早停目标Dice（达到后提前终止训练）\n        'resume_checkpoint': \"/kaggle/working/resume_checkpoint.pth\" # 断点续训检查点路径（可选）\n    }\n    # 创建必要目录（不存在则自动创建，避免报错）\n    os.makedirs(config['save_dir'], exist_ok=True)\n    os.makedirs(config['log_dir'], exist_ok=True)\n    print(\"=\" * 60)\n    print(\"当前训练配置:\")\n    for key, value in config.items():\n        print(f\" {key}: {value}\")\n    print(\"=\" * 60)\n    # 初始化TensorBoard（可视化训练过程：损失曲线、指标变化、预测结果）\n    writer = SummaryWriter(config['log_dir'])\n    # --------------------------\n    # 步骤1：加载并预处理数据\n    # --------------------------\n    print(\"\\n[步骤1/5] 加载LGG脑肿瘤数据集...\")\n    try:\n        # 加载数据集并分割为训练/验证/测试集（7:2:1比例）\n        train_imgs, val_imgs, test_imgs, train_masks, val_masks, test_masks = load_lgg_data(config['data_root'])\n        # 获取数据增强/标准化策略（训练集增强，验证/测试集仅标准化）\n        train_transform, val_transform = get_transforms(config['img_size'])\n        # 创建数据集实例（传入config，供采样器获取参数）\n        train_dataset = LGGDataset(\n            train_imgs, train_masks, train_transform,\n            config=config # 关键：将配置传入数据集，供采样器调用\n        )\n        val_dataset = LGGDataset(val_imgs, val_masks, val_transform, config=config)\n        test_dataset = LGGDataset(test_imgs, test_masks, val_transform, config=config) # 测试集用验证集的转换（无增强）\n        # 计算训练集总batch数（基于1.5倍数据集大小，增加数据多样性）\n        total_train_samples = len(train_dataset) * 1.5\n        total_train_batches = int(total_train_samples // config['batch_size'])\n        # 创建训练集DataLoader（使用鲁棒版平衡采样器）\n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=config['batch_size'],\n            sampler=TumorBalancedSampler(\n                dataset=train_dataset,\n                tumor_ratio=0.6, # 目标肿瘤样本占比（60%，可根据数据不平衡程度调整）\n                file_read_timeout=5, # 掩码读取超时时间（5秒，避免阻塞）\n                min_lesion_pixels=5 # 判定肿瘤样本的最小像素数（过滤微小噪声）\n            ),\n            num_workers=2 if use_cuda else 0, # GPU用2个进程加载数据（加速），CPU用0（避免进程竞争）\n            pin_memory=use_cuda, # GPU启用内存锁定（加速数据从CPU到GPU的传输）\n            drop_last=True # 丢弃最后一个不完整的batch（避免batch大小不一致导致报错）\n        )\n        # 创建验证集/测试集DataLoader（无需平衡采样，按顺序加载）\n        val_loader = DataLoader(\n            val_dataset,\n            batch_size=config['batch_size'],\n            shuffle=False, # 验证集不打乱，便于复现结果\n            num_workers=2 if use_cuda else 0,\n            pin_memory=use_cuda\n        )\n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=config['batch_size'],\n            shuffle=False,\n            num_workers=2 if use_cuda else 0,\n            pin_memory=use_cuda\n        )\n        print(\n            f\"数据加载完成: 训练集{len(train_dataset)}样本 | 验证集{len(val_dataset)}样本 | 测试集{len(test_dataset)}样本\")\n    except Exception as e:\n        # 数据加载失败时（如路径错误、无数据），创建虚拟数据集用于测试代码\n        print(f\"[警告] 真实数据加载失败: {str(e)[:100]}\")\n        print(\"[备用方案] 创建虚拟数据集用于代码测试...\")\n        class DummyDataset(Dataset):\n            \"\"\"虚拟数据集：生成随机图像和掩码（每5个样本含1个肿瘤样本）\"\"\"\n            def __init__(self, size=100, config=None):\n                self.size = size\n                self.config = config or {}\n            def __len__(self):\n                return self.size\n            def __getitem__(self, idx):\n                # 随机生成3通道图像（模拟RGB）\n                img = torch.randn(3, 256, 256)\n                # 生成掩码：每5个样本中1个含肿瘤区域（50x50像素）\n                mask = torch.zeros(1, 256, 256)\n                if idx % 5 == 0:\n                    mask[:, 50:100, 50:100] = 1.0 # 肿瘤区域（固定位置，便于测试）\n                return img, mask\n        # 创建虚拟数据集和加载器\n        train_dataset = DummyDataset(size=200, config=config)\n        val_dataset = DummyDataset(size=50, config=config)\n        test_dataset = DummyDataset(size=30, config=config)\n        total_train_samples = len(train_dataset) * 1.5\n        total_train_batches = int(total_train_samples // config['batch_size'])\n        # 虚拟训练集用平衡采样器（保持逻辑一致）\n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=config['batch_size'],\n            sampler=TumorBalancedSampler(train_dataset, tumor_ratio=0.6),\n            num_workers=0 # 虚拟数据无需多进程加载\n        )\n        val_loader = DataLoader(val_dataset, batch_size=config['batch_size'], shuffle=False)\n        test_loader = DataLoader(test_dataset, batch_size=config['batch_size'], shuffle=False)\n    # --------------------------\n    # 步骤2：初始化模型与优化器\n    # --------------------------\n    print(\"\\n[步骤2/5] 初始化模型、优化器与损失函数...\")\n    # 创建模型实例（轻量级多尺度注意力UNet）\n    model = LightweightMultiScaleAttentionUNet(\n        img_size=config['img_size'],\n        in_channels=3, # 输入3通道（RGB图像）\n        out_channels=1, # 输出1通道（二分类：肿瘤/背景）\n        base_channels=config['base_channels']\n    )\n    # 多GPU支持（若有多个GPU，自动并行训练）\n    if use_cuda and torch.cuda.device_count() > 1:\n        print(f\"检测到{torch.cuda.device_count()}个GPU，启用多GPU并行训练\")\n        model = nn.DataParallel(model) # 包装为多GPU模型\n    model = model.to(device) # 将模型移至计算设备（GPU/CPU）\n    # 计算模型参数数量（评估模型复杂度）\n    total_params = sum(p.numel() for p in model.parameters())\n    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"模型参数统计: 总参数{total_params:,} | 可训练参数{trainable_params:,}\")\n    # 定义学习率调度器（带重启的余弦退火：兼顾收敛速度与泛化性）\n    class CosineAnnealingWithRestartsLR(torch.optim.lr_scheduler.CosineAnnealingLR):\n        def __init__(self, optimizer, T_max, T_mult=1.5, eta_min=5e-6, last_epoch=-1):\n            self.T_mult = T_mult # 每次重启后，周期扩大为原来的1.5倍\n            super().__init__(optimizer, T_max, eta_min, last_epoch)\n        def step(self, epoch=None):\n            # 周期结束时自动重启并扩大周期\n            if self.last_epoch == self.T_max - 1:\n                self.T_max *= self.T_mult\n                self.last_epoch = -1 # 重置计数器\n            super().step(epoch)\n    # 优化器：AdamW（带权重衰减，抑制过拟合）\n    optimizer = optim.AdamW(\n        model.parameters(),\n        lr=config['learning_rate'],\n        weight_decay=1e-4 # 权重衰减系数（控制过拟合）\n    )\n    # 调度器：初始周期15轮，最小学习率5e-6\n    scheduler = CosineAnnealingWithRestartsLR(\n        optimizer,\n        T_max=15,\n        T_mult=1.5,\n        eta_min=5e-6\n    )\n    # 损失函数：Focal-Dice损失（解决类别不平衡，优化分割边界）\n    criterion = FocalDiceLoss(\n        alpha=0.85, # 正样本权重（肿瘤样本占比低，权重更高）\n        gamma=2.0, # 聚焦参数（降低易分类样本权重）\n        bg_penalty=0.3, # 背景假阳性惩罚（减少误判）\n        edge_weight=1.2 # 肿瘤边缘权重（提高边界精度）\n    )\n    # 混合精度训练（仅GPU支持，加速训练并减少内存占用）\n    scaler = None\n    if use_cuda:\n        try:\n            scaler = torch.amp.GradScaler()\n        except ImportError:\n            scaler = GradScaler() # 旧版本兼容\n        print(\"启用混合精度训练（FP16），加速训练并降低内存占用\")\n    else:\n        print(\"CPU环境不支持混合精度训练，禁用该功能\")\n    # --------------------------\n    # 步骤3：断点续训（可选）\n    # --------------------------\n    print(\"\\n[步骤3/5] 检查断点续训状态...\")\n    start_epoch = 0 # 起始训练轮次（默认从0开始）\n    best_dice = 0.0 # 最佳验证集Dice（用于保存最优模型）\n    # 训练记录（用于绘制曲线）\n    train_losses, val_losses = [], []\n    train_dices, val_dices = [], []\n    val_ious, val_sensitivities, val_specificities = [], [], []\n    # 若存在检查点，加载训练状态（模型、优化器、指标记录）\n    if config['resume_checkpoint'] and os.path.exists(config['resume_checkpoint']):\n        print(f\"从检查点恢复训练: {config['resume_checkpoint']}\")\n        checkpoint = torch.load(config['resume_checkpoint'], map_location=device) # 加载到当前设备\n        # 加载模型参数（处理单GPU/多GPU兼容问题）\n        if isinstance(model, nn.DataParallel) and not 'module.' in list(checkpoint['model_state_dict'].keys())[0]:\n            # 多GPU模型加载单GPU训练的参数\n            model.module.load_state_dict(checkpoint['model_state_dict'])\n        else:\n            model.load_state_dict(checkpoint['model_state_dict'])\n        # 加载优化器和调度器状态\n        optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n        scheduler.load_state_dict(checkpoint['scheduler_state_dict'])\n        # 加载混合精度scaler状态（若存在）\n        if 'scaler_state_dict' in checkpoint and scaler is not None:\n            scaler.load_state_dict(checkpoint['scaler_state_dict'])\n        # 加载训练记录\n        start_epoch = checkpoint['epoch'] + 1 # 从下一轮开始\n        best_dice = checkpoint.get('best_dice', 0.0)\n        train_losses = checkpoint.get('train_losses', [])\n        val_losses = checkpoint.get('val_losses', [])\n        train_dices = checkpoint.get('train_dices', [])\n        val_dices = checkpoint.get('val_dices', [])\n        val_ious = checkpoint.get('val_ious', [])\n        val_sensitivities = checkpoint.get('val_sensitivities', [])\n        val_specificities = checkpoint.get('val_specificities', [])\n        print(f\"恢复完成：从第{start_epoch}轮开始，当前最佳Dice={best_dice:.4f}\")\n    else:\n        print(\"未找到有效检查点，将从第0轮开始训练\")\n    # --------------------------\n    # 步骤4：模型测试（确保前向传播正常）\n    # --------------------------\n    print(\"\\n[步骤4/5] 测试模型前向传播是否正常...\")\n    try:\n        model.eval() # 切换到评估模式\n        with torch.no_grad():\n            # 生成随机测试输入（模拟2个样本的RGB图像）\n            test_input = torch.randn(2, 3, config['img_size'], config['img_size']).to(device)\n            test_output = model(test_input) # 模型输出\n            # 验证输出尺寸是否正确（应与输入尺寸一致）\n            assert test_output.shape == (2, 1, config['img_size'], config['img_size']), \\\n                f\"模型输出尺寸错误：预期(2,1,256,256)，实际{test_output.shape}\"\n            # 测试损失函数计算是否正常\n            test_target = (torch.randn_like(test_output).sigmoid() > 0.5).float() # 随机生成目标掩码\n            test_loss = criterion(test_output, test_target)\n            assert not torch.isnan(test_loss), \"损失函数计算出现NaN（异常）\"\n        print(f\"模型测试通过：输入尺寸{test_input.shape} → 输出尺寸{test_output.shape}，测试损失{test_loss.item():.4f}\")\n    except Exception as e:\n        print(f\"[错误] 模型前向传播测试失败：{e}\")\n        return # 模型异常，终止训练\n    # --------------------------\n    # 步骤5：开始训练循环\n    # --------------------------\n    print(\"\\n[步骤5/5] 开始训练（共{}轮）...\".format(config['num_epochs']))\n    start_time = time.time() # 记录训练开始时间\n    try:\n        for epoch in range(start_epoch, config['num_epochs']):\n            print(\"\\n\" + \"=\" * 50)\n            print(f\"Epoch {epoch + 1}/{config['num_epochs']}\")\n            print(\"=\" * 50)\n            # 1. 训练一轮（更新模型参数）\n            train_loss, train_dice = train_epoch(\n                model=model,\n                train_loader=train_loader,\n                criterion=criterion,\n                optimizer=optimizer,\n                device=device,\n                scaler=scaler,\n                total_batches=total_train_batches,\n                writer=writer,\n                epoch=epoch\n            )\n            # 2. 验证一轮（评估模型性能，不更新参数）\n            val_loss, val_dice, val_metrics = validate_epoch(\n                model=model,\n                val_loader=val_loader,\n                criterion=criterion,\n                device=device,\n                writer=writer,\n                epoch=epoch,\n                phase='Validation'\n            )\n            # 3. 更新学习率（按调度器策略）\n            scheduler.step()\n            # 4. 记录训练/验证指标（用于后续绘图）\n            train_losses.append(train_loss)\n            val_losses.append(val_loss)\n            train_dices.append(train_dice)\n            val_dices.append(val_dice)\n            val_ious.append(val_metrics['iou'])\n            val_sensitivities.append(val_metrics['sensitivity'])\n            val_specificities.append(val_metrics['specificity'])\n            # 5. 打印当前轮次结果（直观查看训练进度）\n            print(f\"\\n[本轮结果]\")\n            print(f\" 训练集：损失={train_loss:.4f} | Dice={train_dice:.4f}\")\n            print(f\" 验证集：损失={val_loss:.4f} | Dice={val_dice:.4f}\")\n            print(f\" 关键指标：IoU={val_metrics['iou']:.4f} | 灵敏度={val_metrics['sensitivity']:.4f}\")\n            # 6. 保存最佳模型（基于验证集Dice，仅保存性能更优的模型）\n            if val_dice > best_dice:\n                best_dice = val_dice\n                # 保存的内容：模型参数、优化器状态、训练配置等（便于后续续训）\n                save_dict = {\n                    'epoch': epoch,\n                    'model_state_dict': model.state_dict(),\n                    'optimizer_state_dict': optimizer.state_dict(),\n                    'scheduler_state_dict': scheduler.state_dict(),\n                    'best_dice': best_dice,\n                    'config': config,\n                    'train_losses': train_losses,\n                    'val_losses': val_losses,\n                    'train_dices': train_dices,\n                    'val_dices': val_dices\n                }\n                # 加入混合精度scaler状态（若使用）\n                if scaler is not None:\n                    save_dict['scaler_state_dict'] = scaler.state_dict()\n                # 保存到指定路径\n                best_model_path = os.path.join(config['save_dir'], 'best_model.pth')\n                torch.save(save_dict, best_model_path)\n                print(f\"[模型保存] 最佳模型已更新（Dice={best_dice:.4f}），保存路径：{best_model_path}\")\n            # 7. 定期保存可视化结果和检查点（每10轮或最后一轮，避免频繁IO）\n            if (epoch + 1) % 10 == 0 or epoch == config['num_epochs'] - 1:\n                # a. 保存训练曲线（损失+Dice）\n                curve_path = os.path.join(config['save_dir'], f'training_curves_epoch_{epoch + 1}.png')\n                plot_training_curves(train_losses, val_losses, train_dices, val_dices, curve_path)\n                print(f\"[可视化] 训练曲线已保存：{curve_path}\")\n                # b. 保存预测结果可视化（输入/真实/预测/叠加/热力图）\n                pred_viz_path = os.path.join(config['save_dir'], f'predictions_epoch_{epoch + 1}.png')\n                visualize_predictions(\n                    model=model,\n                    val_loader=val_loader,\n                    device=device,\n                    save_path=pred_viz_path,\n                    writer=writer,\n                    epoch=epoch,\n                    num_samples=4 # 每次可视化4个样本\n                )\n                print(f\"[可视化] 预测结果已保存：{pred_viz_path}\")\n                # c. 保存额外指标曲线（IoU、灵敏度、特异度）\n                metrics_path = os.path.join(config['save_dir'], f'additional_metrics_epoch_{epoch + 1}.png')\n                plot_additional_metrics(val_ious, val_sensitivities, val_specificities, metrics_path)\n                print(f\"[可视化] 额外指标曲线已保存：{metrics_path}\")\n                # d. 保存检查点（用于断点续训）\n                checkpoint_path = os.path.join(config['save_dir'], f'checkpoint_epoch_{epoch + 1}.pth')\n                checkpoint_dict = save_dict.copy() # 复用最佳模型的保存字典，更新epoch和指标\n                checkpoint_dict['epoch'] = epoch\n                checkpoint_dict['train_losses'] = train_losses\n                checkpoint_dict['val_losses'] = val_losses\n                torch.save(checkpoint_dict, checkpoint_path)\n                print(f\"[检查点] 第{epoch + 1}轮检查点已保存：{checkpoint_path}\")\n            # 8. 早停机制（达到目标Dice，提前终止训练，节省时间）\n            if val_dice >= config['target_dice']:\n                print(f\"\\n[早停] 验证集Dice={val_dice:.4f}达到目标{config['target_dice']}，提前终止训练！\")\n                break\n    except Exception as e:\n        # 训练过程中出现异常，打印错误信息并保存当前状态\n        print(f\"\\n[紧急错误] 训练过程中断：{e}\")\n        import traceback\n        traceback.print_exc() # 打印详细错误堆栈（便于定位问题）\n        # 保存紧急检查点（避免训练成果丢失）\n        emergency_path = os.path.join(config['save_dir'], 'emergency_checkpoint.pth')\n        torch.save({\n            'epoch': epoch,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'train_losses': train_losses,\n            'val_losses': val_losses\n        }, emergency_path)\n        print(f\"[紧急保存] 中断前状态已保存至：{emergency_path}\")\n    # --------------------------\n    # 训练结束：测试集最终评估\n    # --------------------------\n    print(\"\\n\" + \"=\" * 60)\n    print(\"训练流程结束，开始在测试集上评估最终性能...\")\n    print(\"=\" * 60)\n    # 在测试集上评估（使用最佳模型参数）\n    # 加载最佳模型（确保用最优参数评估）\n    best_model_path = os.path.join(config['save_dir'], 'best_model.pth')\n    if os.path.exists(best_model_path):\n        best_checkpoint = torch.load(best_model_path, map_location=device)\n        if isinstance(model, nn.DataParallel) and not 'module.' in list(best_checkpoint['model_state_dict'].keys())[0]:\n            model.module.load_state_dict(best_checkpoint['model_state_dict'])\n        else:\n            model.load_state_dict(best_checkpoint['model_state_dict'])\n        print(f\"已加载最佳模型（Dice={best_checkpoint['best_dice']:.4f}）用于测试集评估\")\n    # 测试集评估\n    test_loss, test_dice, test_metrics = validate_epoch(\n        model=model,\n        val_loader=test_loader,\n        criterion=criterion,\n        device=device,\n        writer=writer,\n        epoch=config['num_epochs'],\n        phase='Test' # 标记为测试集，便于TensorBoard区分\n    )\n    # 打印测试集最终结果\n    print(f\"\\n[测试集最终结果]\")\n    print(f\" 损失：{test_loss:.4f}\")\n    print(f\" Dice系数：{test_dice:.4f}（核心分割指标）\")\n    print(f\" IoU：{test_metrics['iou']:.4f}（重叠度指标）\")\n    print(f\" 灵敏度：{test_metrics['sensitivity']:.4f}（肿瘤召回率）\")\n    print(f\" 特异度：{test_metrics['specificity']:.4f}（背景准确率）\")\n    print(f\" 精确率：{test_metrics['precision']:.4f}（预测肿瘤准确率）\")\n    print(f\" F1分数：{test_metrics['f1']:.4f}（精确率与灵敏度平衡）\")\n    # 保存测试集预测结果可视化\n    test_viz_path = os.path.join(config['save_dir'], 'test_set_predictions.png')\n    visualize_predictions(\n        model=model,\n        val_loader=test_loader,\n        device=device,\n        save_path=test_viz_path,\n        writer=writer,\n        epoch=config['num_epochs'],\n        num_samples=4,\n        phase='Test'\n    )\n    print(f\"\\n[结果保存] 测试集预测可视化已保存：{test_viz_path}\")\n    # --------------------------\n    # 生成训练总结报告\n    # --------------------------\n    total_time = (time.time() - start_time) / 60 # 总训练时间（分钟）\n    report_path = os.path.join(config['save_dir'], 'training_report.txt')\n    with open(report_path, 'w', encoding='utf-8') as f:\n        f.write(\"=\" * 60 + \"\\n\")\n        f.write(\"轻量级多尺度注意力UNet训练报告\\n\")\n        f.write(\"=\" * 60 + \"\\n\")\n        f.write(f\"训练时间：{total_time:.2f} 分钟\\n\")\n        f.write(f\"总轮次：{epoch + 1} 轮（原计划{config['num_epochs']}轮）\\n\")\n        f.write(f\"最佳验证集Dice：{best_dice:.4f}\\n\")\n        f.write(\"\\n[测试集最终性能]\\n\")\n        f.write(f\" 损失：{test_loss:.4f}\\n\")\n        f.write(f\" Dice系数：{test_dice:.4f}\\n\")\n        f.write(f\" IoU：{test_metrics['iou']:.4f}\\n\")\n        f.write(f\" 灵敏度：{test_metrics['sensitivity']:.4f}\\n\")\n        f.write(f\" 特异度：{test_metrics['specificity']:.4f}\\n\")\n        f.write(f\" 精确率：{test_metrics['precision']:.4f}\\n\")\n        f.write(f\" F1分数：{test_metrics['f1']:.4f}\\n\")\n        f.write(\"\\n[训练配置]\\n\")\n        for key, value in config.items():\n            f.write(f\" {key}：{value}\\n\")\n    print(f\"\\n[报告保存] 训练总结报告已保存：{report_path}\")\n    print(\"\\n\" + \"=\" * 60)\n    print(\"所有流程完成！可在以下路径查看结果：\")\n    print(f\" - 模型文件：{config['save_dir']}\")\n    print(f\" - 可视化图表：{config['save_dir']}\")\n    print(f\" - TensorBoard日志：{config['log_dir']}\")\n    print(\"=\" * 60)\n    # 关闭TensorBoard（释放资源）\n    writer.close()\n# 程序入口（确保仅在直接运行时执行训练流程）\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}