{"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":"gpu","dataSources":[{"sourceType":"competition","sourceId":91844,"databundleVersionId":11361821},{"sourceType":"datasetVersion","sourceId":11442943,"datasetId":7161252,"databundleVersionId":11882876}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# -*- coding: utf-8 -*-\n\"\"\"\n训练脚本 for BirdCLEF 2025 - EfficientNet B0\n\n基于预计算的多批次梅尔频谱图进行训练。\n数据存储在 /kaggle/input/zwy-bird/precomputed_melspec_batch*.npy\n参考脚本: [Train]-EfficientNet B0 Pytorch .py\n\"\"\"\n\"\"\"\n输出什么时候保存？\n\n1.  **最新检查点 (`..._latest.pth`)**:\n    *   **何时保存**: 在 **每个 Epoch 结束时** 都会保存或更新一次。\n    *   **保存内容**: 包含恢复训练所需的完整信息，包括模型权重、优化器状态、学习率调度器状态、当前 Epoch 数、以及记录的最佳验证 AUC 等。\n    *   **目的**: 主要用于**断点续练**。如果训练中断，可以从这个文件恢复。\n    *   **文件名**: 由 `cfg.checkpoint_filename` 定义，例如 `efficientnet_b0_fold0_latest.pth`。\n\n2.  **最佳模型 (`..._best.pth`)**:\n    *   **何时保存**: 只有当**当前 Epoch 的验证集 AUC (Area Under Curve) 分数 高于 此 Fold 之前所有 Epoch 的最佳 AUC 分数时**，才会保存或更新。\n    *   **保存内容**: 默认只保存了模型的权重 (`model.state_dict()` 或 `model.module.state_dict()`)。脚本中也注释掉了保存完整检查点的选项。只保存权重通常用于后续的推理或评估。\n    *   **目的**: 保存**验证集上表现最好**的模型，通常用于最终的预测或模型集成。\n    *   **文件名**: 例如 `efficientnet_b0_fold0_best.pth`。\n\n这两个文件都会被保存在 `cfg.OUTPUT_DIR` 指定的目录中，默认是 `/kaggle/working/`。\n\n简单来说：\n\n*   **每轮结束**存一个最新的进度 (`_latest.pth`)。\n*   **表现有提升时**存一个最好的模型 (`_best.pth`)。\n\n\"\"\"\nimport os\nimport random\nimport gc\nimport time\nimport glob  # 用于查找匹配的文件\nimport numpy as np\nimport pandas as pd\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import roc_auc_score\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom tqdm.auto import tqdm\nimport warnings\nimport cv2  # 确保导入 cv2\nimport torch.nn.functional as F  # 确保导入 F\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nwarnings.filterwarnings(\"ignore\")\n\n\n# --- 配置 (CFG) ---\nclass CFG:\n    # 基本设置\n    seed = 42\n    debug = False  # 是否开启 Debug 模式 (少量数据, 短 epoch)\n    apex = False  # 是否使用 Apex 进行混合精度训练\n    print_freq = 100  # 打印频率\n    num_workers = 2  # DataLoader 使用的进程数\n\n    # 路径设置\n    OUTPUT_DIR = \"/kaggle/working/\"  # 模型和日志输出目录\n    train_csv = \"/kaggle/input/birdclef-2025/train.csv\"  # 训练标签 CSV 文件\n    taxonomy_csv = \"/kaggle/input/birdclef-2025/taxonomy.csv\"  # 物种分类 CSV 文件\n    # !!! 修改为你包含所有 batch npy 文件的目录 !!!\n    INPUT_SPEC_DIR = \"/kaggle/input/zwy-bird/\"  # 预计算频谱图 batch 文件所在目录\n    # 例如: /kaggle/input/your-merged-dataset-name/\n\n    # 模型设置\n    model_name = \"efficientnet_b0\"  # 使用的模型名称 (来自 timm)\n    pretrained = True  # 是否加载 timm 提供的预训练权重\n    in_channels = 1  # 输入通道数 (灰度梅尔频谱图为 1)\n    num_classes = 182  # 目标类别数 (根据 taxonomy.csv 确定)\n\n    # 数据集设置 (因为是预计算好的，所以不需要音频处理参数)\n    TARGET_SHAPE = (256, 256)  # 预计算频谱图的目标形状 (需要与预计算时一致)\n\n    # 训练设置\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    epochs = 10  # 训练轮数\n    batch_size = 32  # 批处理大小\n    criterion = \"BCEWithLogitsLoss\"  # 损失函数 (多标签分类常用)\n\n    # --- 新增：断点续练设置 ---\n    resume_training = True  # 是否尝试从检查点恢复训练\n    checkpoint_filename = \"{model_name}_fold{fold}_latest.pth\"  # 检查点文件名格式\n    # --- 结束新增 ---\n\n    # 交叉验证设置\n    n_fold = 5  # 交叉验证折数\n    selected_folds = [0]  # 选择训练的折数 (例如 [0, 1, 2, 3, 4] 训练所有折)\n\n    # 优化器设置\n    optimizer = \"AdamW\"  # 优化器类型\n    lr = 1e-3  # 学习率 (AdamW 的推荐值，可以调整)\n    weight_decay = 1e-5  # 权重衰减\n\n    # 学习率调度器设置\n    scheduler = \"CosineAnnealingLR\"  # 学习率调度器类型\n    min_lr = 1e-6  # 最小学习率 (用于 CosineAnnealingLR)\n    T_max = epochs  # CosineAnnealingLR 的周期 (通常设为 epochs)\n\n    # 数据增强设置 (在 Dataset 中实现)\n    aug_prob = 0.5  # 应用频谱图增强的概率\n    mixup_alpha = 0.0  # Mixup 参数 (0 表示不使用 Mixup)\n\n    def update_debug_settings(self):\n        \"\"\"如果 debug=True, 则减少 epochs 和 folds\"\"\"\n        if self.debug:\n            self.epochs = 2\n            self.selected_folds = [0]\n            self.batch_size = 16  # Debug 时减小 batch size\n\n\n# --- 工具函数 ---\ndef set_seed(seed=42):\n    \"\"\"设置随机种子以保证可复现性\"\"\"\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\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\n\ndef calculate_auc(targets, outputs):\n    \"\"\"计算多标签分类的平均 AUC\"\"\"\n    num_classes = targets.shape[1]\n    aucs = []\n\n    # 计算概率 (Sigmoid)\n    probs = 1 / (1 + np.exp(-outputs))\n\n    for i in range(num_classes):\n        # 检查该类别是否有正样本\n        if (\n            np.sum(targets[:, i]) > 0 and np.sum(1 - targets[:, i]) > 0\n        ):  # 需要正负样本才能计算 AUC\n            try:\n                class_auc = roc_auc_score(targets[:, i], probs[:, i])\n                aucs.append(class_auc)\n            except ValueError as e:\n                # print(f\"计算类别 {i} 的 AUC 时出错: {e}\") # 可以取消注释以调试\n                pass  # 如果某个类别无法计算 AUC，则跳过\n\n    return np.mean(aucs) if aucs else 0.0  # 如果没有可计算的 AUC，返回 0\n\n\n# --- !!! 新增：将索引构建移出类外作为辅助函数 !!! ---\ndef build_sample_index(batch_files):\n    \"\"\"加载所有 batch 文件的 keys，建立 sample 到 batch 文件路径的映射。\"\"\"\n    sample_map = {}\n    print(\"正在建立样本索引...\")\n    for batch_path in tqdm(batch_files, desc=\"索引 Batch 文件\"):\n        try:\n            # 只加载 keys，避免加载整个 batch 耗费内存\n            # 注意：标准 np.load 可能仍会加载整个字典，但只取 keys 还是比全加载好\n            batch_keys = list(np.load(batch_path, allow_pickle=True).item().keys())\n            for key in batch_keys:\n                if key in sample_map:\n                    # print(\n                    #     f\"警告: 样本 '{key}' 在多个 batch 文件中找到。将使用路径: {batch_path}\"\n                    # )\n                    pass\n                sample_map[key] = batch_path\n        except Exception as e:\n            print(f\"建立索引时加载 batch 文件 {batch_path} 出错: {e}\")\n    return sample_map\n\n\n# --- 数据集类 (BirdCLEFDatasetFromBatches) ---\nclass BirdCLEFDatasetFromBatches(Dataset):\n    \"\"\"\n    从多个预计算的 .npy batch 文件加载梅尔频谱图的数据集类。\n    采用按需加载策略优化内存使用。\n    接收预构建的索引以避免重复扫描。\n    \"\"\"\n\n    # --- 修改：接收 batch_files 和 sample_index ---\n    def __init__(self, df, cfg, mode=\"train\", batch_files=None, sample_index=None):\n        \"\"\"\n        Args:\n                df (pd.DataFrame): 包含样本信息 (filename, primary_label 等) 的 DataFrame。\n                cfg (CFG): 配置对象。\n                mode (str): 'train' 或 'valid'/'test'。\n                batch_files (list, optional): 预扫描的 batch 文件路径列表。\n                sample_index (dict, optional): 预构建的 sample_name 到 batch_path 的映射。\n        \"\"\"\n        self.df = df.copy()\n        self.cfg = cfg\n        self.mode = mode\n\n        # 加载分类信息\n        taxonomy_df = pd.read_csv(self.cfg.taxonomy_csv)\n        self.species_ids = taxonomy_df[\"primary_label\"].tolist()\n        self.num_classes = len(self.species_ids)\n        self.label_to_idx = {label: idx for idx, label in enumerate(self.species_ids)}\n\n        # --- 优化：使用传入的索引和文件列表 ---\n        if batch_files is not None and sample_index is not None:\n            print(\n                f\"使用预构建的索引 (含 {len(sample_index)} 个样本) 和 {len(batch_files)} 个 batch 文件路径。\"\n            )\n            self.batch_files = batch_files\n            self.sample_to_batch_info = sample_index\n        else:\n            # --- Fallback (理论上不应执行，除非直接调用 Dataset 类) ---\n            print(\n                \"警告: 未提供预构建索引或文件列表，将在 Dataset 内部重新扫描和构建...\"\n            )\n            self.batch_files = sorted(\n                glob.glob(\n                    os.path.join(\n                        self.cfg.INPUT_SPEC_DIR, \"precomputed_melspec_batch*.npy\"\n                    )\n                )\n            )\n            if not self.batch_files:\n                raise FileNotFoundError(\n                    f\"在目录下未找到任何 precomputed_melspec_batch*.npy 文件: {self.cfg.INPUT_SPEC_DIR}\"\n                )\n            # 需要一个内部的 build_index 或调用外部函数\n            # 为避免混淆，这里直接报错，强制要求从外部传入\n            raise ValueError(\"必须通过构造函数提供 batch_files 和 sample_index！\")\n            # self.sample_to_batch_info = build_sample_index(self.batch_files) # 或者这样写，但不推荐\n\n        # --- 根据最终使用的索引过滤 DataFrame ---\n        self._filter_df_by_index()\n\n        # 检查传入的 df 是否确实有 samplename (过滤后检查)\n        if \"samplename\" not in self.df.columns:\n            raise ValueError(\"错误: DataFrame (过滤后) 必须包含 'samplename' 列！\")\n\n        if cfg.debug and mode == \"train\":  # Debug 模式下减少训练数据量\n            sample_size = min(1000, len(self.df))\n            self.df = self.df.sample(sample_size, random_state=cfg.seed).reset_index(\n                drop=True\n            )\n            print(f\"Debug模式：使用 {len(self.df)} 个训练样本\")\n\n        # --- 优化：初始化缓存 ---\n        self.current_batch_path = None\n        self.current_batch_data = None\n\n    # --- 移除 _build_index 方法 ---\n    # def _build_index(self): ... # 不再需要\n\n    def _filter_df_by_index(self, key_column=\"samplename\"):\n        # ... (方法内容不变，使用 self.sample_to_batch_info) ...\n        original_len = len(self.df)\n        if key_column in self.df.columns:\n            available_samples = set(self.sample_to_batch_info.keys())\n            self.df = self.df[self.df[key_column].isin(available_samples)].reset_index(\n                drop=True\n            )\n            filtered_len = len(self.df)\n            if filtered_len < original_len:\n                print(\n                    f\"过滤完成: 从 {original_len} 个样本中移除了 {original_len - filtered_len} 个在预计算数据索引中找不到的样本。剩余 {filtered_len} 个样本。\"\n                )\n            if filtered_len == 0:\n                print(\n                    f\"警告: 过滤后数据集为空！请检查 '{key_column}' 列与 .npy 文件 key 是否匹配。\"\n                )\n        else:\n            print(\n                f\"警告: DataFrame 中缺少用于过滤的列 '{key_column}'。跳过基于索引的过滤。\"\n            )\n\n    def __len__(self):\n        return len(self.df)\n\n    def _load_batch(self, batch_path):\n        \"\"\"加载指定的 batch 文件并更新缓存。\"\"\"\n        # print(f\"缓存未命中，加载 batch: {os.path.basename(batch_path)}\") # Debugging line\n        try:\n            self.current_batch_data = np.load(batch_path, allow_pickle=True).item()\n            self.current_batch_path = batch_path\n            # gc.collect() # 可以取消注释以更积极地释放内存，但可能影响性能\n        except Exception as e:\n            print(f\"加载 batch 文件 {batch_path} 时出错: {e}\")\n            # 如果加载失败，清空缓存避免使用错误数据\n            self.current_batch_path = None\n            self.current_batch_data = None\n            return False\n        return True\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        # 假设 'samplename' 列存在且包含正确的 key\n        samplename = row[\"samplename\"]\n\n        spec = None\n        if samplename in self.sample_to_batch_info:\n            target_batch_path = self.sample_to_batch_info[samplename]\n\n            # --- 优化：检查并加载 batch ---\n            if target_batch_path != self.current_batch_path:\n                if not self._load_batch(target_batch_path):\n                    # 加载失败，返回空谱图\n                    print(\n                        f\"错误: 无法加载样本 {samplename} 所在的 batch 文件 {target_batch_path}。\"\n                    )\n                    spec = np.zeros(self.cfg.TARGET_SHAPE, dtype=np.float32)\n\n            # --- 优化：从缓存中获取数据 ---\n            if (\n                self.current_batch_data is not None\n                and samplename in self.current_batch_data\n            ):\n                spec = self.current_batch_data[samplename]\n            elif spec is None:  # 如果加载成功但key不在里面（理论上不应发生）或加载失败\n                print(\n                    f\"警告: 样本 {samplename} 在其声称的 batch 文件 {target_batch_path} 中未找到！\"\n                )\n            spec = np.zeros(self.cfg.TARGET_SHAPE, dtype=np.float32)\n\n        else:\n            # 理论上不应该发生，因为我们已经根据索引过滤了 df\n            print(f\"警告: 样本 {samplename} 在索引中找不到！返回空频谱图。\")\n            spec = np.zeros(self.cfg.TARGET_SHAPE, dtype=np.float32)\n\n        # 确保频谱图是正确的形状和类型\n        if spec.shape != self.cfg.TARGET_SHAPE:\n            try:\n                # print(f\"警告: 样本 {samplename} 的频谱图形状 {spec.shape} 与目标形状 {self.cfg.TARGET_SHAPE} 不符。尝试调整大小。\")\n                spec = cv2.resize(\n                    spec, self.cfg.TARGET_SHAPE, interpolation=cv2.INTER_LINEAR\n                )\n            except Exception as e:\n                print(f\"错误: 调整样本 {samplename} 大小时出错 ({e})。返回空频谱图。\")\n                spec = np.zeros(self.cfg.TARGET_SHAPE, dtype=np.float32)\n\n        spec = torch.tensor(spec.astype(np.float32)).unsqueeze(0)  # 添加通道维度\n\n        # 应用数据增强 (仅训练模式)\n        if self.mode == \"train\" and random.random() < self.cfg.aug_prob:\n            spec = self.apply_spec_augmentations(spec)\n\n        # 编码标签 (应用新的权重逻辑)\n        target = self.encode_label(\n            row[\"primary_label\"], row.get(\"secondary_labels\")\n        )  # 使用 .get 处理可能不存在的列\n\n        return {\n            \"melspec\": spec,\n            \"target\": torch.tensor(target, dtype=torch.float32),\n            \"samplename\": samplename,  # 或其他你需要的元数据\n        }\n\n    def apply_spec_augmentations(self, spec):\n        \"\"\"对频谱图应用增强 (例如 Time/Frequency Masking)\"\"\"\n        # ... (可以从参考脚本复制或实现你需要的增强) ...\n        # 示例：频率掩码\n        if random.random() < 0.5:\n            num_masks = random.randint(1, 2)\n            for _ in range(num_masks):\n                height = random.randint(5, 25)  # 掩码高度\n                start = random.randint(0, spec.shape[1] - height)  # 频率轴起始位置\n                spec[0, start : start + height, :] = 0  # 将该区域置0\n        # 示例: 时间掩码\n        if random.random() < 0.5:\n            num_masks = random.randint(1, 2)\n            for _ in range(num_masks):\n                width = random.randint(5, 25)  # 掩码宽度\n                start = random.randint(0, spec.shape[2] - width)  # 时间轴起始位置\n                spec[0, :, start : start + width] = 0  # 将该区域置0\n        return spec\n\n    def encode_label(self, primary_label, secondary_labels=None):\n        \"\"\"\n        将主标签和次要标签编码为目标向量。\n        如果存在次要标签，主标签权重0.5，剩余0.5平均分配给次要标签。\n        \"\"\"\n        target = np.zeros(self.num_classes, dtype=np.float32)\n        valid_secondary = []\n\n        # 解析次要标签\n        # 修复 Linter Error: 确保列表格式正确\n        if (\n            secondary_labels is not None\n            and secondary_labels != [\"\"]\n            and not pd.isna(secondary_labels)\n        ):\n            if isinstance(secondary_labels, str):\n                try:\n                    # 尝试解析字符串形式的列表\n                    parsed_labels = eval(secondary_labels)\n                    if isinstance(parsed_labels, list):\n                        secondary_labels = parsed_labels\n                    else:\n                        secondary_labels = []  # 解析结果不是列表，视为空\n                except:\n                    secondary_labels = []  # 解析失败则视为空\n\n            if isinstance(secondary_labels, list):\n                # 筛选出有效的次要标签\n                valid_secondary = [\n                    l for l in secondary_labels if l in self.label_to_idx\n                ]\n\n        # 根据是否有有效的次要标签来分配权重\n        if primary_label in self.label_to_idx:\n            primary_idx = self.label_to_idx[primary_label]\n            if valid_secondary:\n                # 有次要标签：主标签权重 0.5\n                target[primary_idx] = 0.5\n                # 剩余 0.5 平均分配给次要标签\n                # 确保除数不为零\n                if len(valid_secondary) > 0:\n                    sec_weight = 0.5 / len(valid_secondary)\n                else:\n                    sec_weight = 0  # 不应该发生，但作为保险\n\n                for label in valid_secondary:\n                    # 注意：如果次要标签和主标签相同，这里会覆盖主标签的0.5权重\n                    # 如果不希望覆盖，可以用 max(target[...], sec_weight) 或其他逻辑\n                    target[\n                        self.label_to_idx[label]\n                    ] += sec_weight  # 使用 += 避免覆盖主标签权重（如果次要标签包含主标签）\n                # 修正可能存在的重复标签累加问题（例如主标签也在次要里）\n                target[primary_idx] = max(\n                    target[primary_idx], 0.5\n                )  # 确保主标签权重至少为0.5\n                # 归一化（可选，但推荐，以防权重总和略超1）\n                # current_sum = target.sum()\n                # if current_sum > 1e-6: # 避免除以零\n                #     target = target / current_sum\n\n            else:\n                # 没有次要标签：主标签权重 1.0\n                target[primary_idx] = 1.0\n\n        elif valid_secondary:  # 如果主标签无效，但有次要标签\n            print(\n                f\"警告: 主标签 '{primary_label}' 无效，但存在有效的次要标签。仅分配次要标签权重。\"\n            )\n            # 可以选择将所有权重 (1.0) 分配给次要标签\n            if len(valid_secondary) > 0:\n                sec_weight = 1.0 / len(valid_secondary)\n            else:\n                sec_weight = 0\n            for label in valid_secondary:\n                target[self.label_to_idx[label]] = sec_weight\n\n        return target\n\n\n# --- 模型类 (BirdCLEFModel) ---\nclass BirdCLEFModel(nn.Module):\n    \"\"\"\n    使用 timm 库创建的 EfficientNet B0 模型。\n    \"\"\"\n\n    def __init__(self, cfg):\n        super().__init__()\n        self.cfg = cfg\n        # 动态获取类别数\n        taxonomy_df = pd.read_csv(cfg.taxonomy_csv)\n        self.num_classes = len(taxonomy_df)\n        self.cfg.num_classes = self.num_classes  # 更新 CFG 中的 num_classes\n\n        self.backbone = timm.create_model(\n            cfg.model_name,\n            pretrained=cfg.pretrained,\n            in_chans=cfg.in_channels,\n            drop_rate=0.2,  # 可以调整\n            drop_path_rate=0.2,  # 可以调整\n        )\n\n        # 获取 backbone 输出特征维度并替换分类器\n        if hasattr(self.backbone, \"classifier\"):\n            backbone_out = self.backbone.classifier.in_features\n            self.backbone.classifier = nn.Identity()\n        elif hasattr(self.backbone, \"fc\"):\n            backbone_out = self.backbone.fc.in_features\n            self.backbone.fc = nn.Identity()\n        else:\n            # 尝试通用的 get_classifier() 方法\n            try:\n                backbone_out = self.backbone.get_classifier().in_features\n                self.backbone.reset_classifier(0, \"\")  # 移除分类器\n            except AttributeError:\n                # 如果以上都不行，可能需要针对特定模型结构进行调整\n                # 或者检查 timm 的文档\n                raise ValueError(\n                    f\"无法自动确定模型 {cfg.model_name} 的输出特征维度或移除分类器。\"\n                )\n\n        self.pooling = nn.AdaptiveAvgPool2d(1)  # 全局平均池化\n        self.classifier = nn.Linear(backbone_out, self.num_classes)  # 最终分类层\n\n        # Mixup 相关 (如果启用)\n        self.mixup_enabled = hasattr(cfg, \"mixup_alpha\") and cfg.mixup_alpha > 0\n        if self.mixup_enabled:\n            self.mixup_alpha = cfg.mixup_alpha\n\n    def forward(self, x, targets=None):\n        \"\"\"模型前向传播\"\"\"\n        # Mixup 处理 (如果启用且在训练模式)\n        if self.training and self.mixup_enabled and targets is not None:\n            mixed_x, targets_a, targets_b, lam = self.mixup_data(x, targets)\n            x = mixed_x\n        else:\n            targets_a, targets_b, lam = None, None, None  # 确保定义了这些变量\n\n        # === 修改：显式调用 forward_features ===\n        features = self.backbone.forward_features(x)\n        # 现在 features 的形状应该是 (B, C, H, W), 例如 (B, 1280, 8, 8)\n\n        # 如果 backbone 输出是字典 (某些 timm 模型会这样，虽然 forward_features 通常不会)\n        if isinstance(features, dict):\n            features = features[\"features\"]  # 或其他合适的 key\n\n        # 应用池化和展平\n        pooled_features = self.pooling(features).flatten(1)  # 形状变为 (B, C)\n        logits = self.classifier(\n            pooled_features\n        )  # 通过分类器得到 logits (C 应该等于 1280)\n\n        # 如果启用了 Mixup，计算混合损失\n        if self.training and self.mixup_enabled and targets is not None:\n            # 确保 mixup_criterion 定义在类中\n            loss = self.mixup_criterion(\n                F.binary_cross_entropy_with_logits,  # 使用 nn.functional\n                logits,\n                targets_a,\n                targets_b,\n                lam,\n            )\n            return logits, loss\n\n        # 如果没有 Mixup 或在评估模式，只返回 logits\n        return logits\n\n    # --- Mixup 辅助函数 (如果启用 Mixup) ---\n    def mixup_data(self, x, targets):\n        \"\"\"应用 Mixup\"\"\"\n        if self.mixup_alpha > 0:\n            lam = np.random.beta(self.mixup_alpha, self.mixup_alpha)\n        else:\n            lam = 1.0\n\n        batch_size = x.size(0)\n        indices = torch.randperm(batch_size).to(x.device)\n\n        mixed_x = lam * x + (1 - lam) * x[indices]\n        targets_a, targets_b = targets, targets[indices]\n        return mixed_x, targets_a, targets_b, lam\n\n    def mixup_criterion(self, criterion, pred, y_a, y_b, lam):\n        \"\"\"计算 Mixup 损失\"\"\"\n        return lam * criterion(pred, y_a, reduction=\"mean\") + (1 - lam) * criterion(\n            pred, y_b, reduction=\"mean\"\n        )\n\n\n# --- 训练与验证循环 ---\ndef train_one_epoch(\n    model, loader, optimizer, criterion, device, scheduler=None, cfg=None\n):\n    \"\"\"训练一个 epoch\"\"\"\n    model.train()\n    losses = []\n    all_targets = []\n    all_outputs = []\n\n    pbar = tqdm(enumerate(loader), total=len(loader), desc=\"训练中\")\n    for step, batch in pbar:\n        inputs = batch[\"melspec\"].to(device)\n        targets = batch[\"target\"].to(device)\n\n        optimizer.zero_grad()\n\n        # 处理 Mixup (如果模型返回 logits 和 loss)\n        if hasattr(cfg, \"mixup_alpha\") and cfg.mixup_alpha > 0:\n            logits, loss = model(inputs, targets)\n        else:\n            logits = model(inputs)\n            loss = criterion(logits, targets)\n\n        loss.backward()\n        optimizer.step()\n\n        # 如果使用 OneCycleLR，则每个 step 更新 scheduler\n        if scheduler is not None and cfg.scheduler == \"OneCycleLR\":\n            scheduler.step()\n\n        losses.append(loss.item())\n        all_outputs.append(logits.detach().cpu().numpy())\n        all_targets.append(targets.detach().cpu().numpy())\n\n        pbar.set_postfix(\n            {\n                \"loss\": np.mean(losses[-10:]) if losses else 0,\n                \"lr\": optimizer.param_groups[0][\"lr\"],\n            }\n        )\n\n    # 计算 epoch 的平均损失和 AUC\n    avg_loss = np.mean(losses)\n    all_outputs = np.concatenate(all_outputs)\n    all_targets = np.concatenate(all_targets)\n    auc = calculate_auc(all_targets, all_outputs)  # 现在调用补全的函数\n    return avg_loss, auc  # 返回平均损失和 AUC\n\n\ndef validate(model, loader, criterion, device):\n    \"\"\"验证模型在一个 epoch 上的表现\"\"\"\n    model.eval()\n    losses = []\n    all_targets = []\n    all_outputs = []\n\n    with torch.no_grad():\n        pbar = tqdm(enumerate(loader), total=len(loader), desc=\"验证中\")\n        for step, batch in pbar:\n            inputs = batch[\"melspec\"].to(device)\n            targets = batch[\"target\"].to(device)\n\n            logits = model(inputs)\n            loss = criterion(logits, targets)\n\n            losses.append(loss.item())\n            all_outputs.append(logits.detach().cpu().numpy())\n            all_targets.append(targets.detach().cpu().numpy())\n\n            pbar.set_postfix({\"loss\": np.mean(losses[-10:]) if losses else 0})\n\n    # 计算验证集的平均损失和 AUC\n    avg_loss = np.mean(losses)\n    all_outputs = np.concatenate(all_outputs)\n    all_targets = np.concatenate(all_targets)\n    auc = calculate_auc(all_targets, all_outputs)  # 现在调用补全的函数\n    return avg_loss, auc  # 返回平均损失和 AUC\n\n\n# --- 主训练函数 ---\ndef run_training(cfg):\n    \"\"\"执行完整的训练流程 (包括交叉验证)\"\"\"\n    set_seed(cfg.seed)\n\n    # --- 优化：只在开始时扫描文件和构建索引一次 ---\n    print(\"开始扫描 Batch 文件并构建全局索引...\")\n    batch_files = sorted(\n        glob.glob(os.path.join(cfg.INPUT_SPEC_DIR, \"precomputed_melspec_batch*.npy\"))\n    )\n    if not batch_files:\n        raise FileNotFoundError(f\"在 {cfg.INPUT_SPEC_DIR} 未找到任何 batch 文件！\")\n    print(f\"找到 {len(batch_files)} 个 batch 文件。\")\n\n    # 构建全局索引\n    global_sample_index = build_sample_index(batch_files)\n    all_keys = list(global_sample_index.keys())  # 从索引获取 keys\n\n    if not all_keys:\n        raise ValueError(\"未能从任何 batch 文件中加载 keys 或构建索引！\")\n    print(f\"全局索引构建完成，包含 {len(global_sample_index)} 个唯一样本 key。\")\n\n    # --- !!! 关键步骤：解析 key 以获取原始文件名 !!! ---\n    # (这部分逻辑保持不变，但现在基于 all_keys)\n    def parse_key_to_filename_example1(key):\n        # ... (你的解析逻辑) ...\n        parts = key.split(\"_chunk_\")[0]\n        # 检查拆分结果是否符合预期\n        if \"-\" not in parts:\n            # print(f\"警告: Key '{key}' 不符合 'dirname-basename_chunk_' 格式，跳过解析。\")\n            return None  # 或引发错误，或返回特殊值\n        dirname, basename = parts.split(\"-\", 1)\n        return f\"{dirname}/{basename}.ogg\"\n\n    parsed_data = []\n    for key in tqdm(all_keys, desc=\"解析 Keys\"):\n        try:\n            original_filename = parse_key_to_filename_example1(key)\n            if original_filename:  # 确保解析成功\n                parsed_data.append(\n                    {\"samplename\": key, \"original_filename_parsed\": original_filename}\n                )\n        except Exception as e:\n            print(f\"解析 key '{key}' 时出错: {e}\")\n\n    if not parsed_data:\n        raise ValueError(\"解析 key 失败，无法创建 key DataFrame。请检查解析逻辑！\")\n\n    key_df = pd.DataFrame(parsed_data)\n    # ... (加载 original_train_df, 合并 df 的逻辑保持不变) ...\n    original_train_df = pd.read_csv(cfg.train_csv)\n    df = pd.merge(\n        key_df,\n        original_train_df,\n        left_on=\"original_filename_parsed\",\n        right_on=\"filename\",\n        how=\"left\",\n    )\n    # ... (检查 missing_labels 的逻辑保持不变) ...\n    missing_labels = df[\"primary_label\"].isnull().sum()\n    if missing_labels > 0:\n        print(f\"警告: 合并后发现 {missing_labels} 行缺少 primary_label。\")\n        # 可选：移除或填充缺失标签的行\n        # df = df.dropna(subset=['primary_label']).reset_index(drop=True)\n        # df['primary_label'] = df['primary_label'].fillna('unknown_species') # 或其他填充策略\n\n    # 确保 primary_label 列存在且可用于分层抽样\n    if \"primary_label\" not in df.columns or df[\"primary_label\"].isnull().any():\n        raise ValueError(\n            \"错误: 合并后 DataFrame 缺少 'primary_label' 或存在空值，无法进行分层 K 折拆分。请检查数据或合并逻辑。\"\n        )\n\n    # --- 现在 df 包含了正确的 'samplename' 和对应的标签 ---\n\n    if cfg.debug:\n        cfg.update_debug_settings()\n        print(\"Debug 模式已启用。\")\n\n    # 交叉验证设置\n    # --- 修正：确保标签列中至少有 n_splits 个不同的类，或者处理无法分层的情况 ---\n    if df[\"primary_label\"].nunique() < cfg.n_fold:\n        print(\n            f\"警告: 数据集中不同主标签的数量 ({df['primary_label'].nunique()}) 少于指定的折数 ({cfg.n_fold})。将使用非分层 KFold。\"\n        )\n        from sklearn.model_selection import KFold\n\n        kf = KFold(n_splits=cfg.n_fold, shuffle=True, random_state=cfg.seed)\n        split_iterator = kf.split(df)  # 非分层拆分\n    else:\n        skf = StratifiedKFold(n_splits=cfg.n_fold, shuffle=True, random_state=cfg.seed)\n        split_iterator = skf.split(df, df[\"primary_label\"])  # 分层拆分\n\n    oof_auc_scores = []\n\n    # --- 修改：使用 split_iterator ---\n    for fold, (train_idx, val_idx) in enumerate(split_iterator):\n        if fold not in cfg.selected_folds:\n            continue\n\n        print(f'\\n{\"=\"*30} Fold {fold} {\"=\"*30}')\n        train_df = df.iloc[train_idx].reset_index(drop=True)\n        val_df = df.iloc[val_idx].reset_index(drop=True)\n\n        # 创建 Dataset 和 DataLoader (传递索引)\n        print(\"创建训练数据集 (使用全局索引)...\")\n        train_dataset = BirdCLEFDatasetFromBatches(\n            train_df,\n            cfg,\n            mode=\"train\",\n            batch_files=batch_files,\n            sample_index=global_sample_index,\n        )\n        print(\"创建验证数据集 (使用全局索引)...\")\n        val_dataset = BirdCLEFDatasetFromBatches(\n            val_df,\n            cfg,\n            mode=\"valid\",\n            batch_files=batch_files,\n            sample_index=global_sample_index,\n        )\n\n        # ... (检查数据集是否为空，创建 DataLoader, 初始化模型等逻辑保持不变) ...\n        if len(train_dataset) == 0 or len(val_dataset) == 0:\n            print(f\"错误: Fold {fold} 的训练或验证数据集在过滤后为空。跳过此 Fold。\")\n            continue\n\n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=cfg.batch_size,\n            shuffle=True,\n            num_workers=cfg.num_workers,\n            pin_memory=True,\n            drop_last=True,\n        )\n        val_loader = DataLoader(\n            val_dataset,\n            batch_size=cfg.batch_size * 2,\n            shuffle=False,  # 验证时 batch 可以大一些\n            num_workers=cfg.num_workers,\n            pin_memory=True,\n        )\n\n        # 初始化模型、优化器、损失函数、调度器\n        model = BirdCLEFModel(cfg).to(cfg.device)\n        optimizer = get_optimizer(model, cfg)\n        criterion = get_criterion(cfg)\n        scheduler = get_scheduler(optimizer, cfg)  # scheduler 可能为 None\n\n        # --- 修改：添加检查点加载逻辑 ---\n        start_epoch = 0\n        best_val_auc = 0.0\n        checkpoint_path = os.path.join(\n            cfg.OUTPUT_DIR,\n            cfg.checkpoint_filename.format(model_name=cfg.model_name, fold=fold),\n        )\n\n        if cfg.resume_training and os.path.exists(checkpoint_path):\n            print(f\"发现检查点: {checkpoint_path}，尝试恢复训练...\")\n            try:\n                checkpoint = torch.load(checkpoint_path, map_location=cfg.device)\n                model.load_state_dict(checkpoint[\"model_state_dict\"])\n                optimizer.load_state_dict(checkpoint[\"optimizer_state_dict\"])\n                # 加载 scheduler 状态时需要检查是否存在以及类型是否匹配\n                if (\n                    scheduler  # 确保 scheduler 被初始化了\n                    and \"scheduler_state_dict\" in checkpoint\n                    and checkpoint[\"scheduler_state_dict\"]  # 确保检查点里存了这个状态\n                ):\n                    try:\n                        scheduler.load_state_dict(checkpoint[\"scheduler_state_dict\"])\n                        print(\"成功加载 Scheduler 状态\")\n                    except Exception as scheduler_load_error:\n                        print(\n                            f\"警告：加载 Scheduler 状态时出错 (可能是类型不匹配或 scheduler 配置已更改): {scheduler_load_error}。将使用新的 Scheduler 状态。\"\n                        )\n                        # 可选：根据需要重置 scheduler 或使用新的\n\n                start_epoch = checkpoint[\"epoch\"] + 1\n                # 确保从检查点恢复 best_val_auc，即使 scheduler 加载失败\n                best_val_auc = checkpoint.get(\n                    \"best_val_auc\", 0.0\n                )  # 使用 .get 以兼容旧的检查点\n                print(\n                    f\"恢复成功，将从 Epoch {start_epoch} 开始训练。之前的最佳 AUC: {best_val_auc:.4f}\"\n                )\n                # 清理 checkpoint 占用的内存\n                del checkpoint\n                gc.collect()\n                if cfg.device == \"cuda\":\n                    torch.cuda.empty_cache()\n            except Exception as e:\n                print(f\"加载检查点失败: {e}。将从头开始训练。\")\n                start_epoch = 0\n                best_val_auc = 0.0\n        else:\n            if cfg.resume_training:\n                print(f\"未找到检查点 {checkpoint_path} 或检查点无效。\")\n            print(f\"将从 Epoch 0 开始训练。\")\n            start_epoch = 0\n            best_val_auc = 0.0\n        # --- 结束修改 ---\n\n        best_epoch = start_epoch - 1  # 初始化 best_epoch，如果从0开始则是-1\n\n        # --- 修改：调整 Epoch 循环范围 ---\n        for epoch in range(start_epoch, cfg.epochs):\n            print(f\"\\nEpoch {epoch+1}/{cfg.epochs}\")  # 打印时仍用 epoch+1\n            epoch_start_time = time.time()\n\n            # 训练和验证\n            train_loss, train_auc = train_one_epoch(\n                model, train_loader, optimizer, criterion, cfg.device, scheduler, cfg\n            )\n            val_loss, val_auc = validate(model, val_loader, criterion, cfg.device)\n\n            # --- 补充：更新学习率 ---\n            if scheduler is not None:\n                if isinstance(scheduler, lr_scheduler.ReduceLROnPlateau):\n                    scheduler.step(val_auc)  # 基于验证 AUC 更新\n                elif (\n                    cfg.scheduler != \"OneCycleLR\"\n                ):  # OneCycleLR 在 train_one_epoch 中更新\n                    scheduler.step()\n\n            epoch_time = time.time() - epoch_start_time\n            print(f\"耗时: {epoch_time:.2f}s\")\n            print(f\"训练损失: {train_loss:.4f}, 训练 AUC: {train_auc:.4f}\")\n            print(f\"验证损失: {val_loss:.4f}, 验证 AUC: {val_auc:.4f}\")\n\n            # --- 修改：保存逻辑 ---\n            # 1. 在每个 epoch 结束后都保存最新检查点 (包含所有状态)\n            latest_checkpoint_path = os.path.join(\n                cfg.OUTPUT_DIR,\n                cfg.checkpoint_filename.format(model_name=cfg.model_name, fold=fold),\n            )\n            save_dict = {\n                \"epoch\": epoch,\n                \"model_state_dict\": model.state_dict(),\n                \"optimizer_state_dict\": optimizer.state_dict(),\n                \"best_val_auc\": best_val_auc,  # 保存当前的最佳 AUC，以便恢复时知道历史最佳\n                \"current_val_auc\": val_auc,  # 保存当前的 AUC 便于查看\n            }\n            # 只有当 scheduler 存在时才保存其状态\n            if scheduler is not None:\n                save_dict[\"scheduler_state_dict\"] = scheduler.state_dict()\n\n            torch.save(save_dict, latest_checkpoint_path)\n            # print(f\"已保存最新检查点到: {latest_checkpoint_path}\") # 可以取消注释以确认保存\n\n            # 2. 如果当前是最佳模型，额外保存一个 _best.pth (只含模型权重)\n            if val_auc > best_val_auc:\n                best_val_auc = val_auc\n                best_epoch = epoch + 1  # 记录最佳 epoch (从 1 开始计数)\n                print(\n                    f\"*** Fold {fold} 找到新的最佳 AUC: {best_val_auc:.4f} at epoch {best_epoch} ***\"\n                )\n                best_model_path = os.path.join(\n                    cfg.OUTPUT_DIR, f\"{cfg.model_name}_fold{fold}_best.pth\"\n                )\n                # 只保存模型权重，方便后续推理或集成\n                torch.save(model.state_dict(), best_model_path)\n                print(f\"已保存最佳模型权重到: {best_model_path}\")\n\n            # 移除旧的 _last.pth 保存逻辑 (如果存在)\n            # --- 逻辑已包含在 best_val_auc 更新中 ---\n\n        # --- 结束 Epoch 循环 ---\n\n        print(\n            f\"\\nFold {fold} 训练完成。最佳验证 AUC: {best_val_auc:.4f} (Epoch {best_epoch if best_epoch > 0 else 'N/A'})\"\n        )\n        oof_auc_scores.append(best_val_auc)  # 记录该 fold 的最佳分数\n\n        # --- 补充：清理内存 ---\n        print(f\"Fold {fold} 结束，清理内存...\")\n        del (\n            model,\n            train_dataset,\n            val_dataset,\n            train_loader,\n            val_loader,\n            optimizer,\n            criterion,\n        )\n        if scheduler is not None:\n            del scheduler\n        gc.collect()  # 强制垃圾回收\n        if cfg.device == \"cuda\":\n            torch.cuda.empty_cache()  # 清空未使用的 CUDA 缓存\n\n    # --- 结束 Fold 循环 ---\n\n    # --- 补充：打印 OOF AUC 结果 ---\n    if oof_auc_scores:  # 确保至少完成了一个 fold\n        mean_oof_auc = np.mean(oof_auc_scores)\n        print(f'\\n{\"=\"*30} 训练结束 {\"=\"*30}')\n        print(\n            f\"所有选定 Folds ({cfg.selected_folds}) 的 OOF AUC 分数: {oof_auc_scores}\"\n        )\n        print(f\"平均 OOF AUC: {mean_oof_auc:.4f}\")\n    else:\n        print(\"\\n没有 Fold 被训练或完成，无法计算 OOF AUC。\")\n\n\n# --- Helper Functions (补全实现) ---\ndef get_optimizer(model, cfg):\n    \"\"\"根据配置返回优化器实例\"\"\"\n    if cfg.optimizer == \"AdamW\":\n        optimizer = optim.AdamW(\n            model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay\n        )\n    elif cfg.optimizer == \"Adam\":\n        optimizer = optim.Adam(\n            model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay\n        )\n    # 可以根据需要添加 SGD 等其他优化器选项\n    # elif cfg.optimizer == 'SGD':\n    #     optimizer = optim.SGD(\n    #         model.parameters(),\n    #         lr=cfg.lr,\n    #         momentum=0.9,\n    #         weight_decay=cfg.weight_decay\n    #     )\n    else:\n        raise ValueError(f\"不支持的优化器: {cfg.optimizer}\")\n    return optimizer\n\n\ndef get_scheduler(optimizer, cfg):\n    \"\"\"根据配置返回学习率调度器实例\"\"\"\n    if cfg.scheduler == \"CosineAnnealingLR\":\n        scheduler = lr_scheduler.CosineAnnealingLR(\n            optimizer, T_max=cfg.T_max, eta_min=cfg.min_lr\n        )\n    elif cfg.scheduler == \"ReduceLROnPlateau\":\n        scheduler = lr_scheduler.ReduceLROnPlateau(\n            optimizer,\n            mode=\"max\",  # 通常监控验证集 AUC，所以用 max\n            factor=0.5,  # 学习率衰减因子\n            patience=2,  # 多少个 epoch AUC 没有提升则降低学习率\n            min_lr=cfg.min_lr,\n            verbose=True,\n        )\n    # 可以根据需要添加 StepLR 等其他调度器选项\n    # elif cfg.scheduler == 'StepLR':\n    #     scheduler = lr_scheduler.StepLR(\n    #         optimizer,\n    #         step_size=cfg.epochs // 3,\n    #         gamma=0.5\n    #     )\n    elif cfg.scheduler == \"OneCycleLR\":\n        # OneCycleLR 需要在每个 step 更新，特殊处理\n        # 初始化在主训练循环中进行\n        scheduler = None\n    elif cfg.scheduler is None:\n        scheduler = None\n    else:\n        raise ValueError(f\"不支持的调度器: {cfg.scheduler}\")\n    return scheduler\n\n\ndef get_criterion(cfg):\n    \"\"\"根据配置返回损失函数实例\"\"\"\n    if cfg.criterion == \"BCEWithLogitsLoss\":\n        criterion = nn.BCEWithLogitsLoss()\n    # 可以根据需要添加 CrossEntropyLoss 等其他损失函数选项\n    # elif cfg.criterion == 'CrossEntropyLoss':\n    #     criterion = nn.CrossEntropyLoss()\n    else:\n        raise ValueError(f\"不支持的损失函数: {cfg.criterion}\")\n    return criterion\n\n\n# --- 主程序入口 (保持不变) ---\nif __name__ == \"__main__\":\n    print(\"初始化配置...\")\n    cfg = CFG()\n\n    # 创建输出目录\n    os.makedirs(cfg.OUTPUT_DIR, exist_ok=True)\n\n    print(\"开始训练流程...\")\n    run_training(cfg)\n    print(\"训练流程结束。\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-17T07:35:37.083182Z","iopub.execute_input":"2025-04-17T07:35:37.083509Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}