{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":11351027,"sourceType":"datasetVersion","datasetId":7102743},{"sourceId":11433559,"sourceType":"datasetVersion","datasetId":7161252}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport random\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset\nimport librosa\nimport cv2\n\nimport pickle\nimport random\n\n\nclass Config:\n    def __init__(self):\n        # 其他配置参数...\n        self.FS = 32000  # 采样率\n        self.taxonomy_csv = \"/kaggle/input/birdclef-2025/taxonomy.csv\"\n        self.train_datadir = \"/kaggle/input/birdclef-2025/train_audio\"\n        self.voice_data_path = (\n            \"/kaggle/input/humen-voice/train_voice_data.pkl\"  # 人声数据文件路径\n        )\n\n        # 预计算相关参数\n        self.INPUT_SPEC_DIR = (\n            \"/kaggle/input/zwy-bird\"  # 预计算频谱图输入目录\n        )\n        self.OUTPUT_DIR = \"/kaggle/working/\"  # 新batch的输出目录\n        self.BATCH_SIZE = 5000  # 每个batch包含的频谱图数量\n        # 预计算频谱图路径 - 只从input文件夹读取\n        self.PRECOMPUTED_SPEC_PATH = (\n            \"/kaggle/input/precomputed_melspec/precomputed_melspec.npy\"\n        )\n        # 保存的路径 - 只能保存到working目录\n        self.SAVE_SPEC_PATH = \"/kaggle/working/precomputed_melspec.npy\"\n        self.LOAD_DATA = False  # 设置为False以触发预计算\n\n        # 音频处理参数\n        self.TARGET_SHAPE = (256, 256)  # 目标频谱图形状\n        self.N_FFT = 1024\n        self.HOP_LENGTH = 512\n        self.N_MELS = 128\n        self.FMIN = 50\n        self.FMAX = 14000\n\n        # 人声数据相关配置\n        self.voice_data = None  # 初始化为None\n        self.use_10sec_chunks = True\n\n        # 调试模式\n        self.debug = False\n        self.seed = 42\n\n        # 添加数据增强相关参数\n        self.aug_prob = 0.5  # 应用数据增强的概率，与 EfficientNet 代码中默认值相同\n\n\n# 创建配置对象\ncfg = Config()\nwith open(cfg.voice_data_path, \"rb\") as f:\n    voice_timestamps = pickle.load(f)\n    print(\"成功加载人声时间戳数据\")\n# 将加载的人声时间戳数据赋值给cfg.voice_data\ncfg.voice_data = voice_timestamps\n\n\nclass BirdCLEFDataset(Dataset):\n    def __init__(self, df, cfg, spectrograms=None, mode=\"train\"):\n        \"\"\"\n        增强版鸟类声音数据集类\n\n        Args:\n            df: 包含音频文件信息的DataFrame\n            cfg: 配置对象\n            spectrograms: 预计算的频谱图字典 (可选)\n            mode: 'train', 'valid', 或 'test'\n        \"\"\"\n        self.df = df.copy()  # 创建副本以避免修改原始df\n        self.cfg = cfg\n        self.mode = mode\n        self.spectrograms = spectrograms\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 \"filepath\" not in self.df.columns:\n            self.df[\"filepath\"] = self.df.filename.apply(\n                lambda x: os.path.join(self.cfg.train_datadir, x)\n            )\n\n        # 为每个样本创建唯一标识符\n        if \"samplename\" not in self.df.columns:\n            self.df[\"samplename\"] = self.df.filename.map(\n                lambda x: x.split(\"/\")[0] + \"-\" + x.split(\"/\")[-1].split(\".\")[0]\n            )\n\n        # 处理人声过滤（如果有人声检测结果）\n        if hasattr(cfg, \"voice_data\") and cfg.voice_data is not None:\n            print(\"应用人声过滤...\")\n            self.voice_data = cfg.voice_data\n        else:\n            self.voice_data = None\n\n        # 准备10秒段落（每个由两个5秒chunk组成）\n        if (\n            mode == \"train\"\n            and hasattr(cfg, \"use_10sec_chunks\")\n            and cfg.use_10sec_chunks\n        ):\n            print(\"准备10秒训练段...\")\n            self._prepare_10sec_chunks()\n\n        # 处理频谱图 - 在限制数据集大小之前先处理频谱图\n        if self.spectrograms is None:  # 如果没有传入预计算的频谱图\n            if cfg.LOAD_DATA and os.path.exists(cfg.PRECOMPUTED_SPEC_PATH):\n                # 尝试从input目录加载预计算频谱图\n                try:\n                    print(f\"尝试加载预计算频谱图从 {cfg.PRECOMPUTED_SPEC_PATH}...\")\n                    self.spectrograms = np.load(\n                        cfg.PRECOMPUTED_SPEC_PATH, allow_pickle=True\n                    ).item()\n                    print(f\"成功加载 {len(self.spectrograms)} 个预计算频谱图\")\n                except Exception as e:\n                    print(f\"加载预计算频谱图时出错: {e}\")\n                    self.spectrograms = None\n            else:\n                # 如果PRECOMPUTED_SPEC_PATH不存在或LOAD_DATA为False\n                print(f\"预计算频谱图文件不存在或未启用加载，将计算并保存频谱图...\")\n                # 注意：这里使用完整的df进行预计算\n                self.spectrograms = self._precompute_spectrograms(self.df)\n\n        # Debug模式下减少数据量 - 移到频谱图处理之后\n        if cfg.debug:\n            print(\n                f\"Debug模式：从 {len(self.df)} 个样本中抽样 {min(1000, len(self.df))} 个\"\n            )\n            self.df = self.df.sample(\n                min(1000, len(self.df)), random_state=cfg.seed\n            ).reset_index(drop=True)\n\n        # # 报告与预计算频谱图的匹配情况\n        # if self.spectrograms:\n        #     sample_names = set(self.df[\"samplename\"])\n        #     loaded_names = set(self.spectrograms.keys())\n\n        #     # 显示两个集合的交集大小\n        #     intersection = sample_names.intersection(loaded_names)\n        #     print(f\"数据集样本名称与加载的频谱图有 {len(intersection)} 个交集\")\n\n        #     # 如果交集为空，打印一些样本名称进行对比\n        #     if len(intersection) == 0:\n        #         print(\"数据集中的前5个样本名称:\", list(sample_names)[:5])\n        #         print(\"加载的频谱图中的前5个键:\", list(loaded_names)[:5])\n\n        #     found_samples = sum(1 for name in sample_names if name in self.spectrograms)\n        #     print(f\"找到 {found_samples} 个匹配的频谱图，共 {len(self.df)} 个样本\")\n\n    def _prepare_10sec_chunks(self):\n        \"\"\"准备10秒的音频段（2个相邻的5秒chunk）\"\"\"\n        new_rows = []\n        grouped = self.df.groupby(\"filename\")\n\n        for filename, group in grouped:\n            # 获取音频长度\n            try:\n                audio_path = group.iloc[0][\"filepath\"]\n                audio_length = librosa.get_duration(path=audio_path)\n\n                # 创建10秒的滑动窗口（每次移动5秒）\n                for start_time in range(0, int(audio_length) - 10 + 1, 5):\n                    chunk_row = group.iloc[0].copy()\n                    chunk_row[\"chunk_start\"] = start_time\n                    chunk_row[\"chunk_end\"] = start_time + 10\n                    new_rows.append(chunk_row)\n\n            except Exception as e:\n                print(f\"处理文件 {filename} 时出错: {e}\")\n\n        if new_rows:\n            self.df = pd.DataFrame(new_rows)\n            print(f\"创建了 {len(self.df)} 个10秒段\")\n\n    def _process_audio_chunk(self, filepath, start_time=0, duration=5):\n        \"\"\"处理单个音频片段，返回梅尔频谱图\"\"\"\n        try:\n            # 使用指定的起始位置和持续时间加载音频\n            audio_data, sr = librosa.load(\n                filepath, sr=self.cfg.FS, offset=start_time, duration=duration\n            )\n\n            # 如果音频太短，循环填充\n            if len(audio_data) < duration * self.cfg.FS:\n                n_repeat = int(np.ceil(duration * self.cfg.FS / len(audio_data)))\n                audio_data = np.tile(audio_data, n_repeat)[\n                    : int(duration * self.cfg.FS)\n                ]\n\n            # 过滤人声（如果有人声数据）\n            if self.voice_data and filepath in self.voice_data:\n                audio_data = self._filter_voice(\n                    audio_data, filepath, start_time, duration\n                )\n\n            # 生成梅尔频谱图\n            mel_spec = self._audio_to_melspec(audio_data)\n\n            # 调整大小以匹配目标形状\n            if mel_spec.shape != self.cfg.TARGET_SHAPE:\n                mel_spec = cv2.resize(\n                    mel_spec, self.cfg.TARGET_SHAPE, interpolation=cv2.INTER_LINEAR\n                )\n\n            return mel_spec.astype(np.float32)\n\n        except Exception as e:\n            print(\n                f\"处理 {filepath} 的 {start_time}-{start_time+duration}秒片段时出错: {e}\"\n            )\n            return None\n\n    def _filter_voice(self, audio_data, filepath, start_time, duration):\n        \"\"\"移除音频中的人声部分\"\"\"\n        # 获取当前时间片段内的人声时间戳\n        voice_timestamps = self.voice_data.get(filepath, [])\n        relevant_voices = [\n            v\n            for v in voice_timestamps\n            if (v[\"start\"] < start_time + duration and v[\"end\"] > start_time)\n        ]\n\n        if not relevant_voices:\n            return audio_data\n\n        # 创建一个掩码，0表示保留，1表示人声\n        mask = np.zeros_like(audio_data)\n        sr = self.cfg.FS\n\n        for v in relevant_voices:\n            # 转换为相对于当前片段的样本索引\n            rel_start = max(0, int((v[\"start\"] - start_time) * sr))\n            rel_end = min(len(audio_data), int((v[\"end\"] - start_time) * sr))\n\n            if rel_start < rel_end:\n                # 创建淡入淡出效果（避免突变）\n                fade_len = min(int(0.1 * sr), (rel_end - rel_start) // 2)\n                if fade_len > 0:\n                    # 淡入\n                    mask[rel_start : rel_start + fade_len] = np.linspace(0, 1, fade_len)\n                    # 中间部分\n                    mask[rel_start + fade_len : rel_end - fade_len] = 1\n                    # 淡出\n                    mask[rel_end - fade_len : rel_end] = np.linspace(1, 0, fade_len)\n                else:\n                    mask[rel_start:rel_end] = 1\n\n        # 应用掩码（将人声部分设为很小的值或用环境噪声替换）\n        filtered_audio = audio_data * (1 - mask)\n        return filtered_audio\n\n    def _audio_to_melspec(self, audio_data):\n        \"\"\"将音频数据转换为梅尔频谱图\"\"\"\n        if np.isnan(audio_data).any():\n            mean_signal = np.nanmean(audio_data)\n            audio_data = np.nan_to_num(audio_data, nan=mean_signal)\n\n        mel_spec = librosa.feature.melspectrogram(\n            y=audio_data,\n            sr=self.cfg.FS,\n            n_fft=self.cfg.N_FFT,\n            hop_length=self.cfg.HOP_LENGTH,\n            n_mels=self.cfg.N_MELS,\n            fmin=self.cfg.FMIN,\n            fmax=self.cfg.FMAX,\n            power=2.0,\n        )\n\n        mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max)\n        mel_spec_norm = (mel_spec_db - mel_spec_db.min()) / (\n            mel_spec_db.max() - mel_spec_db.min() + 1e-8\n        )\n\n        return mel_spec_norm\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        spec = None\n\n        # 检查是否有10秒段信息\n        has_chunks = hasattr(row, \"chunk_start\") and hasattr(row, \"chunk_end\")\n\n        # 1. 尝试从预计算频谱图加载\n        if (\n            not has_chunks\n            and self.spectrograms\n            and row[\"samplename\"] in self.spectrograms\n        ):\n            spec = self.spectrograms[row[\"samplename\"]]\n\n        # 2. 处理10秒段\n        elif has_chunks:\n            spec = self._process_audio_chunk(\n                row[\"filepath\"], start_time=row[\"chunk_start\"], duration=10  # 10秒段\n            )\n\n        # 3. 处理常规5秒段\n        elif not self.cfg.LOAD_DATA:\n            if hasattr(row, \"start_time\") and hasattr(row, \"duration\"):\n                spec = self._process_audio_chunk(\n                    row[\"filepath\"],\n                    start_time=row[\"start_time\"],\n                    duration=row[\"duration\"],\n                )\n            else:\n                # 获取音频长度\n                audio_length = librosa.get_duration(path=row[\"filepath\"])\n                # 计算中心5秒的起始位置（如果音频长度>5秒）\n                if audio_length > 5:\n                    start_time = (audio_length - 5) / 2\n                else:\n                    start_time = 0\n                # 从中心提取5秒\n                spec = self._process_audio_chunk(row[\"filepath\"], start_time=start_time)\n\n        # 如果以上方法都失败，创建一个空的频谱图\n        if spec is None:\n            spec = np.zeros(self.cfg.TARGET_SHAPE, dtype=np.float32)\n            if self.mode == \"train\":  # 只在训练时显示警告\n                print(f\"警告: 未找到 {row['samplename']} 的频谱图，无法生成\")\n\n        # 转换为张量并添加通道维度\n        spec = torch.tensor(spec, dtype=torch.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(row[\"primary_label\"])\n\n        # 处理次要标签\n        if \"secondary_labels\" in row and row[\"secondary_labels\"] not in [\n            [\"\"],\n            None,\n            np.nan,\n        ]:\n            if isinstance(row[\"secondary_labels\"], str):\n                secondary_labels = eval(row[\"secondary_labels\"])\n            else:\n                secondary_labels = row[\"secondary_labels\"]\n\n            # 如果有次要标签，根据2024年第一名方案，主标签权重0.5，其余0.5平分给次要标签\n            if secondary_labels and len(secondary_labels) > 0:\n                valid_secondary = [\n                    l for l in secondary_labels if l in self.label_to_idx\n                ]\n                if valid_secondary:\n                    # 将主标签权重设为0.5\n                    main_idx = self.label_to_idx[row[\"primary_label\"]]\n                    target[main_idx] = 0.5\n\n                    # 剩余0.5平均分配给次要标签\n                    sec_weight = 0.5 / len(valid_secondary)\n                    for label in valid_secondary:\n                        target[self.label_to_idx[label]] = sec_weight\n\n        return {\n            \"melspec\": spec,\n            \"target\": torch.tensor(target, dtype=torch.float32),\n            \"filename\": row[\"filename\"],\n            \"samplename\": row[\"samplename\"] if \"samplename\" in row else \"\",\n        }\n\n    def apply_spec_augmentations(self, spec):\n        \"\"\"对频谱图应用增强\"\"\"\n        # 1. 时间掩码 (水平条纹) - XY masking中的Y方向\n        if random.random() < 0.5:\n            num_masks = random.randint(1, 3)\n            for _ in range(num_masks):\n                width = random.randint(5, 20)\n                start = random.randint(0, spec.shape[2] - width)\n                spec[0, :, start : start + width] = 0\n\n        # 2. 频率掩码 (垂直条纹) - XY masking中的X方向\n        if random.random() < 0.5:\n            num_masks = random.randint(1, 3)\n            for _ in range(num_masks):\n                height = random.randint(5, 20)\n                start = random.randint(0, spec.shape[1] - height)\n                spec[0, start : start + height, :] = 0\n\n        # 3. 随机亮度/对比度调整\n        if random.random() < 0.5:\n            gain = random.uniform(0.8, 1.2)\n            bias = random.uniform(-0.1, 0.1)\n            spec = spec * gain + bias\n            spec = torch.clamp(spec, 0, 1)\n\n        # 4. 水平cutmix (根据2024年第一名方案添加)\n        if random.random() < 0.3:  # 30%的概率应用\n            cut_width = random.randint(5, int(spec.shape[2] * 0.3))  # 最多切30%宽度\n            cut_start = random.randint(0, spec.shape[2] - cut_width)\n\n            # 将切出的部分水平镜像翻转\n            spec[0, :, cut_start : cut_start + cut_width] = torch.flip(\n                spec[0, :, cut_start : cut_start + cut_width], dims=[-1]\n            )\n\n        return spec\n\n    def encode_label(self, label):\n        \"\"\"将标签编码为one-hot向量\"\"\"\n        target = np.zeros(self.num_classes)\n        if label in self.label_to_idx:\n            target[self.label_to_idx[label]] = 1.0\n        return target\n\n    def _precompute_spectrograms(self, df):\n        \"\"\"内部方法：预计算并保存频谱图，支持断点续传\"\"\"\n        import time\n        from tqdm import tqdm\n        import glob  # 添加glob模块用于文件匹配\n\n        # 初始化存储所有频谱图的字典\n        all_spectrograms = {}\n        processed_samples = set()\n        errors = []\n\n        # 1. 首先加载已有的batch文件\n        input_dir = self.cfg.INPUT_SPEC_DIR\n        print(\"检查已有的预计算频谱图...\")\n        last_batch_num = 0\n\n        # 使用glob自动查找所有batch文件\n        batch_files = glob.glob(f\"{input_dir}/precomputed_melspec_batch*.npy\")\n        # 从文件名中提取batch号并排序\n        batch_numbers = sorted(\n            [int(f.split(\"batch\")[-1].split(\".\")[0]) for f in batch_files]\n        )\n\n        if not batch_numbers:\n            print(\"未找到任何预计算的频谱图文件\")\n        else:\n            print(f\"找到 {len(batch_numbers)} 个batch文件\")\n\n        # 按顺序加载每个batch\n        for i in batch_numbers:\n            batch_path = f\"{input_dir}/precomputed_melspec_batch{i}.npy\"\n            try:\n                batch_data = np.load(batch_path, allow_pickle=True).item()\n                # 提取基础样本名（去掉_chunk_部分）\n                base_samples = set(key.split(\"_chunk_\")[0] for key in batch_data.keys())\n                processed_samples.update(base_samples)\n                last_batch_num = i\n                print(f\"已加载batch {i}: 包含 {len(base_samples)} 个基础样本\")\n            except Exception as e:\n                print(f\"batch {i} 不完整或不存在: {e}\")\n                break\n\n        print(f\"已处理 {len(processed_samples)} 个样本\")\n        print(f\"最后完整的batch编号: {last_batch_num}\")\n\n        # 2. 筛选出未处理的样本\n        df[\"samplename\"] = df.filename.map(\n            lambda x: x.split(\"/\")[0] + \"-\" + x.split(\"/\")[-1].split(\".\")[0]\n        )\n        remaining_df = df[~df[\"samplename\"].isin(processed_samples)].reset_index(\n            drop=True\n        )\n        print(f\"剩余 {len(remaining_df)} 个样本待处理\")\n\n        if len(remaining_df) == 0:\n            print(\"所有样本都已处理完成！\")\n            # 加载并返回所有处理好的频谱图\n            final_spectrograms = {}\n            for i in range(1, last_batch_num + 1):\n                batch_path = f\"{input_dir}/precomputed_melspec_batch{i}.npy\"\n                batch_data = np.load(batch_path, allow_pickle=True).item()\n                final_spectrograms.update(batch_data)\n            return final_spectrograms\n\n        # 3. 继续处理剩余样本\n        output_dir = self.cfg.OUTPUT_DIR\n        output_base = \"precomputed_melspec\"\n        start_time = time.time()\n\n        # 从上一个batch号继续\n        batch_num = last_batch_num + 1\n        batch_size = 5000\n        processed_count = 0\n        current_spectrograms = {}\n\n        # 确保输出目录存在\n        os.makedirs(output_dir, exist_ok=True)\n\n        for i, row in tqdm(\n            remaining_df.iterrows(), total=len(remaining_df), desc=\"处理剩余音频文件\"\n        ):\n            try:\n                samplename = row[\"samplename\"]\n                filepath = row[\"filepath\"]\n\n                # 获取音频长度\n                audio_length = librosa.get_duration(path=filepath)\n\n                # 创建10秒的滑动窗口\n                for start_time in range(0, int(audio_length) - 10 + 1, 5):\n                    chunk_id = f\"{samplename}_chunk_{start_time}_{start_time+10}\"\n\n                    # 处理10秒音频段\n                    mel_spec = self._process_audio_chunk(filepath, start_time, 10)\n\n                    if mel_spec is not None:\n                        current_spectrograms[chunk_id] = mel_spec\n                        processed_count += 1\n\n                # 分批保存\n                if len(current_spectrograms) >= batch_size:\n                    batch_path = f\"{output_dir}/{output_base}_batch{batch_num}.npy\"\n                    print(\n                        f\"\\n保存批次 {batch_num}，包含 {len(current_spectrograms)} 个频谱图...\"\n                    )\n\n                    try:\n                        np.save(batch_path, current_spectrograms)\n                        print(f\"成功将批次 {batch_num} 保存到 {batch_path}\")\n                        all_spectrograms.update(current_spectrograms)\n                        current_spectrograms = {}\n                        batch_num += 1\n                    except Exception as e:\n                        print(f\"保存批次 {batch_num} 时出错: {e}\")\n                        import traceback\n\n                        traceback.print_exc()\n\n            except Exception as e:\n                print(f\"处理 {row.get('filepath', 'unknown')} 时出错: {e}\")\n                errors.append((row.get(\"filepath\", \"unknown\"), str(e)))\n\n        # 保存最后一批数据\n        if current_spectrograms:\n            batch_path = f\"{output_dir}/{output_base}_batch{batch_num}.npy\"\n            try:\n                np.save(batch_path, current_spectrograms)\n                print(\n                    f\"\\n保存最后批次 {batch_num}，包含 {len(current_spectrograms)} 个频谱图\"\n                )\n                all_spectrograms.update(current_spectrograms)\n            except Exception as e:\n                print(f\"保存最后批次时出错: {e}\")\n\n        end_time = time.time()\n        print(f\"\\n预处理完成，耗时 {end_time - start_time:.2f} 秒\")\n        print(f\"新处理了 {processed_count} 个频谱图，失败 {len(errors)} 个\")\n\n        if errors:\n            print(\"\\n前几个错误:\")\n            for filepath, error_msg in errors[:5]:\n                print(f\"  - {filepath}: {error_msg}\")\n\n        # 合并所有频谱图\n        final_spectrograms = {}\n        print(\"\\n加载并合并所有批次...\")\n\n        # 加载已有的batch\n        for i in range(1, last_batch_num + 1):\n            batch_path = f\"{input_dir}/precomputed_melspec_batch{i}.npy\"\n            try:\n                batch_data = np.load(batch_path, allow_pickle=True).item()\n                final_spectrograms.update(batch_data)\n            except Exception as e:\n                print(f\"加载已有batch {i} 时出错: {e}\")\n\n        # 加载新计算的batch\n        for i in range(last_batch_num + 1, batch_num + 1):\n            batch_path = f\"{output_dir}/{output_base}_batch{i}.npy\"\n            try:\n                batch_data = np.load(batch_path, allow_pickle=True).item()\n                final_spectrograms.update(batch_data)\n            except Exception as e:\n                print(f\"加载新batch {i} 时出错: {e}\")\n\n        print(f\"最终共有 {len(final_spectrograms)} 个频谱图\")\n        return final_spectrograms\n\n\ndef collate_fn(batch):\n    \"\"\"自定义整理函数，处理不同大小的频谱图\"\"\"\n    batch = [item for item in batch if item is not None]\n    if len(batch) == 0:\n        return {}\n\n    result = {key: [] for key in batch[0].keys()}\n\n    for item in batch:\n        for key, value in item.items():\n            result[key].append(value)\n\n    for key in result:\n        if key in [\"target\", \"melspec\"] and isinstance(result[key][0], torch.Tensor):\n            try:\n                result[key] = torch.stack(result[key])\n            except:\n                # 如果形状不一致，保持为列表\n                pass\n\n    return result\n\n\ndef main():\n    \"\"\"使用示例\"\"\"\n    import pandas as pd\n\n    # 创建配置\n    cfg = Config()\n\n    # 加载训练数据\n    train_df = pd.read_csv(\"/kaggle/input/birdclef-2025/train.csv\")\n\n    # 加载人声数据（如果需要）\n    if os.path.exists(cfg.voice_data_path):\n        try:\n            with open(cfg.voice_data_path, \"rb\") as f:\n                cfg.voice_data = pickle.load(f)\n                print(f\"成功加载人声时间戳数据，包含 {len(cfg.voice_data)} 个文件\")\n        except Exception as e:\n            print(f\"加载人声数据时出错: {e}\")\n\n    # 创建数据集 - 频谱图的加载或计算现在在BirdCLEFDataset的__init__方法中处理\n    train_dataset = BirdCLEFDataset(train_df, cfg, spectrograms=None, mode=\"train\")\n\n    # 示例：获取一个样本\n    sample = train_dataset[0]\n    print(f\"样本形状: {sample['melspec'].shape}\")\n    print(f\"标签形状: {sample['target'].shape}\")\n\n    return train_dataset.spectrograms, train_dataset\n\n\nif __name__ == \"__main__\":\n    spectrograms, dataset = main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T07:07:30.156319Z","iopub.execute_input":"2025-04-15T07:07:30.156803Z","iopub.status.idle":"2025-04-15T07:16:14.473402Z","shell.execute_reply.started":"2025-04-15T07:07:30.156777Z","shell.execute_reply":"2025-04-15T07:16:14.472364Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(spectrograms)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 在main函数的末尾添加以下代码\nprint(\"\\n最终dataset中的df信息:\")\nprint(f\"行数: {len(dataset.df)}\")  \nprint(f\"列名: {dataset.df.columns.tolist()}\")\nprint(\"\\n前几行数据:\")\ndataset.df.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 如果启用了10秒段处理，查看chunk相关列\nif 'chunk_start' in dataset.df.columns:\n    print(\"\\n10秒段信息示例:\")\n    chunk_info = dataset.df[['filename', 'chunk_start', 'chunk_end']].head(5)\n    print(chunk_info)\n    ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"在debug模式下，数据随机选了1000条","metadata":{}},{"cell_type":"code","source":"# 在main函数末尾添加以下代码验证预计算频谱图文件\ndef test_precomputed_file():\n    import numpy as np\n    import os\n    \n    precomputed_path = \"/kaggle/working/precomputed_melspec.npy\"\n    \n    if not os.path.exists(precomputed_path):\n        print(\"预计算频谱图文件不存在!\")\n        return\n    \n    try:\n        # 加载预计算频谱图\n        spectrograms = np.load(precomputed_path, allow_pickle=True).item()\n        \n        print(\"\\n预计算频谱图文件验证:\")\n        print(f\"包含 {len(spectrograms)} 个预计算频谱图\")\n        \n        # 检查频谱图的形状\n        sample_keys = list(spectrograms.keys())[:5]  # 取前5个样本\n        print(\"\\n前几个频谱图示例:\")\n        for key in sample_keys:\n            spec = spectrograms[key]\n            print(f\"样本 {key}: 形状 {spec.shape}, 类型 {spec.dtype}\")\n            \n        # 检查是否包含10秒段标识\n        has_chunks = any(\"_chunk_\" in key for key in spectrograms.keys())\n        if has_chunks:\n            chunk_keys = [key for key in spectrograms.keys() if \"_chunk_\" in key][:3]\n            print(\"\\n检测到10秒段频谱图，示例:\")\n            for key in chunk_keys:\n                print(f\"  - {key}\")\n        else:\n            print(\"\\n未检测到10秒段频谱图，可能是常规5秒段模式\")\n            \n        # 检查数值范围\n        sample_spec = spectrograms[sample_keys[0]]\n        print(f\"\\n数值范围检查 (应在0-1之间): 最小值 {sample_spec.min():.4f}, 最大值 {sample_spec.max():.4f}\")\n        \n        return spectrograms\n        \n    except Exception as e:\n        print(f\"验证预计算频谱图文件时出错: {e}\")\n        return None\n\n# 调用测试函数\nprecomputed_specs = test_precomputed_file()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}