{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.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":70203,"databundleVersionId":8068726,"sourceType":"competition"},{"sourceId":8090934,"sourceType":"datasetVersion","datasetId":4776799},{"sourceId":154204277,"sourceType":"kernelVersion"},{"sourceId":167220511,"sourceType":"kernelVersion"}],"dockerImageVersionId":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport gc\nimport sys\nimport cv2\nimport math\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import KFold\nimport librosa\nfrom scipy import signal as sci_signal\n \nimport torch\nfrom torch import nn\nfrom torchvision.models import efficientnet\n \nimport albumentations as albu\n \nimport pytorch_lightning as pl\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts\nfrom pytorch_lightning.callbacks import ModelCheckpoint, TQDMProgressBar\n \n#引入竞赛的打分方式\nsys.path.append('/kaggle/usr/lib/birdclef-roc-auc/')\nsys.path.append('/kaggle/usr/lib/kaggle_metric_utilities')\n\nfrom metric import score","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-06-19T01:01:34.434302Z","iopub.execute_input":"2024-06-19T01:01:34.434940Z","iopub.status.idle":"2024-06-19T01:01:45.511652Z","shell.execute_reply.started":"2024-06-19T01:01:34.434904Z","shell.execute_reply":"2024-06-19T01:01:45.510666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta_df = pd.read_csv('/kaggle/input/birdclef-2024/train_metadata.csv')\nmeta_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:03:39.277172Z","iopub.execute_input":"2024-06-19T01:03:39.277765Z","iopub.status.idle":"2024-06-19T01:03:39.476024Z","shell.execute_reply.started":"2024-06-19T01:03:39.277732Z","shell.execute_reply":"2024-06-19T01:03:39.475002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta_df.info()\nnum_classes = len(meta_df['primary_label'].unique())\nprint(f'鸟类类别总数：{num_classes}')","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:03:42.218222Z","iopub.execute_input":"2024-06-19T01:03:42.218615Z","iopub.status.idle":"2024-06-19T01:03:42.264004Z","shell.execute_reply.started":"2024-06-19T01:03:42.218588Z","shell.execute_reply":"2024-06-19T01:03:42.263106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"* **primary_label:       目标鸟类标签**\n* secondary_labels:    音频中含有的其他鸟类标签\n* type:                鸟类音频鸣叫类型 e.g. song,call\n* latitude&longtitude: 音频采集的经纬度\n* scientific_name:     学名\n* common_name:         俗名\n* author&license&url:  音频作者，许可，下载地址\n* rating:              音频可信度,[0,5]之间的浮点数\n* **filename:            对应的训练数据文件名**","metadata":{"execution":{"iopub.status.busy":"2024-06-18T10:11:40.870140Z","iopub.execute_input":"2024-06-18T10:11:40.871143Z","iopub.status.idle":"2024-06-18T10:11:40.879011Z","shell.execute_reply.started":"2024-06-18T10:11:40.871104Z","shell.execute_reply":"2024-06-18T10:11:40.877321Z"}}},{"cell_type":"code","source":"def count_audio_files(path):\n    folder_count = 0\n    audio_counts = {}\n\n    for folder, subfolders, files in os.walk(path):\n        if files:\n            folder_count += 1\n            folder_name = os.path.basename(folder)\n            audio_counts[folder_name] = len([file for file in files if file.endswith(\".ogg\")])\n    audio_file_paths = [os.path.join(folder, file) for folder, _, files in os.walk(path) for file in files if file.endswith(\".ogg\")]\n    return folder_count, audio_counts, audio_file_paths\n\ndirectory_path = \"/kaggle/input/birdclef-2024/train_audio\"\nfolder_count, audio_counts, audio_file_paths = count_audio_files(directory_path)\n\nprint(f\"总共 \\x1b[34m{folder_count}\\x1b[0m 种不同的鸟类音频文件\")","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:03:50.152020Z","iopub.execute_input":"2024-06-19T01:03:50.152431Z","iopub.status.idle":"2024-06-19T01:03:58.314696Z","shell.execute_reply.started":"2024-06-19T01:03:50.152397Z","shell.execute_reply":"2024-06-19T01:03:58.313712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"audio_file_path = '/kaggle/input/birdclef-2024/train_audio/asbfly/XC134896.ogg'\naudio_data, sample_rate = librosa.load(audio_file_path, sr=32000)\nplt.plot(audio_data)\nprint(len(audio_data) / sample_rate)","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:36:11.844956Z","iopub.execute_input":"2024-06-19T01:36:11.845853Z","iopub.status.idle":"2024-06-19T01:36:12.582907Z","shell.execute_reply.started":"2024-06-19T01:36:11.845816Z","shell.execute_reply":"2024-06-19T01:36:12.581980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import Audio\nAudio(audio_data,rate=32000)","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:37:29.267932Z","iopub.execute_input":"2024-06-19T01:37:29.268556Z","iopub.status.idle":"2024-06-19T01:37:29.308830Z","shell.execute_reply.started":"2024-06-19T01:37:29.268527Z","shell.execute_reply":"2024-06-19T01:37:29.307855Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spectrogram = librosa.amplitude_to_db(librosa.stft(audio_data), ref=np.max)\nlibrosa.display.specshow(spectrogram, sr=sample_rate, x_axis='time', y_axis='hz', cmap='jet')","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:50:02.264637Z","iopub.execute_input":"2024-06-19T01:50:02.265359Z","iopub.status.idle":"2024-06-19T01:50:04.084207Z","shell.execute_reply.started":"2024-06-19T01:50:02.265328Z","shell.execute_reply":"2024-06-19T01:50:04.083246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#配置参数\nclass config:\n    \n    #通用参数\n    SEED = 20241 #随机种子\n    DEVICE = 'cuda'  #竞赛要求必须用cpu\n    MIXED_PRECISION = False  #不使用混合精度训练，容易影响模型的泛化能力\n    OUTPUT_DIR = '/kaggle/working/'  #输出文件夹\n    \n    #数据参数\n    DATA_ROOT = '/kaggle/input/birdclef-2024'  #原始数据目录\n    PREPROCESSED_DATA_ROOT = '/kaggle/input/birdclef24-spectrograms-via-cupy'\n    LOAD_DATA = True  #使用预训练的数据，提高训练效率\n    FS = 32000  #采样率\n    N_FFT = 1095  #FFT点数\n    WIN_SIZE = 412  #频谱每段样本数量\n    WIN_LAP = 100  #频谱每段重叠样本数\n    MIN_FREQ = 40  #最小频率\n    MAX_FREQ = 15000  #最大频率\n    \n    #模型参数\n    MODEL_TYPE = 'efficientnet_b0'\n    \n    #数据集参数\n    BATCH_SIZE = 64\n    N_WORKERS = 4\n    \n    #数据增强\n    USE_XYMASKING = True\n    \n    #训练参数\n    FOLDS = 5  #k折参数\n    EPOCHS = 10  #最大迭代轮次\n    LR = 1e-3  #学习率\n    WEIGHT_DECAY = 1e-5  #优化器权重衰减\n    \n    #其他参数\n    VISUALIZE = True  # whether to visualize data and batch\n    \nprint('fix seed')\npl.seed_everything(config.SEED, workers=True)","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:07:08.847008Z","iopub.execute_input":"2024-06-19T01:07:08.847387Z","iopub.status.idle":"2024-06-19T01:07:08.862873Z","shell.execute_reply.started":"2024-06-19T01:07:08.847356Z","shell.execute_reply":"2024-06-19T01:07:08.861797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_list = sorted(os.listdir(os.path.join(config.DATA_ROOT, 'train_audio')))\nlabel_id_list = list(range(len(label_list)))\nlabel2id = dict(zip(label_list, label_id_list))\nid2label = dict(zip(label_id_list, label_list))","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:04:23.538562Z","iopub.execute_input":"2024-06-19T01:04:23.538939Z","iopub.status.idle":"2024-06-19T01:04:23.546024Z","shell.execute_reply.started":"2024-06-19T01:04:23.538910Z","shell.execute_reply":"2024-06-19T01:04:23.544972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0,1\"\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint('Using', torch.cuda.device_count(), 'GPU(s)')","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:04:24.796272Z","iopub.execute_input":"2024-06-19T01:04:24.797127Z","iopub.status.idle":"2024-06-19T01:04:24.829516Z","shell.execute_reply.started":"2024-06-19T01:04:24.797094Z","shell.execute_reply":"2024-06-19T01:04:24.828524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def oog2spec_via_cupy(audio_data):\n    \n    import cupy as cp\n    from cupyx.scipy import signal as cupy_signal\n    \n    audio_data = cp.array(audio_data)\n    \n    # handles NaNs\n    mean_signal = cp.nanmean(audio_data)\n    audio_data = cp.nan_to_num(audio_data, nan=mean_signal) if cp.isnan(audio_data).mean() < 1 else cp.zeros_like(audio_data)\n    \n    # to spec.\n    frequencies, times, spec_data = cupy_signal.spectrogram(\n        audio_data, \n        fs=config.FS, \n        nfft=config.N_FFT, \n        nperseg=config.WIN_SIZE, \n        noverlap=config.WIN_LAP, \n        window='hann'\n    )\n    \n    # Filter frequency range\n    valid_freq = (frequencies >= config.MIN_FREQ) & (frequencies <= config.MAX_FREQ)\n    spec_data = spec_data[valid_freq, :]\n    \n    # Log\n    spec_data = cp.log10(spec_data + 1e-20)\n    \n    # min/max normalize\n    spec_data = spec_data - spec_data.min()\n    spec_data = spec_data / spec_data.max()\n    \n    return spec_data.get()","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:04:26.186144Z","iopub.execute_input":"2024-06-19T01:04:26.187098Z","iopub.status.idle":"2024-06-19T01:04:26.194911Z","shell.execute_reply.started":"2024-06-19T01:04:26.187062Z","shell.execute_reply":"2024-06-19T01:04:26.193807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if config.LOAD_DATA:\n    print('load from file')\n    all_bird_data = np.load(f'{config.PREPROCESSED_DATA_ROOT}/spec_center_5sec_256_256.npy', allow_pickle=True).item()\nelse:\n    all_bird_data = dict()\n    for i, row_metadata in tqdm(train_df.iterrows()):\n\n        # load ogg\n        audio_data, _ = librosa.load(row_metadata.filepath, sr=config.FS)\n\n        # crop\n        n_copy = math.ceil(5 * config.FS / len(audio_data))\n        if n_copy > 1: audio_data = np.concatenate([audio_data]*n_copy)\n\n        start_idx = int(len(audio_data) / 2 - 2.5 * config.FS)\n        end_idx = int(start_idx + 5.0 * config.FS)\n        input_audio = audio_data[start_idx:end_idx]\n\n        # ogg to spec.\n        input_spec = oog2spec_via_cupy(input_audio)\n        \n        input_spec = cv2.resize(input_spec, (256, 256), interpolation=cv2.INTER_AREA)\n\n        all_bird_data[row_metadata.samplename] = input_spec.astype(np.float32)\n\n    # save to file\n    np.save(os.path.join(config.OUTPUT_DIR, f'spec_center_5sec_256_256.npy'), all_bird_data)","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:04:30.100543Z","iopub.execute_input":"2024-06-19T01:04:30.100868Z","iopub.status.idle":"2024-06-19T01:05:30.461959Z","shell.execute_reply.started":"2024-06-19T01:04:30.100845Z","shell.execute_reply":"2024-06-19T01:05:30.461143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i, (name, keys) in enumerate(all_bird_data.items()):\n    print(f'{name}: {keys}')\n    print(keys.shape)\n    if i == 5:\n        break","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:06:10.226345Z","iopub.execute_input":"2024-06-19T01:06:10.226725Z","iopub.status.idle":"2024-06-19T01:06:10.236135Z","shell.execute_reply.started":"2024-06-19T01:06:10.226694Z","shell.execute_reply":"2024-06-19T01:06:10.235046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = meta_df[['primary_label', 'rating', 'filename']].copy()\n\n# create target\ntrain_df['target'] = train_df.primary_label.map(label2id)\n# create filepath\ntrain_df['filepath'] = config.DATA_ROOT + '/train_audio/' + train_df.filename\n# create new sample name\ntrain_df['samplename'] = train_df.filename.map(lambda x: x.split('/')[0] + '-' + x.split('/')[-1].split('.')[0])\n\nprint(f'find {len(train_df)} samples')\n\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:06:13.717413Z","iopub.execute_input":"2024-06-19T01:06:13.718327Z","iopub.status.idle":"2024-06-19T01:06:13.777482Z","shell.execute_reply.started":"2024-06-19T01:06:13.718284Z","shell.execute_reply":"2024-06-19T01:06:13.776543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass BirdDataset(torch.utils.data.Dataset):\n    \n    def __init__(\n        self,\n        metadata,\n        augmentation=None,\n        mode='train'\n    ):\n        super().__init__()\n        self.metadata = metadata\n        self.augmentation = augmentation\n        self.mode = mode\n    \n    def __len__(self):\n        return len(self.metadata)\n    \n    def __getitem__(self, index):\n        \n        # 根据index获取频谱图数据\n        row_metadata = self.metadata.iloc[index]\n        input_spec = all_bird_data[row_metadata.samplename]\n        \n        # 数据增强\n        if self.augmentation is not None:\n            input_spec = self.augmentation(image=input_spec)['image']\n        \n        # 标签\n        target = row_metadata.target\n        \n        return torch.tensor(input_spec, dtype=torch.float32), torch.tensor(target, dtype=torch.long)","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:06:16.599187Z","iopub.execute_input":"2024-06-19T01:06:16.599839Z","iopub.status.idle":"2024-06-19T01:06:16.607820Z","shell.execute_reply.started":"2024-06-19T01:06:16.599807Z","shell.execute_reply":"2024-06-19T01:06:16.606633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transforms(_type):\n    \n    if _type == 'train':\n        return albu.Compose([\n            albu.HorizontalFlip(0.5),\n            albu.XYMasking(\n                p=0.3,\n                num_masks_x=(1, 3),\n                num_masks_y=(1, 3),\n                mask_x_length=(1, 10),\n                mask_y_length=(1, 20),\n            ) if config.USE_XYMASKING else albu.NoOp()\n        ])\n    elif _type == 'valid':\n        return albu.Compose([])","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:06:19.180546Z","iopub.execute_input":"2024-06-19T01:06:19.180980Z","iopub.status.idle":"2024-06-19T01:06:19.188942Z","shell.execute_reply.started":"2024-06-19T01:06:19.180947Z","shell.execute_reply":"2024-06-19T01:06:19.187771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_batch(ds, row=3, col=3):\n    fig = plt.figure(figsize=(10, 10))\n    img_index = np.random.randint(0, len(ds)-1, row*col)\n    \n    for i in range(len(img_index)):\n        img, label = dummy_dataset[img_index[i]]\n        \n        if isinstance(img, torch.Tensor):\n            img = img.detach().numpy()\n        \n        ax = fig.add_subplot(row, col, i + 1, xticks=[], yticks=[])\n        ax.imshow(img, cmap='jet')\n        ax.set_title(f'ID: {img_index[i]}; Target: {label}')\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:49:37.176950Z","iopub.execute_input":"2024-06-19T01:49:37.177664Z","iopub.status.idle":"2024-06-19T01:49:37.185295Z","shell.execute_reply.started":"2024-06-19T01:49:37.177634Z","shell.execute_reply":"2024-06-19T01:49:37.184325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dummy_dataset = BirdDataset(train_df, get_transforms('train'))\n\ntest_input, test_target = dummy_dataset[0]\nprint(test_input.detach().numpy().shape)\n\nif config.VISUALIZE:\n    show_batch(dummy_dataset)\n\ndel dummy_dataset","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:49:43.634602Z","iopub.execute_input":"2024-06-19T01:49:43.635258Z","iopub.status.idle":"2024-06-19T01:49:44.759999Z","shell.execute_reply.started":"2024-06-19T01:49:43.635226Z","shell.execute_reply":"2024-06-19T01:49:44.758963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EffNet(nn.Module):\n    \n    def __init__(self, model_type, n_classes, pretrained=True):\n        super().__init__()\n        \n        if model_type == 'efficientnet_b0':\n            if pretrained: weights = efficientnet.EfficientNet_B0_Weights.DEFAULT\n            else: weights = None\n            self.base_model = efficientnet.efficientnet_b0(weights=weights)\n        elif model_type == 'efficientnet_b1':\n            if pretrained: weights = efficientnet.EfficientNet_B1_Weights.DEFAULT\n            else: weights = None\n            self.base_model = efficientnet.efficientnet_b1(weights=weights)\n        elif model_type == 'efficientnet_b2':\n            if pretrained: weights = efficientnet.EfficientNet_B2_Weights.DEFAULT\n            else: weights = None\n            self.base_model = efficientnet.efficientnet_b2(weights=weights)\n        elif model_type == 'efficientnet_b3':\n            if pretrained: weights = efficientnet.EfficientNet_B3_Weights.DEFAULT\n            else: weights = None\n            self.base_model = efficientnet.efficientnet_b3(weights=weights)\n        else:\n            raise ValueError('model type not supported')\n        \n        self.base_model.classifier[1] = nn.Linear(self.base_model.classifier[1].in_features, n_classes, dtype=torch.float32)\n    \n    def forward(self, x):\n        x = x.unsqueeze(-1)\n        x = torch.cat([x, x, x], dim=3).permute(0, 3, 1, 2)\n        return self.base_model(x)","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:07:51.234159Z","iopub.execute_input":"2024-06-19T01:07:51.235108Z","iopub.status.idle":"2024-06-19T01:07:51.244905Z","shell.execute_reply.started":"2024-06-19T01:07:51.235072Z","shell.execute_reply":"2024-06-19T01:07:51.244024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#LightningModule 是 PyTorch Lightning 框架的核心，用于构建易于训练、验证和测试的模型\nclass BirdModel(pl.LightningModule):    \n    def __init__(self):\n        super().__init__()\n        \n        #主干网络\n        self.backbone = EffNet(config.MODEL_TYPE, n_classes=len(label_list))\n        \n        #交叉熵损失函数\n        self.loss_fn = nn.CrossEntropyLoss()\n        \n        #初始化一个列表，用于存储验证步骤的输出\n        self.validation_step_outputs = []\n        \n    def forward(self, images):\n        return self.backbone(images)\n    #定义优化器和学习率调度器\n    def configure_optimizers(self):\n        \n        #定义 Adam 优化器\n        model_optimizer = torch.optim.Adam(\n            filter(lambda p: p.requires_grad, self.parameters()),\n            lr=config.LR,\n            weight_decay=config.WEIGHT_DECAY\n        )\n        \n        #定义余弦退火学习率调度器\n        lr_scheduler = CosineAnnealingWarmRestarts(\n            model_optimizer,\n            T_0=config.EPOCHS,\n            T_mult=1,\n            eta_min=1e-6,\n            last_epoch=-1\n        )\n        \n        return {\n            'optimizer': model_optimizer,\n            'lr_scheduler': {\n                'scheduler': lr_scheduler,\n                'interval': 'epoch',\n                'monitor': 'val_loss',\n                'frequency': 1\n            }\n        }\n    #定义训练步骤\n    def training_step(self, batch, batch_idx):\n        \n        #获取输入\n        image, target = batch\n        image = image.to(self.device)\n        target = target.to(self.device)\n        \n        #前向传播\n        y_pred = self(image)\n        \n        #计算损失\n        train_loss = self.loss_fn(y_pred, target)\n        \n        #记录训练损失\n        self.log('train_loss', train_loss, True)\n        \n        return train_loss\n    #定义验证步骤\n    def validation_step(self, batch, batch_idx):\n        \n        #从验证批次中获取输入图像和目标标签\n        image, target = batch\n        image = image.to(self.device)\n        target = target.to(self.device)\n        \n        #在不计算梯度的情况下进行前向传播\n        with torch.no_grad():\n            y_pred = self(image)\n        #将预测输出和目标存储在validation_step_outputs列表中    \n        self.validation_step_outputs.append({\"logits\": y_pred, \"targets\": target})\n    #返回训练和验证数据的 DataLoader\n    def train_dataloader(self):\n        return self._train_dataloader\n \n    def validation_dataloader(self):\n        return self._validation_dataloader\n   #在每个验证周期结束时调用，用于处理整个验证集的结果\n    def on_validation_epoch_end(self):\n        \n        #合并验证批次数据\n        outputs = self.validation_step_outputs\n        \n        output_val = torch.cat([x['logits'] for x in outputs], dim=0).cpu().detach()\n        target_val = torch.cat([x['targets'] for x in outputs], dim=0).cpu().detach()\n        \n        #计算验证损失，使用存储的目标标签和模型的预测输出\n        val_loss = self.loss_fn(output_val, target_val)\n        \n        #将目标标签转换为独热编码格式\n        target_val = torch.nn.functional.one_hot(target_val, len(label_list))\n        \n        #评估指标\n        gt_df = pd.DataFrame(target_val.numpy().astype(np.float32), columns=label_list)\n        pred_df = pd.DataFrame(output_val.numpy().astype(np.float32), columns=label_list)\n        \n        gt_df['id'] = [f'id_{i}' for i in range(len(gt_df))]\n        pred_df['id'] = [f'id_{i}' for i in range(len(pred_df))]\n        \n        val_score = score(gt_df, pred_df, row_id_column_name='id')\n        #使用self.log记录验证损失和分数，以便在Lightning日志中跟踪\n        self.log(\"val_score\", val_score, True)\n        \n        #清除验证集输出\n        self.validation_step_outputs = list()\n        \n        return {'val_loss': val_loss, 'val_score': val_score}","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:07:58.475819Z","iopub.execute_input":"2024-06-19T01:07:58.476746Z","iopub.status.idle":"2024-06-19T01:07:58.496940Z","shell.execute_reply.started":"2024-06-19T01:07:58.476711Z","shell.execute_reply":"2024-06-19T01:07:58.495796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#使用给定的模型对 data_loader 中的数据进行预测，并且收集真实标签\ndef predict(data_loader, model):\n    model.to(config.DEVICE)\n    model.eval()\n    predictions = []\n    gts = []\n    for batch in tqdm(data_loader):\n        with torch.no_grad():\n            x, y = batch\n            x = x.cuda()\n            outputs = model(x)\n            outputs = nn.Softmax(dim=1)(outputs)\n        #收集预测结果和真实标签\n        predictions.append(outputs.detach().cpu())\n        gts.append(y.detach().cpu())\n    #合并预测结果和真实标签\n    predictions = torch.cat(predictions, dim=0).cpu().detach()\n    gts = torch.cat(gts, dim=0).cpu().detach()\n    #将真实标签转换为独热编码\n    gts = torch.nn.functional.one_hot(gts, len(label_list))\n    \n    return predictions.numpy().astype(np.float32), gts.numpy().astype(np.float32)","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:08:01.469409Z","iopub.execute_input":"2024-06-19T01:08:01.469753Z","iopub.status.idle":"2024-06-19T01:08:01.479614Z","shell.execute_reply.started":"2024-06-19T01:08:01.469728Z","shell.execute_reply":"2024-06-19T01:08:01.478395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#执行机器学习模型的训练和预测过程，用于 K 折交叉验证中的一个折（fold）\ndef run_training(fold_id, total_df):\n    #打印训练的相关信息，包括当前折的编号\n    print('================================================================')\n    print(f\"==== Running training for fold {fold_id} ====\")\n    \n    #创建数据集和数据加载器\n    train_df = total_df[total_df['fold'] != fold_id].copy()\n    valid_df = total_df[total_df['fold'] == fold_id].copy()\n    \n    print(f'Train Samples: {len(train_df)}')\n    print(f'Valid Samples: {len(valid_df)}')\n    \n    train_ds = BirdDataset(train_df, get_transforms('train'), 'train')\n    val_ds = BirdDataset(valid_df, get_transforms('valid'), 'valid')\n    \n    train_dl = torch.utils.data.DataLoader(\n        train_ds,\n        batch_size=config.BATCH_SIZE,\n        shuffle=True,\n        num_workers=config.N_WORKERS,\n        pin_memory=True,\n        persistent_workers=True\n    )\n    \n    val_dl = torch.utils.data.DataLoader(\n        val_ds,\n        batch_size=config.BATCH_SIZE * 2,\n        shuffle=False,\n        num_workers=config.N_WORKERS,\n        pin_memory=True,\n        persistent_workers=True\n    )\n    \n    #初始化模型\n    bird_model = BirdModel()\n    \n    #创建一个 ModelCheckpoint 回调，用于保存最佳模型\n    checkpoint_callback = ModelCheckpoint(monitor='val_score',\n                                          dirpath=config.OUTPUT_DIR,\n                                          save_top_k=1,\n                                          save_last=False,\n                                          save_weights_only=True,\n                                          filename=f\"fold_{fold_id}\",\n                                          mode='max')\n    callbacks_to_use = [checkpoint_callback, TQDMProgressBar(refresh_rate=1)]\n    \n    #初始化训练器\n    trainer = pl.Trainer(\n        max_epochs=config.EPOCHS,\n        val_check_interval=0.5,\n        callbacks=callbacks_to_use,\n        enable_model_summary=False,\n        accelerator=\"gpu\",\n        deterministic=True,\n        precision='16-mixed' if config.MIXED_PRECISION else 32,\n    )\n    \n    #训练模型\n    trainer.fit(bird_model, train_dataloaders=train_dl, val_dataloaders=val_dl)\n    \n    #预测\n    best_model_path = checkpoint_callback.best_model_path\n    weights = torch.load(best_model_path)['state_dict']\n    bird_model.load_state_dict(weights)\n    \n    preds, gts = predict(val_dl, bird_model)\n    \n    #创建包含预测结果和真实标签的 DataFrame\n    pred_df = pd.DataFrame(preds, columns=label_list)\n    pred_df['id'] = np.arange(len(pred_df))\n    gt_df = pd.DataFrame(gts, columns=label_list)\n    gt_df['id'] = np.arange(len(gt_df))\n    \n    #计算分数\n    val_score = score(gt_df, pred_df, row_id_column_name='id')\n    \n    #保存结果\n    pred_cols = [f'pred_{t}' for t in label_list]\n    valid_df = pd.concat([valid_df.reset_index(), pd.DataFrame(np.zeros((len(valid_df), len(label_list)*2)).astype(np.float32), columns=label_list+pred_cols)], axis=1)\n    valid_df[label_list] = gts\n    valid_df[pred_cols] = preds\n    valid_df.to_csv(f\"{config.OUTPUT_DIR}/pred_df_f{fold_id}.csv\", index=False)\n    \n    return preds, gts, val_score","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:08:02.929767Z","iopub.execute_input":"2024-06-19T01:08:02.930104Z","iopub.status.idle":"2024-06-19T01:08:02.945942Z","shell.execute_reply.started":"2024-06-19T01:08:02.930080Z","shell.execute_reply":"2024-06-19T01:08:02.944892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kf = KFold(n_splits=config.FOLDS, shuffle=True, random_state=config.SEED)\ntrain_df['fold'] = 0\nfor fold, (train_idx, val_idx) in enumerate(kf.split(train_df)):\n    train_df.loc[val_idx, 'fold'] = fold\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-06-19T01:08:05.776063Z","iopub.execute_input":"2024-06-19T01:08:05.776502Z","iopub.status.idle":"2024-06-19T01:08:05.803367Z","shell.execute_reply.started":"2024-06-19T01:08:05.776460Z","shell.execute_reply":"2024-06-19T01:08:05.802491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#训练\ntorch.set_float32_matmul_precision('high')\n \n#记录\nfold_val_score_list = list()\noof_df = train_df.copy()\npred_cols = [f'pred_{t}' for t in label_list]\noof_df = pd.concat([oof_df, pd.DataFrame(np.zeros((len(oof_df), len(pred_cols)*2)).astype(np.float32), columns=label_list+pred_cols)], axis=1)\n \nfor f in range(config.FOLDS):\n    \n    #获取当前折的验证集索引\n    val_idx = list(train_df[train_df['fold'] == f].index)\n    \n    #调用 run_training 函数进行训练和验证，获取验证集的预测 val_preds、真实标签 val_gts 和验证分数 val_score\n    val_preds, val_gts, val_score = run_training(f, train_df)\n    \n    #将当前折的验证集真实标签和预测结果更新到 oof_df，并记录当前折的验证分数\n    oof_df.loc[val_idx, label_list] = val_gts\n    oof_df.loc[val_idx, pred_cols] = val_preds\n    fold_val_score_list.append(val_score)\n\n    \nfor idx, val_score in enumerate(fold_val_score_list):\n    print(f'Fold {idx} Val Score: {val_score:.5f}')\n#计算整个 OOF 的评估分数\noof_gt_df = oof_df[['samplename'] + label_list].copy()\noof_pred_df = oof_df[['samplename'] + pred_cols].copy()\noof_pred_df.columns = ['samplename'] + label_list\noof_score = score(oof_gt_df, oof_pred_df, 'samplename')\nprint(f'OOF Score: {oof_score:.5f}')\n# 将包含训练过程中所有预测结果的 oof_df 保存为 CSV 文件\noof_df.to_csv(f\"{config.OUTPUT_DIR}/oof_pred.csv\", index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# #训练\n# torch.set_float32_matmul_precision('high')\n \n# #记录\n# fold_val_score_list = list()\n# oof_df = train_df.copy()\n# pred_cols = [f'pred_{t}' for t in label_list]\n# oof_df = pd.concat([oof_df, pd.DataFrame(np.zeros((len(oof_df), len(pred_cols)*2)).astype(np.float32), columns=label_list+pred_cols)], axis=1)\n \n# for f in range(config.FOLDS):\n    \n#     #获取当前折的验证集索引\n#     val_idx = list(train_df[train_df['fold'] == f].index)\n    \n#     #调用 run_training 函数进行训练和验证，获取验证集的预测 val_preds、真实标签 val_gts 和验证分数 val_score\n#     val_preds, val_gts, val_score = run_training(f, train_df)\n    \n#     #将当前折的验证集真实标签和预测结果更新到 oof_df，并记录当前折的验证分数\n#     oof_df.loc[val_idx, label_list] = val_gts\n#     oof_df.loc[val_idx, pred_cols] = val_preds\n#     fold_val_score_list.append(val_score)\n    \n#     break\n\n    \n# for idx, val_score in enumerate(fold_val_score_list):\n#     print(f'Fold {idx} Val Score: {val_score:.5f}')\n# # #计算整个 OOF 的评估分数\n# # oof_gt_df = oof_df[['samplename'] + label_list].copy()\n# # oof_pred_df = oof_df[['samplename'] + pred_cols].copy()\n# # oof_pred_df.columns = ['samplename'] + label_list\n# # oof_score = score(oof_gt_df, oof_pred_df, 'samplename')\n# # print(f'OOF Score: {oof_score:.5f}')\n# #将包含训练过程中所有预测结果的 oof_df 保存为 CSV 文件\n# oof_df.to_csv(f\"{config.OUTPUT_DIR}/oof_pred.csv\", index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}