{"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":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":11549194,"sourceType":"datasetVersion","datasetId":7242630},{"sourceId":355731,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":296690,"modelId":317294}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport math\nimport numpy as np\nimport pandas as pd\nimport librosa\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchaudio\nimport torch.nn.functional as F\nfrom torch.optim.lr_scheduler import LambdaLR\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchaudio.transforms import MelSpectrogram\nfrom torchvision import models\nfrom sklearn.model_selection import train_test_split\nfrom matplotlib import pyplot as plt\nfrom tqdm import tqdm, trange\nfrom torchmetrics.classification import BinaryAUROC, MulticlassAUROC","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-24T19:19:28.816676Z","iopub.execute_input":"2025-04-24T19:19:28.817413Z","iopub.status.idle":"2025-04-24T19:19:33.641280Z","shell.execute_reply.started":"2025-04-24T19:19:28.817383Z","shell.execute_reply":"2025-04-24T19:19:33.640664Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 配置参数","metadata":{}},{"cell_type":"code","source":"if not os.path.exists(\"/kaggle/working/birdclef-2025\"):\n    os.makedirs(\"/kaggle/working/birdclef-2025\")\nos.listdir('/kaggle/working')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T19:19:33.642519Z","iopub.execute_input":"2025-04-24T19:19:33.643250Z","iopub.status.idle":"2025-04-24T19:19:33.649131Z","shell.execute_reply.started":"2025-04-24T19:19:33.643231Z","shell.execute_reply":"2025-04-24T19:19:33.648493Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    seed = 42\n    img_size = [128, 384]\n    batch_size = 64\n    duration = 15  # 秒\n    sample_rate = 32000\n    audio_len = duration * sample_rate\n    dim = sample_rate * 5\n    nfft = 2028\n    window = 2048\n    hop_length = audio_len // (img_size[1] - 1)\n    fmin = 20\n    fmax = 16000\n    epochs = 300  # 减少epoch便于测试\n    # preset = 'efficientnetv2_b2_imagenet'\n    BASE_PATH_work = '/kaggle/working/birdclef-2025'\n    # model_weight = \"/kaggle/input/best_model_weights.pth/pytorch/default/1/best_model_weights.pth\"\n    # BASE_PATH_input = '/kaggle/input/birdclef-2025'\n    BASE_PATH_input = '/kaggle/input/birdclef-2025'\n    augment = True\n    class_names = sorted(os.listdir(f'{BASE_PATH_input}/train_audio/'))  # 确保路径正确\n    num_classes = len(class_names)\n    class_labels = list(range(num_classes))\n    label2name = dict(zip(class_labels, class_names))\n    name2label = {v: k for k, v in label2name.items()}\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T19:19:33.650091Z","iopub.execute_input":"2025-04-24T19:19:33.650359Z","iopub.status.idle":"2025-04-24T19:19:33.668382Z","shell.execute_reply.started":"2025-04-24T19:19:33.650336Z","shell.execute_reply":"2025-04-24T19:19:33.667704Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 数据增强模块","metadata":{}},{"cell_type":"code","source":"class MixUp(nn.Module):\n    def __init__(self, alpha=0.4):\n        super().__init__()\n        self.alpha = alpha\n\n    def forward(self, images, labels):\n        if self.alpha <= 0 or torch.rand(1) > 0.35:\n            return images, labels\n\n        lam = np.random.beta(self.alpha, self.alpha)\n        batch_size = images.size(0)\n        indices = torch.randperm(batch_size)\n        shuffled_images = images[indices]\n        shuffled_labels = labels[indices]\n\n        mixed_images = lam * images + (1 - lam) * shuffled_images\n        mixed_labels = torch.stack([\n            lam * labels,\n            (1 - lam) * shuffled_labels\n        ], dim=1)\n        return mixed_images, mixed_labels\n\nclass RandomMask(nn.Module):\n    def __init__(self, height_ratio, width_ratio, is_time_mask=True):\n        super().__init__()\n        self.height_ratio = height_ratio\n        self.width_ratio = width_ratio\n        self.is_time_mask = is_time_mask\n\n    def forward(self, img):\n        if torch.rand(1) > 0.35:\n            return img\n\n        c, h, w = img.shape\n        if self.is_time_mask:\n            mask_width = int(w * np.random.uniform(*self.width_ratio))\n            mask_start = np.random.randint(0, w - mask_width)\n            img[:, :, mask_start:mask_start + mask_width] = 0\n        else:\n            mask_height = int(h * np.random.uniform(*self.height_ratio))\n            mask_start = np.random.randint(0, h - mask_height)\n            img[:, mask_start:mask_start + mask_height, :] = 0\n        return img\nclass AudioAugmenter(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.augmenters = nn.ModuleList([\n            MixUp(alpha=0.4),\n            RandomMask((1.0, 1.0), (0.06, 0.12), True),\n            RandomMask((0.06, 0.1), (1.0, 1.0), False)\n        ])\n\n    def forward(self, img, label):\n        if img.ndim == 3 and img.shape[2] in [1, 3]:\n            img = img.permute(2, 0, 1)\n\n        for aug in self.augmenters:\n            if isinstance(aug, MixUp):\n                img, label = aug(img, label)\n            else:\n                img = aug(img)\n\n        if img.shape[0] in [1, 3]:\n            img = img.permute(1, 2, 0)\n        return img, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T19:19:33.669863Z","iopub.execute_input":"2025-04-24T19:19:33.670058Z","iopub.status.idle":"2025-04-24T19:19:33.688051Z","shell.execute_reply.started":"2025-04-24T19:19:33.670042Z","shell.execute_reply":"2025-04-24T19:19:33.687337Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 数据预处理与加载","metadata":{}},{"cell_type":"code","source":"def build_decoder(with_labels=True):\n    def get_audio(filepath: str) -> torch.Tensor:\n        \"\"\"\n        PyTorch实现音频加载与预处理\n\n        :param filepath: .ogg音频文件路径\n        :return: 单声道音频张量，形状为 [num_samples]\n        \"\"\"\n        # 读取.ogg文件\n        waveform, sample_rate = torchaudio.load(filepath)\n\n        # 单声道转换（如果音频是立体声，则取左声道）\n        if waveform.shape[0] > 1:  # 检查是否为立体声\n            waveform = waveform[0:1, :]  # 取第一个声道\n\n        # 压缩维度并标准化\n        waveform = waveform.squeeze(0)  # 移除通道维度，形状变为 [samples]\n        return waveform\n\n\n\n    def crop_or_pad(audio: torch.Tensor,\n                    target_len: int,\n                    pad_mode: str = \"constant\") -> torch.Tensor:\n        \"\"\"\n        音频长度统一化 (PyTorch 实现)\n\n        动态调整音频长度至固定长度：\n        1. 音频过短：随机分配填充位置进行补零\n        2. 音频过长：随机选取起始位置截取片段\n        3. 支持多种填充模式（constant/reflect/replicate）\n\n        参数：\n            audio: 输入音频张量，形状为 [L]\n            target_len: 目标长度（样本点数）\n            pad_mode: 填充模式，支持 \"constant\"/\"reflect\"/\"replicate\"\n\n        返回：\n            调整后的音频张量，形状为 [target_len]\n        \"\"\"\n        audio_len = audio.size(0)\n        diff_len = abs(target_len - audio_len)\n\n        # 短音频补零\n        if audio_len < target_len:\n            pad1 = torch.randint(0, diff_len + 1, ()).item()\n            pad2 = diff_len - pad1\n            audio = F.pad(audio, (pad1, pad2), mode=pad_mode)\n\n        # 长音频随机截取\n        elif audio_len > target_len:\n            start_idx = torch.randint(0, diff_len + 1, ()).item()\n            audio = audio[start_idx: start_idx + target_len]\n\n        # 确保形状正确 [target_len]\n        return audio.reshape(-1)[:target_len]\n\n    def apply_preproc(spec: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        频谱图标准化与归一化 (PyTorch 实现)\n\n        功能说明：\n        1. Z-Score标准化：消除量纲差异，使数据均值为0，标准差为1\n        2. Min-Max归一化：压缩值域到[0,1]，提升模型训练稳定性\n        3. 鲁棒性处理：对静音段等特殊情况降级处理\n\n        参数：\n            spec: 输入频谱图张量，形状任意\n\n        返回：\n            预处理后的张量，保持原始形状\n        \"\"\"\n        # ================= Z-Score 标准化 =================\n        mean = torch.mean(spec)\n        std = torch.std(spec)\n\n        # 处理零标准差情况（静音段）\n        if torch.is_nonzero(std):\n            spec = (spec - mean) / std\n        else:  # 退化为中心化操作\n            spec = spec - mean\n\n        # ================= Min-Max 归一化 =================\n        min_val = torch.min(spec)\n        max_val = torch.max(spec)\n        delta = max_val - min_val\n\n        # 处理零值域情况（全零或全同值）\n        if torch.is_nonzero(delta):\n            spec = (spec - min_val) / delta\n        else:  # 退化到零基线\n            spec = spec - min_val\n\n        return spec\n\n    def get_target(target: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        One-Hot 标签编码 (PyTorch 实现)\n\n        功能：\n        1. 将整数类别标签转换为 One-Hot 编码\n        2. 确保输出形状为 [CFG.num_classes]\n        3. 支持标量输入和批量输入（需与损失函数兼容）\n\n        参数：\n            target: 输入标签张量，形状为 []（标量）或 [batch_size]\n\n        返回：\n            One-Hot 编码张量，形状为 [CFG.num_classes] 或 [batch_size, CFG.num_classes]\n        \"\"\"\n        # 确保输入为整型张量\n        target = target.long()\n\n        # 生成 One-Hot 编码\n        one_hot = torch.nn.functional.one_hot(target, num_classes=CFG.num_classes)\n\n        # 调整形状和数据类型\n        return one_hot.reshape(-1, CFG.num_classes).float().squeeze()\n\n    def decode(path: str) -> torch.Tensor:\n        \"\"\"\n        频谱图生成与处理 (PyTorch 实现)\n\n        功能：\n        1. 加载音频并统一长度\n        2. 生成Mel频谱图\n        3. 标准化与归一化\n        4. 三通道适配\n\n        参数：\n            path: 音频文件路径\n\n        返回：\n            3通道频谱图张量，形状为 [3, CFG.img_size[0], CFG.img_size[1]]\n        \"\"\"\n        # 加载音频 (假设已实现 get_audio 返回单声道张量)\n        audio = get_audio(path)  # 形状 [L]\n\n        # 统一音频长度\n        audio = crop_or_pad(audio, CFG.dim)  # 形状 [CFG.dim]\n\n        # 生成Mel频谱图\n        mel_spec_transform = MelSpectrogram(\n            sample_rate=CFG.sample_rate,\n            n_fft=CFG.nfft,\n            hop_length=CFG.hop_length,\n            n_mels=CFG.img_size[0],\n            power=2.0  # 与TensorFlow的MelSpectrogram默认参数对齐[7](@ref)\n        )\n        spec = mel_spec_transform(audio)  # 形状 [n_mels, time_steps]\n\n        # 转换为对数刻度（模拟TensorFlow的return_decibel=True）\n        spec = torchaudio.functional.amplitude_to_DB(spec, multiplier=20, amin=1e-5, db_multiplier=0)\n\n        # 标准化与归一化\n        spec = apply_preproc(spec)\n\n        # 三通道适配 (模仿tf.tile)\n        spec = spec.unsqueeze(0)  # 添加通道维度 [1, n_mels, time_steps]\n        spec = spec.repeat(3, 1, 1)  # 复制为3通道 [3, n_mels, time_steps]\n\n        # 调整形状至目标尺寸（需确保 hop_length 参数设置正确）\n        spec = spec[..., :CFG.img_size[1]]  # 截断时间轴至目标长度\n        return spec\n\n    def decode_with_labels(path, label):\n        label = get_target(label)\n        return decode(path), label\n\n    return decode_with_labels if with_labels else decode\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T19:19:33.688732Z","iopub.execute_input":"2025-04-24T19:19:33.688918Z","iopub.status.idle":"2025-04-24T19:19:33.703001Z","shell.execute_reply.started":"2025-04-24T19:19:33.688903Z","shell.execute_reply":"2025-04-24T19:19:33.702444Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 数据集与数据加载器","metadata":{}},{"cell_type":"code","source":"class MixUp(nn.Module):\n    \"\"\"\n    MixUp：以alpha=0.4的参数混合两个样本的频谱图和标签，创造新的训练样本\n    \"\"\"\n    def __init__(self, alpha=0.4):\n        super().__init__()\n        self.alpha = alpha\n\n    def forward(self, images, labels):\n        if self.alpha <= 0 or torch.rand(1) > 0.35:\n            return images, labels\n\n        # 生成混合比例\n        lam = np.random.beta(self.alpha, self.alpha)\n        batch_size = images.size(0)\n\n        # 随机打乱样本顺序\n        indices = torch.randperm(batch_size)\n        shuffled_images = images[indices]\n        shuffled_labels = labels[indices]\n\n        # 混合图像和标签\n        mixed_images = lam * images + (1 - lam) * shuffled_images\n        mixed_labels = torch.stack([\n            lam * labels,\n            (1 - lam) * shuffled_labels\n        ], dim=1)\n\n        return mixed_images, mixed_labels\n\n\nclass RandomMask(nn.Module):\n    def __init__(self, height_ratio, width_ratio, is_time_mask=True):\n        super().__init__()\n        self.height_ratio = height_ratio\n        self.width_ratio = width_ratio\n        self.is_time_mask = is_time_mask\n\n    def forward(self, img):\n        if torch.rand(1) > 0.35:\n            return img\n\n        # 频谱图维度处理（假设输入为[C, Freq, Time]）\n        c, h, w = img.shape\n\n        # 时间掩码（垂直条）\n        if self.is_time_mask:\n            mask_width = int(w * np.random.uniform(*self.width_ratio))\n            mask_start = np.random.randint(0, w - mask_width)\n            img[:, :, mask_start:mask_start + mask_width] = 0\n        # 频率掩码（水平条）\n        else:\n            mask_height = int(h * np.random.uniform(*self.height_ratio))\n            mask_start = np.random.randint(0, h - mask_height)\n            img[:, mask_start:mask_start + mask_height, :] = 0\n\n        return img\n\n\nclass AudioAugmenter(nn.Module):\n    \"\"\"可序列化的增强流水线\"\"\"\n\n    def __init__(self):\n        super().__init__()\n        self.augmenters = nn.ModuleList([\n            MixUp(alpha=0.4),\n            RandomMask((1.0, 1.0), (0.06, 0.12), True),\n            RandomMask((0.06, 0.1), (1.0, 1.0), False)\n        ])\n\n    def forward(self, img, label):\n        # 统一维度格式为 [C, H, W]\n        if img.ndim == 3 and img.shape[2] in [1, 3]:  # [H,W,C] -> [C,H,W]\n            img = img.permute(2, 0, 1)\n\n        for aug in self.augmenters:\n            if isinstance(aug, MixUp):\n                img, label = aug(img, label)\n            else:\n                img = aug(img)\n\n        # 恢复原始维度格式\n        if img.shape[0] in [1, 3]:  # [C,H,W] -> [H,W,C]\n            img = img.permute(1, 2, 0)\n\n        return img, label\n\n\ndef build_augmenter():\n    \"\"\"构建可序列化的增强器\"\"\"\n    return AudioAugmenter()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T19:19:33.703989Z","iopub.execute_input":"2025-04-24T19:19:33.704255Z","iopub.status.idle":"2025-04-24T19:19:33.725082Z","shell.execute_reply.started":"2025-04-24T19:19:33.704232Z","shell.execute_reply":"2025-04-24T19:19:33.724505Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 模型定义","metadata":{}},{"cell_type":"code","source":"    class ImageClassifier(nn.Module):\n        def __init__(self, num_classes, preset='efficientnet_v2_s'):\n            super().__init__()\n            # 加载预训练主干网络\n            self.backbone = models.efficientnet_v2_s(pretrained=True)\n            in_features = self.backbone.classifier[1].in_features\n\n            # 替换分类头\n            self.backbone.classifier = nn.Sequential(\n                nn.Dropout(p=0.2, inplace=True),\n                nn.Linear(in_features, num_classes)\n            )\n\n        def forward(self, x):\n            return self.backbone(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T19:19:33.725835Z","iopub.execute_input":"2025-04-24T19:19:33.726332Z","iopub.status.idle":"2025-04-24T19:19:33.742652Z","shell.execute_reply.started":"2025-04-24T19:19:33.726308Z","shell.execute_reply":"2025-04-24T19:19:33.741945Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 训练准备","metadata":{}},{"cell_type":"code","source":"os.listdir(\"/kaggle/working\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T19:19:51.626662Z","iopub.execute_input":"2025-04-24T19:19:51.627390Z","iopub.status.idle":"2025-04-24T19:19:51.632413Z","shell.execute_reply.started":"2025-04-24T19:19:51.627359Z","shell.execute_reply":"2025-04-24T19:19:51.631863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(f'{CFG.BASE_PATH_input}/train.csv')\ndf['filepath'] = CFG.BASE_PATH_input + '/train_audio/' + df.filename\ndf['target'] = df.primary_label.map(CFG.name2label)\ndf['filename'] = df.filepath.map(lambda x: x.split('/')[-1])\ndf['xc_id'] = df.filepath.map(lambda x: x.split('/')[-1].split('.')[0])\ntrain_df, valid_df = train_test_split(df, test_size=0.2, stratify=df['target'], random_state=CFG.seed)\ntrain_df = train_df.reset_index(drop=True).copy()\nvalid_df = valid_df.reset_index(drop=True).copy()\nprint(f'{CFG.BASE_PATH_work}/train_df.csv')\ntrain_df.to_csv(f'{CFG.BASE_PATH_work}/train_df.csv', index=False)\nvalid_df.to_csv(f'{CFG.BASE_PATH_work}/valid_df.csv', index=False)\n# print(df)\n# Display rwos\n# df.head(2)\ntrain_mel_df = []\nvalid_mel_df = []\ntrain_labels = []\nvalid_labels = []\nif not os.path.exists(f'{CFG.BASE_PATH_work}/train_mel_path.npy'):\n    \"\"\"训练集\"\"\"\n    for (idx, row_i) in tqdm(train_df.iterrows(), total=train_df.shape[0], desc=\"训练集\"):\n        # print(row_i)\n        train_label = row_i['target']\n        file_path = row_i['filepath']\n        train_mel_path = file_path.replace(\"train_audio\", \"train_mel\")\n        train_mel_path = train_mel_path.replace(CFG.BASE_PATH_input, CFG.BASE_PATH_work)\n        train_mel_path = train_mel_path.replace(\".ogg\", \".npy\")\n        train_label = torch.tensor([int(train_label)])\n        decode_with_labels = build_decoder()\n        mel, target = decode_with_labels(file_path, train_label)\n        train_labels.append(target)\n        # torch.save(mel, mel_path)\n        if not os.path.exists(os.path.dirname(train_mel_path)):\n            os.makedirs(os.path.dirname(train_mel_path))\n        np.save(train_mel_path, mel.numpy())\n        train_mel_df.append([train_mel_path])\n        # if idx == 3:\n        #     break\n    train_mel_df = pd.DataFrame(train_mel_df, columns=[\"mel_path\"])\n    train_mel_df.to_csv(f'{CFG.BASE_PATH_work}/train_mel_path.csv', index=None)\n    train_labels = np.array(train_labels)\n    np.save(f'{CFG.BASE_PATH_work}/train_labels.npy', train_labels)\n\nif not os.path.exists(f'{CFG.BASE_PATH_work}/valid_mel_path.npy'):\n    \"\"\"测试集\"\"\"\n    for (idx, row_i) in tqdm(valid_df.iterrows(), total=valid_df.shape[0], desc=\"测试集\"):\n        # print(row_i)\n        valid_label = row_i['target']\n        file_path = row_i['filepath']\n        valid_mel_path = file_path.replace(\"train_audio\", \"valid_mel\")\n        valid_mel_path = train_mel_path.replace(CFG.BASE_PATH_input, CFG.BASE_PATH_work)\n        valid_mel_path = valid_mel_path.replace(\".ogg\", \".npy\")\n        valid_label = torch.tensor([int(valid_label)])\n        decode_with_labels = build_decoder()\n        mel, target = decode_with_labels(file_path, valid_label)\n        valid_labels.append(target)\n        # torch.save(mel, mel_path)\n        if not os.path.exists(os.path.dirname(valid_mel_path)):\n            os.makedirs(os.path.dirname(valid_mel_path))\n        np.save(valid_mel_path, mel.numpy())\n        valid_mel_df.append([valid_mel_path])\n        # if idx == 3:\n        #     break\n    valid_mel_df = pd.DataFrame(valid_mel_df, columns=[\"mel_path\"])\n    valid_mel_df.to_csv(f'{CFG.BASE_PATH_work}/valid_mel_path.csv', index=None)\n    valid_labels = np.array(valid_labels)\n    np.save(f'{CFG.BASE_PATH_work}/valid_labels.npy', valid_labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T19:20:07.404524Z","iopub.execute_input":"2025-04-24T19:20:07.405055Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 初始化配置\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# 加载数据路径\ntrain_df = pd.read_csv(f'{CFG.BASE_PATH_work}/train_df.csv')\nvalid_df = pd.read_csv(f'{CFG.BASE_PATH_work}/valid_df.csv')\n\ntrain_mel_path = pd.read_csv(f'{CFG.BASE_PATH_work}/train_mel_path.csv').to_numpy().ravel().tolist()\nvalid_mel_path = pd.read_csv(f'{CFG.BASE_PATH_work}/valid_mel_path.csv').to_numpy().ravel().tolist()\n\n# train_mel_path = train_df['filepath'].to_numpy().ravel().tolist()\n# valid_mel_path = valid_df['filepath'].to_numpy().ravel().tolist()\n\n# 加载标签\ntrain_labels = torch.LongTensor(np.load(f\"{CFG.BASE_PATH_work}/train_labels.npy\")).argmax(dim=1)\nvalid_labels = torch.LongTensor(np.load(f\"{CFG.BASE_PATH_work}/valid_labels.npy\")).argmax(dim=1)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# train_mel_path[0]\nvalid_mel_path[0]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AudioDataset(Dataset):\n    def __init__(self,\n                 spec_paths: list,  # 新增频谱图路径参数\n                 labels=None,\n                 augment_fn=None):\n        \"\"\"\n        改进后的音频数据集类\n        :param spec_paths: 预先生成的梅尔频谱图路径列表\n        :param labels: 标签列表（可选）\n        :param augment_fn: 数据增强函数\n        \"\"\"\n        self.spec_paths = spec_paths\n        self.labels = labels\n        self.augment_fn = augment_fn or None\n\n    def _load_mel_spec(self, path: str) -> torch.Tensor:\n        \"\"\"加载预存的梅尔频谱图\"\"\"\n        spec = torch.from_numpy(np.load(path)).float()\n\n        # 通道处理（单通道转三通道）\n        if spec.shape[0] == 1:\n            spec = spec.repeat(3, 1, 1)\n        return spec\n\n    def __getitem__(self, idx):\n        # 加载预存频谱\n        spec = self._load_mel_spec(self.spec_paths[idx])\n        label = self.labels[idx] if (self.labels is not None) else None\n        # 应用数据增强\n        if self.augment_fn:\n            label = self.labels[idx] if (self.labels is not None) else None\n            spec, label = self.augment_fn(spec, label)\n\n        return (spec, label) if (self.labels is not None) else spec\n\n    def __len__(self):\n        return len(self.spec_paths)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_dataloader(paths, labels=None, batch_size=32,\n                     augment_fn=None,\n                     cache=False, augment=False, shuffle=2048, num_workers=0):\n    \"\"\"\n    PyTorch数据加载管道\n    :param shuffle: 设置为True启用随机洗牌，False则关闭（原TF的shuffle参数为缓冲区大小）\n    \"\"\"\n    # 创建数据集实例\n    dataset = AudioDataset(paths, labels, augment_fn if augment else None)\n\n    # 缓存机制\n    if cache:\n        # 预加载所有数据到内存（适合小数据集）\n        specs, labels = zip(*[dataset[i] for i in range(len(dataset))])\n        dataset = TensorDataset(torch.stack(specs), torch.tensor(labels))\n\n    # 构建DataLoader\n    loader = DataLoader(\n        dataset,\n        batch_size=batch_size,\n        shuffle=bool(shuffle),  # 原TF的shuffle参数转换为布尔值\n        num_workers=num_workers,  # 并行加载进程数（替代TF的AUTOTUNE）\n        pin_memory=True,  # 加速GPU传输\n        drop_last=True  # 对应原TF的drop_remainder\n    )\n    return loader","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 创建数据加载器\naugmenter = AudioAugmenter()\ntrain_loader = build_dataloader(train_mel_path, train_labels, CFG.batch_size, \n                               augment_fn=None, augment=False)\nvalid_loader = build_dataloader(valid_mel_path, valid_labels, CFG.batch_size)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 初始化模型\nmodel = ImageClassifier(CFG.num_classes).to(device)\n# 加载自定义权重\n# model.load_state_dict(torch.load(CFG.model_weight))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 训练循环","metadata":{}},{"cell_type":"code","source":"# 损失函数\nclass LabelSmoothCrossEntropy(nn.Module):\n    def __init__(self, smoothing=0.02):\n        super().__init__()\n        self.smoothing = smoothing\n\n    def forward(self, logits, targets):\n        targets = targets.to(logits.device).long()\n        log_probs = torch.log_softmax(logits, dim=-1)\n        nll_loss = -log_probs.gather(dim=-1, index=targets.unsqueeze(1)).squeeze(1)\n        smooth_loss = -log_probs.mean(dim=-1)\n        loss = (1 - self.smoothing) * nll_loss + self.smoothing * smooth_loss\n        return loss.mean()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 优化器与评估指标\noptimizer = optim.Adam(model.parameters())\ncriterion = LabelSmoothCrossEntropy(smoothing=0.02)\nauc_metric = MulticlassAUROC(num_classes=CFG.num_classes, average=\"macro\").to(device)\n\n# 学习率调度\ndef get_lr_scheduler(optimizer, batch_size=8, mode='cos', epochs=30):\n    lr_start, lr_max, lr_min = 5e-5, 8e-6 * batch_size, 1e-5\n    lr_ramp_ep, lr_sus_ep = 3, 0\n\n    def lr_lambda(epoch):\n        if epoch < lr_ramp_ep:\n            return (lr_max - lr_start) / lr_ramp_ep * epoch + lr_start\n        elif epoch < lr_ramp_ep + lr_sus_ep:\n            return lr_max\n        elif mode == 'cos':\n            decay_total_epochs = epochs - lr_ramp_ep - lr_sus_ep + 3\n            decay_epoch_index = epoch - lr_ramp_ep - lr_sus_ep\n            phase = math.pi * decay_epoch_index / decay_total_epochs\n            return (lr_max - lr_min) * 0.5 * (1 + math.cos(phase)) + lr_min\n    return LambdaLR(optimizer, lr_lambda)\n\nscheduler = get_lr_scheduler(optimizer, CFG.batch_size)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 训练过程\nbest_auc = 0\nfor epoch in trange(CFG.epochs, desc=\"Epochs\"):\n    # 训练阶段\n    model.train()\n    # for inputs, labels in tqdm(train_loader, desc=f\"Train Epoch {epoch+1}\"):\n    for inputs, labels in train_loader:\n        inputs, labels = inputs.to(device), labels.to(device)\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n    \n    # 验证阶段\n    model.eval()\n    val_auc = 0\n    with torch.no_grad():\n        for inputs, labels in valid_loader:\n            inputs, labels = inputs.to(device), labels.to(device)\n            outputs = model(inputs)\n            auc_metric.update(outputs, labels)\n        \n        val_auc = auc_metric.compute()\n        if val_auc > best_auc:\n            best_auc = val_auc\n            torch.save(model.state_dict(), f\"{CFG.BASE_PATH_work}best_model_weights.pth\")\n        auc_metric.reset()\n    \n    scheduler.step()\n    print(f\"Epoch {epoch+1}/{CFG.epochs} | Val AUC: {val_auc:.4f} | Best AUC: {best_auc:.4f}\")\n\nprint(\"训练完成!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}