{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":21669,"databundleVersionId":1692278,"sourceType":"competition"}],"dockerImageVersionId":30805,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"#  Идея решения\n\n1. Аудиоинтервал преобразуется в mel-спектрограмму, которая представляется в виде изображения\n2. На полученных изображения обучается нейронная сеть resnet101\n3. Для предсказаний файл со звуком будет разбиваться на равномерные интревалы\n\n# Зависимости\n\n* `os` -- поиск по директориям\n* `copy` -- необходимо для функции `deepcopy`, чтобы скопировать веса лучшей модели в процессе обучения\n* `pandas` и `numpy` -- работа с массивами и csv\n* `torch` -- обучение и запуск модели resnet\n* `librosa` -- обработка звуковых данных\n* `skimage` -- необходимо для функции `resize`, чтобы преобразовывать полученное изображение mel-спектрограммы в размер, который на вход принимает resnet\n* `torchvision` -- реализация модели resnet101\n* `tqdm` -- красивый вывод прогресса\n* `sklearn` -- необходим для kfold для разбиения процесса обучения\n* `matplotlib` -- построение графиков, нужно один раз, чтобы посмотреть на изображения mel-спектрограммы\n","metadata":{}},{"cell_type":"code","source":"import os\nimport copy\nimport numpy as np\nimport pandas as pd\nimport torch\nimport librosa\nfrom skimage.transform import resize\nfrom torchvision.models import resnet101\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm\nfrom sklearn.model_selection import KFold\n\nimport matplotlib.pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T22:56:16.822423Z","iopub.execute_input":"2024-12-15T22:56:16.822743Z","iopub.status.idle":"2024-12-15T22:56:16.827728Z","shell.execute_reply.started":"2024-12-15T22:56:16.822716Z","shell.execute_reply":"2024-12-15T22:56:16.826846Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Входные данные\n\nДанные для обучения лежат по пути `input/rfcx-species-audio-detection/train/`, а разметка данных в файле `input/rfcx-species-audio-detection/train_tp.csv`.\nДля отправки решения необходимо выполнить предсказание для аудиофайлов в папке `input/rfcx-species-audio-detection/test/`","metadata":{}},{"cell_type":"code","source":"train_audio_dir = \"../input/rfcx-species-audio-detection/train/\"\ntest_audio_dir = \"../input/rfcx-species-audio-detection/test/\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T22:56:16.832556Z","iopub.execute_input":"2024-12-15T22:56:16.832856Z","iopub.status.idle":"2024-12-15T22:56:16.839767Z","shell.execute_reply.started":"2024-12-15T22:56:16.832793Z","shell.execute_reply":"2024-12-15T22:56:16.838729Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Функции чтения аудиофайлов\n\n## Функция get_audio\n\nПо переданным пути до папки с аудиофайлами (`audio_dir`) и id записи (`fileid`) конструирует путь до файла и считывает его в оперативную память\n\n## Фунцкция get_audio_part\n\nАналогично функции get_audio только дополнительно принимает начальный и конечный временные метки с интересующим интервалом. Также принимается параметр `duration`, который определяет длину выдаваемого интервала. Работает следующим образом: берётся середина интервала в \\[*начальный временная метка*; *конечная временная метка*\\]. Затем относительного этой центральной точки в обе стороны временной линии берётся значание в половину желаемой длины.\n\n## Кэширование\n\nДля переиспользования данных, все возвращаемые данные складываются в кэш и затем берутся оттуда. Для отключения использования кэша предусмотрен параметр `nocache`. Может быть полезен при больших объемах данных, которые не имеет смысла кэшировать (например, предсказание тестовых данных)\n\n### Функция clear_cache\n\nОчищает кэш, чтобы освободить оперативную память\n","metadata":{}},{"cell_type":"code","source":"# fileid -> (wave_form, rate)\ncache_audio = {}\n# fileid_tmin_tmax -> (wave_form part, rate)\ncache_splitted = {}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T22:56:16.841121Z","iopub.execute_input":"2024-12-15T22:56:16.841381Z","iopub.status.idle":"2024-12-15T22:56:16.882730Z","shell.execute_reply.started":"2024-12-15T22:56:16.841349Z","shell.execute_reply":"2024-12-15T22:56:16.882076Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def clear_cache():\n    global cache_audio, cache_splitted\n\n    cache_audio = {}\n    cache_splitted = {}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T22:56:16.883641Z","iopub.execute_input":"2024-12-15T22:56:16.883896Z","iopub.status.idle":"2024-12-15T22:56:16.893642Z","shell.execute_reply.started":"2024-12-15T22:56:16.883871Z","shell.execute_reply":"2024-12-15T22:56:16.892988Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_audio(audio_dir: str, fileid: str, nocache: bool = False):\n    if fileid in cache_audio:\n        return cache_audio[fileid]\n    filename = f\"{audio_dir}/{fileid}.flac\"\n    if not os.path.isfile(filename):\n        raise ValueError(f\"Cannot find {fileid} in audio dir\")\n    waveform, sample_rate = librosa.load(filename, sr=None)\n    if not nocache:\n        cache_audio[fileid] = (waveform, sample_rate)\n    return (waveform, sample_rate)\n\n\ndef get_audio_part(audio_dir: str, fileid: str, tmin: float, tmax: float, duration: int, nocache: bool = False):\n    ts = round((float(tmax) + float(tmin)) / 2, 3)\n    key = f\"{fileid}_{ts}\"\n    if key in cache_splitted:\n        return cache_splitted[key]\n    wave, rate = get_audio(audio_dir, fileid)\n    hd = duration / 2\n    start = int((ts - hd) * rate)\n    end = int((ts + hd) * rate)\n    wave_len = len(wave)\n    # Sometimes start+end do not equals duration due float multiplication\n    # Just add diff to the end\n    dur_diff = end - start - rate * duration\n    if dur_diff != 0:\n        end -= dur_diff\n    # If we out of bounds, shift audio window\n    if start < 0:\n        end -= start\n        start = 0\n    if end > wave_len:\n        start -= (end - wave_len)\n        end = wave_len\n    if start < 0 or end > wave_len:\n        raise ValueError(f\"Start or end beyond the wave length: start={start} end={end} len={wave_len}\")\n    waveform_part = wave[start:end]\n    if not nocache:\n        cache_splitted[key] = (waveform_part, rate)\n    return (waveform_part, rate)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T22:56:16.895180Z","iopub.execute_input":"2024-12-15T22:56:16.895453Z","iopub.status.idle":"2024-12-15T22:56:16.905276Z","shell.execute_reply.started":"2024-12-15T22:56:16.895428Z","shell.execute_reply":"2024-12-15T22:56:16.904440Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Загрузчик данных\n\n## Класс CustomDataset\n\nПо индексу считывает аудиоданные и преобразует их в изображение, которое принимает на вход модель для обучения. Также используется кэширование, чтобы ускорить предобработку данных\n\n## Функция melspectrogram\n\nВычисляет mel-спекторграмму по полученным данным\n\n## Функция to_image\n\nПреобразует mel-спекторграмму в rgb изображение необходимого размера\n","metadata":{}},{"cell_type":"code","source":"def melspectrogram(w, r, fmin, fmax):\n    return librosa.power_to_db(\n        librosa.feature.melspectrogram(\n            y=w, sr=r, n_mels=128, fmin=fmin, fmax=fmax\n        ), top_db=80,\n    ).astype(np.float32)\n\n\ndef to_image(X):\n    X = resize(X, (224, 400))\n    eps = 1e-6\n    mean = X.mean()\n    std = X.std()\n    X = (X - mean) / (std + eps)\n    _min, _max = X.min(), X.max()\n    if (_max - _min) > eps:\n        V = np.clip(X, _min, _max)\n        V = 255 * (V - _min) / (_max - _min)\n        V = V.astype(np.uint8)\n    else:\n        V = np.zeros_like(X, dtype=np.uint8)\n    return np.stack([V.astype(np.float32) / 255.0] * 3)\n\n\nclass TrainDataset(Dataset):\n    def __init__(self, data, audio_dir, fmin, fmax, duration: int = 10):\n        self.data = data\n        self.audio_dir = audio_dir\n        self.fmin = fmin\n        self.fmax = fmax\n        self.duration = duration\n        self.cache = {}\n    \n    def __len__(self):\n        return len(self.data)\n    \n    def audio_to_image(self, audio, sr):\n        return to_image(\n            melspectrogram(audio, sr, self.fmin, self.fmax)\n        )\n    \n    def __getitem__(self, idx):\n        if idx in self.cache:\n            return self.cache[idx]\n\n        row_data = self.data.iloc[idx]\n        fileid = row_data[\"recording_id\"]\n        tmin = row_data[\"t_min\"]\n        tmax = row_data[\"t_max\"]\n        s_id = row_data[\"species_id\"]\n        wave, sr = get_audio_part(self.audio_dir, fileid, tmin, tmax, self.duration)\n        self.cache[idx] = (self.audio_to_image(wave, sr), s_id)\n        return self.cache[idx]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T22:56:16.906265Z","iopub.execute_input":"2024-12-15T22:56:16.906557Z","iopub.status.idle":"2024-12-15T22:56:16.963025Z","shell.execute_reply.started":"2024-12-15T22:56:16.906532Z","shell.execute_reply":"2024-12-15T22:56:16.962350Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(\"../input/rfcx-species-audio-detection/train_tp.csv\")\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T22:56:16.964400Z","iopub.execute_input":"2024-12-15T22:56:16.964639Z","iopub.status.idle":"2024-12-15T22:56:16.987357Z","shell.execute_reply.started":"2024-12-15T22:56:16.964616Z","shell.execute_reply":"2024-12-15T22:56:16.986552Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Глобальные параметры\n\n* `num_labels` -- количество классов. Вычисляется на основе размеченных данных\n* `fmin` -- минимальная частота для mel-спектрограммы. Берётся минимальная из представленных, в `train_tp.csv` файле для обучения\n* `fmax` -- максимальная частота для mel-спектрограммы. Верхняя граница взята эмпирическим путём, что получить достаточно данных для mel-спектрограммы при большом количестве mel-фильтров\n* `device` -- устройство для запуска модели (CPU или GPU)\n* `FOLD_NUM` -- количество разбиений для кросс-валидации","metadata":{}},{"cell_type":"code","source":"num_labels = len(df[\"species_id\"].unique())\nfmin = df[\"f_min\"].min() * 0.9\nfmax = 24000\ndevice = \"cuda:0\" if torch.cuda.is_available() else \"cpu\"\nFOLD_NUM = 5\nprint(num_labels, fmin, fmax, device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T22:56:16.988286Z","iopub.execute_input":"2024-12-15T22:56:16.988530Z","iopub.status.idle":"2024-12-15T22:56:16.994005Z","shell.execute_reply.started":"2024-12-15T22:56:16.988506Z","shell.execute_reply":"2024-12-15T22:56:16.993105Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Проверим работу созданного загрузчика данных и посмотрим на изображения. Получаются трехканальные нормализованные изображения 224x440. Видно наличие пения птиц, а также общий шумовой фон леса","metadata":{}},{"cell_type":"code","source":"train_data = TrainDataset(df, train_audio_dir, fmin, fmax)\ntest_item = train_data[14][0]\nprint(test_item.shape)\n\ncount = 9\nfig, axs = plt.subplots(count, 1, figsize=(10, 20))\nfor i in range(count):\n    axs[i].imshow(np.transpose(train_data[1 + i][0], axes=[1, 2, 0]))\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T22:56:16.995139Z","iopub.execute_input":"2024-12-15T22:56:16.995472Z","iopub.status.idle":"2024-12-15T22:56:20.080619Z","shell.execute_reply.started":"2024-12-15T22:56:16.995436Z","shell.execute_reply":"2024-12-15T22:56:20.079782Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Функция get_model\n\nЗагружает чистую модель resnet101 и меняет ей \"наконечник\" на полносвязный слой, который выдаёт требуемое количество классов","metadata":{}},{"cell_type":"code","source":"def get_model(num_labels, device):\n    model = resnet101(pretrained=True)\n    num_ftrs = model.fc.in_features\n    model.fc = torch.nn.Linear(num_ftrs, num_labels)\n    model = model.to(device)\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T22:56:20.081615Z","iopub.execute_input":"2024-12-15T22:56:20.081891Z","iopub.status.idle":"2024-12-15T22:56:20.086358Z","shell.execute_reply.started":"2024-12-15T22:56:20.081865Z","shell.execute_reply":"2024-12-15T22:56:20.085453Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Функция train_model\n\nОбычное обучение модели в рамках одной итерации кросс-валидации. По результатам обучения сохраняется и возвращается наилучшая модель\n\n# Функция train_with_kfold\n\nПроводит кросс-валидацию и сохраняет модель в файловую систему для последующего использования","metadata":{}},{"cell_type":"code","source":"def train_model(model, loss_fn, train_loader, valid_loader, epochs, optimizer, scheduler):\n    \n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_acc = 0.0\n    train_losses = []\n    valid_losses = []\n    \n    for epoch in tqdm(range(epochs)):\n        model.train()\n        batch_losses=[]\n        for x, y in train_loader:\n            optimizer.zero_grad()\n            x = x.to(device, dtype=torch.float32)\n            y = y.to(device, dtype=torch.long)\n            y_hat = model(x)\n            loss = loss_fn(y_hat, y)\n            loss.backward()\n            batch_losses.append(loss.item())\n            optimizer.step()\n        train_losses.append(batch_losses)\n\n        model.eval()\n        batch_losses=[]\n        trace_y = []\n        trace_yhat = []\n        \n        for x, y in valid_loader:\n            x = x.to(device, dtype=torch.float32)\n            y = y.to(device, dtype=torch.long)\n            y_hat = model(x)\n            loss = loss_fn(y_hat, y)\n            trace_y.append(y.cpu().detach().numpy())\n            trace_yhat.append(y_hat.cpu().detach().numpy())      \n            batch_losses.append(loss.item())\n        valid_losses.append(batch_losses)\n        trace_y = np.concatenate(trace_y)\n        trace_yhat = np.concatenate(trace_yhat)\n        accuracy = np.mean(trace_yhat.argmax(axis=1)==trace_y)\n        \n        print(f\"epoch = {epoch}, train_loss = {np.mean(train_losses[-1]):.5f}, val_loss = {np.mean(valid_losses[-1]):.5f}, val_accuracy = {accuracy:.5f}\")\n\n        scheduler.step(np.mean(valid_losses[-1]))\n        if accuracy > best_acc:\n            best_acc = accuracy\n            best_model_wts = copy.deepcopy(model.state_dict())\n\n    model.load_state_dict(best_model_wts)\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T23:04:51.191682Z","iopub.execute_input":"2024-12-15T23:04:51.192088Z","iopub.status.idle":"2024-12-15T23:04:51.201893Z","shell.execute_reply.started":"2024-12-15T23:04:51.192044Z","shell.execute_reply":"2024-12-15T23:04:51.201146Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_with_kfold(k: int):\n    kfold = KFold(n_splits=k, shuffle=True, random_state=42)\n    \n    for fold_id, (train_index, val_index) in enumerate(kfold.split(df, df[\"species_id\"].tolist())):\n        print(\"Fold\", fold_id)\n    \n        learning_rate = 2e-4\n        epochs = 20\n        loss_fn = torch.nn.CrossEntropyLoss()\n    \n        train_data = TrainDataset(df.loc[train_index], train_audio_dir, fmin, fmax)\n        valid_data = TrainDataset(df.loc[val_index], train_audio_dir, fmin, fmax)\n        train_loader = DataLoader(train_data, batch_size=8, shuffle=True, drop_last=True)\n        valid_loader = DataLoader(valid_data, batch_size=8, shuffle=True, drop_last=True)\n        model = get_model(num_labels, device)\n        optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate)\n        scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=3)\n        model = train_model(model, loss_fn, train_loader, valid_loader, epochs, optimizer, scheduler)\n        torch.save(model.state_dict(), f\"./model{fold_id}.pt\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T23:04:51.663706Z","iopub.execute_input":"2024-12-15T23:04:51.665134Z","iopub.status.idle":"2024-12-15T23:04:51.674037Z","shell.execute_reply.started":"2024-12-15T23:04:51.665078Z","shell.execute_reply":"2024-12-15T23:04:51.672863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_with_kfold(FOLD_NUM)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T23:04:55.750749Z","iopub.execute_input":"2024-12-15T23:04:55.751193Z","iopub.status.idle":"2024-12-15T23:13:22.841045Z","shell.execute_reply.started":"2024-12-15T23:04:55.751153Z","shell.execute_reply":"2024-12-15T23:13:22.840096Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Очистка памяти\n\nНа kaggle может не хватать памяти для предсказания тестовых данных, поэтому необходимо освободить кэши и очистить все возможные объекты из памяти","metadata":{}},{"cell_type":"code","source":"def clear_mem():\n    import gc\n\n    # Reset cache to empty state\n    clear_cache()\n    # Also drop cache from gpu\n    torch.cuda.empty_cache()\n    # Manually trigger garbage collection\n    collected = gc.collect()\n    # Verify memory release\n    print(f\"Garbage collector collected {collected} objects.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T23:14:30.212628Z","iopub.execute_input":"2024-12-15T23:14:30.213062Z","iopub.status.idle":"2024-12-15T23:14:30.218507Z","shell.execute_reply.started":"2024-12-15T23:14:30.213020Z","shell.execute_reply":"2024-12-15T23:14:30.217502Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"clear_mem()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T23:14:30.220220Z","iopub.execute_input":"2024-12-15T23:14:30.220610Z","iopub.status.idle":"2024-12-15T23:14:30.636538Z","shell.execute_reply.started":"2024-12-15T23:14:30.220567Z","shell.execute_reply":"2024-12-15T23:14:30.635667Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Загружаем все полученные модели с кросс-валидации","metadata":{}},{"cell_type":"code","source":"models = []\nfor i in range(FOLD_NUM):\n    model = get_model(num_labels, device)\n    model.load_state_dict(torch.load(f\"./model{i}.pt\"))\n    models.append(model.eval())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T23:14:30.639305Z","iopub.execute_input":"2024-12-15T23:14:30.639865Z","iopub.status.idle":"2024-12-15T23:14:35.671128Z","shell.execute_reply.started":"2024-12-15T23:14:30.639824Z","shell.execute_reply":"2024-12-15T23:14:35.670420Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Загрузчик данных (ещё один)\n\n## Класс TestDataset\n\nАналогичен загрузчику `TrainDataset` за исключением двух моментов:\n* Содержит опцию для отключения кэша\n* Возвращает не отдельный временной интервал и метку, а имя файла и набор mel-спектрограмм для этого файла","metadata":{}},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, data, audio_dir, fmin, fmax, duration: int = 10, nocache: bool = False):\n        self.data = data\n        self.audio_dir = audio_dir\n        self.fmin = fmin\n        self.fmax = fmax\n        self.duration = duration\n        self.cache = {}\n        self.nocache = nocache\n    \n    def __len__(self):\n        return len(self.data)\n    \n    def audio_to_image(self, audio, sr):\n        return to_image(\n            melspectrogram(audio, sr, self.fmin, self.fmax)\n        )\n    \n    def __getitem__(self, idx):\n        if idx in self.cache:\n            return self.cache[idx]\n\n        fileid = self.data[idx]\n        fileid = fileid[:fileid.rfind(\".flac\")]\n        wave, rate = get_audio(self.audio_dir, fileid, self.nocache)\n        wave_len = len(wave) / rate\n        output = []\n        for i in range(int(np.ceil(wave_len / self.duration))):\n            tmin = i * self.duration\n            tmax = tmin + self.duration\n            wave, sr = get_audio_part(self.audio_dir, fileid, tmin, tmax, self.duration, self.nocache)\n            output.append(self.audio_to_image(wave, sr))\n        if not self.nocache:\n            self.cache[idx] = (fileid, np.array(output))\n            return self.cache[idx]\n        return (fileid, np.array(output))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T23:14:35.672251Z","iopub.execute_input":"2024-12-15T23:14:35.672592Z","iopub.status.idle":"2024-12-15T23:14:35.681098Z","shell.execute_reply.started":"2024-12-15T23:14:35.672552Z","shell.execute_reply":"2024-12-15T23:14:35.680160Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Функция predict\n\nПроходит по всем файлам в папке с тестовым аудиофайлами и высчитывает вероятность нахождения в данном файле звука определенной птицы. Используются сразу все несколько полученных моделей, их результат комбинируется при помощи среднего между всеми моделями\n\n## Функция make_submission\n\nЗаписывает результаты в файл `submission.csv` согласно требуемому шаблону","metadata":{}},{"cell_type":"code","source":"def predict(models, device):\n    test_dataset = TestDataset(os.listdir(test_audio_dir), test_audio_dir, fmin, fmax, nocache=True)\n    res_indexes = []\n    res_probabilities = []\n    with torch.no_grad():\n        for i in tqdm(range(0, len(test_dataset))):\n            fileid, data = test_dataset[i]\n            data = torch.tensor(data).float()\n            if device == \"cuda:0\":\n                data = data.cuda()\n    \n            pred = [torch.max(mdl(data), dim=0)[0].cpu().detach() for mdl in models]\n            avg_pred = torch.mean(torch.stack(pred), dim=0)\n            res_indexes.append(fileid)\n            res_probabilities.append([avg.item() for avg in avg_pred])\n    return res_indexes, res_probabilities","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T23:14:35.682501Z","iopub.execute_input":"2024-12-15T23:14:35.682747Z","iopub.status.idle":"2024-12-15T23:14:35.700267Z","shell.execute_reply.started":"2024-12-15T23:14:35.682722Z","shell.execute_reply":"2024-12-15T23:14:35.699495Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_submission(preds):\n    res_indexes, res_probabilities = preds\n    df_sub = pd.DataFrame(res_probabilities,\n                          columns=[f\"s{i}\" for i in range(num_labels)])\n    df_sub[\"recording_id\"] = res_indexes\n    df_sub.to_csv(\"submission.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T23:14:35.701074Z","iopub.execute_input":"2024-12-15T23:14:35.701287Z","iopub.status.idle":"2024-12-15T23:14:35.709330Z","shell.execute_reply.started":"2024-12-15T23:14:35.701266Z","shell.execute_reply":"2024-12-15T23:14:35.708636Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"make_submission(predict(models, device))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-15T23:14:35.710305Z","iopub.execute_input":"2024-12-15T23:14:35.710614Z","iopub.status.idle":"2024-12-15T23:14:40.151514Z","shell.execute_reply.started":"2024-12-15T23:14:35.710578Z","shell.execute_reply":"2024-12-15T23:14:40.150651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}