{"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":"gpu","dataSources":[{"sourceId":25954,"databundleVersionId":2091745,"sourceType":"competition"}],"dockerImageVersionId":30786,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Data preparation","metadata":{}},{"cell_type":"code","source":"import os\nfrom pathlib import Path\n\nhome_dir = Path('/kaggle/input/birdclef-2021')\nwork_dir = Path('/kaggle/working')\nprint(*os.listdir(home_dir), sep='\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T11:54:06.039715Z","iopub.execute_input":"2024-11-19T11:54:06.040412Z","iopub.status.idle":"2024-11-19T11:54:06.053041Z","shell.execute_reply.started":"2024-11-19T11:54:06.040375Z","shell.execute_reply":"2024-11-19T11:54:06.052078Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Вся обработка будет предполагать, что дорожка имеет следующую частоту дискретизации","metadata":{}},{"cell_type":"code","source":"sr = 32000          # 32 kHz\nn_mfcc = 40         # 40 mel coeffs\nn_fft = 1024        # for melspec\nmax_sample_dur = 5  # sec","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T11:54:07.334955Z","iopub.execute_input":"2024-11-19T11:54:07.335283Z","iopub.status.idle":"2024-11-19T11:54:07.339388Z","shell.execute_reply.started":"2024-11-19T11:54:07.335253Z","shell.execute_reply":"2024-11-19T11:54:07.338544Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Данные можно разделить на 2 типа:\n1. Короткие адиозаписи, каждая лишь с одной поющей птицей. Тестовых данных нет\n2. Записи по 10 минут, в которых поёт множество птиц. Тестовые есть, но скрыты\n\nВ связи с этими различиями, предлагается разделить обучение на 2 этапа: \n1. Обучение распознаванию одной птицы на коротких записях\n2. Обучение распознаванию множества птиц на длинных записях","metadata":{}},{"cell_type":"markdown","source":" ## Short audio data","metadata":{}},{"cell_type":"code","source":"train_short_dir = home_dir / 'train_short_audio'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T11:54:08.649165Z","iopub.execute_input":"2024-11-19T11:54:08.649509Z","iopub.status.idle":"2024-11-19T11:54:08.653406Z","shell.execute_reply.started":"2024-11-19T11:54:08.649465Z","shell.execute_reply":"2024-11-19T11:54:08.652619Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(*os.listdir(train_short_dir), sep=' ')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T11:54:09.033656Z","iopub.execute_input":"2024-11-19T11:54:09.034446Z","iopub.status.idle":"2024-11-19T11:54:09.055964Z","shell.execute_reply.started":"2024-11-19T11:54:09.034414Z","shell.execute_reply":"2024-11-19T11:54:09.055009Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Длительность записей может превосходить минуту -- для моего подхода это много. Нужно будет нарезать запись на перекрывающиеся (или нет?) отрезки где-то по секунде, иначе памяти не хватит.","metadata":{}},{"cell_type":"code","source":"import numpy as np\n\ndef cut_into_slices(audio: np.ndarray, duration=5.0, overlap=0.2, at_most=None) -> [np.ndarray]:\n    '''\n    `duration` in seconds  \n    `overlap` in [0, 0.9)\n    '''\n    assert 0.1 < duration < len(audio)/sr\n    assert 0 <= overlap < 0.9\n    \n    duration = round(duration*sr)\n    interval = round(duration * (1-overlap))\n    \n    start_idxs = list(range(0, len(audio)-duration, interval))\n    stop_idxs = list(range(duration, len(audio), interval))\n\n    start_idxs.append(start_idxs[-1]+duration)\n    stop_idxs.append(len(audio))\n    at_most = at_most if (at_most and len(stop_idxs) > at_most) else len(stop_idxs)\n\n    slices = [audio[i:j] for i, j in list(zip(start_idxs, stop_idxs))[:at_most]]\n\n    return slices","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T11:54:09.877778Z","iopub.execute_input":"2024-11-19T11:54:09.878453Z","iopub.status.idle":"2024-11-19T11:54:09.884677Z","shell.execute_reply.started":"2024-11-19T11:54:09.878422Z","shell.execute_reply":"2024-11-19T11:54:09.883699Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import librosa as lr\n\naudio = lr.load(train_short_dir / 'acafly/XC109605.ogg', sr=sr)[0]\nslices = cut_into_slices(audio, 1.0, 0, 5)\nprint(len(slices))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T11:54:10.457546Z","iopub.execute_input":"2024-11-19T11:54:10.457876Z","iopub.status.idle":"2024-11-19T11:54:22.872173Z","shell.execute_reply.started":"2024-11-19T11:54:10.457847Z","shell.execute_reply":"2024-11-19T11:54:22.8712Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calc_stft_powers(audio):\n  ''' Расчёт мощности модуля спектра STFT в децибелах '''\n  n_fft = 2048\n  # проводим оконное преобразование нашего сигнала с заданным окном\n  stft_complex = lr.stft(audio, n_fft=n_fft)\n  stft_mag = np.abs(stft_complex)\n  stft_ang = np.angle(stft_complex)\n  # приводим к логарифмической шкале\n  powers = lr.amplitude_to_db(stft_mag, ref=np.max(np.abs(stft_mag))) # np.min\n  # axis:\n  #  0: freq\n  #  1: time\n  return powers\n\ndef plot_powers(powers):\n  ''' Вывод спектрограммы мощности спектра '''\n  fig, ax = plt.subplots()\n  img = lr.display.specshow(powers, y_axis='log', x_axis='time', ax=ax)\n  ax.set_title('Power spectrogram')\n  fig.colorbar(img, ax=ax, format='%+2.0f dB')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T11:54:22.873968Z","iopub.execute_input":"2024-11-19T11:54:22.874821Z","iopub.status.idle":"2024-11-19T11:54:22.881117Z","shell.execute_reply.started":"2024-11-19T11:54:22.874776Z","shell.execute_reply":"2024-11-19T11:54:22.880177Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def filter_audio(audio, threshold=0.3):\n  ''' Фильтрация шума и тишины '''\n  # Считаем мощность спектра (для каждого окна)\n  powers = calc_stft_powers(audio)\n  # Суммируем мощность по каждой частоте, нормализуем\n  psums = np.sum(powers, axis=0)\n  psums -= np.min(psums)\n  psums /= np.max(psums)\n    \n\n  # Сглаживаем график, чтобы сохранить короткие пропуски и паузы, убрать резкость\n  # для этого применияем свёртку с прямоугольным сигналом\n  psums = np.convolve(psums, np.ones((10,)), 'same')\n  # Нормализуем снова\n  psums -= np.min(psums)\n  psums /= np.max(psums)\n\n  # Устанавливаем порог в treshold*100%\n  clear_sound_bool_vec_window = np.array(psums > threshold, dtype=bool)\n  # имеем clear_sound_bool_vec_window = [1 1 1 0 0 1 1]\n\n  # дублируем расчёты для каждого окна на размер окна,\n  # формируя битовую маску\n  clear_sound_bool_vec = np.repeat(clear_sound_bool_vec_window, np.ceil(len(audio) / len(clear_sound_bool_vec_window)))[:len(audio)]\n  # clear_sound_bool_vec = [1*window 1*window 1*window 0*window 0*window 1*window 1*window]\n\n  # фильтруем\n  return audio[clear_sound_bool_vec]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T11:54:22.882253Z","iopub.execute_input":"2024-11-19T11:54:22.882595Z","iopub.status.idle":"2024-11-19T11:54:22.898677Z","shell.execute_reply.started":"2024-11-19T11:54:22.882567Z","shell.execute_reply":"2024-11-19T11:54:22.898002Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import IPython.display as ipd\nipd.Audio(audio, rate=sr) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T11:54:22.900216Z","iopub.execute_input":"2024-11-19T11:54:22.900475Z","iopub.status.idle":"2024-11-19T11:54:23.02343Z","shell.execute_reply.started":"2024-11-19T11:54:22.90045Z","shell.execute_reply":"2024-11-19T11:54:23.022112Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ipd.Audio(filter_audio(audio, 0.7), rate=sr) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T11:54:23.024517Z","iopub.execute_input":"2024-11-19T11:54:23.024762Z","iopub.status.idle":"2024-11-19T11:54:23.219187Z","shell.execute_reply.started":"2024-11-19T11:54:23.024738Z","shell.execute_reply":"2024-11-19T11:54:23.218249Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"К сожалению, фильтрация оказалась слишком затратной по времени, поэтому мы не будем её использовать","metadata":{}},{"cell_type":"markdown","source":"Короткие аудиозаписи содержат пение лишь одной определённой птицы. Эти записи распиханы по папкам с названиями птиц. Удобно. Используем `DatasetFolder` для формирования датасетов","metadata":{}},{"cell_type":"code","source":"from torchvision.datasets.vision import VisionDataset\nfrom torchvision.datasets.folder import find_classes\nimport torch\n\nimport os.path\nfrom pathlib import Path\nfrom typing import Any, Callable, cast, Dict, List, Optional, Tuple, Union\nfrom sklearn.preprocessing import MultiLabelBinarizer\n\nimport torchaudio\n\nclass DatasetShort(VisionDataset):\n    def __init__(\n        self,\n        root: Union[str, Path],\n        at_most_samples_per_class: Optional[int] = None,\n        transform: Optional[Callable] = None,\n    ) -> None:\n        super().__init__(root, transform=transform)\n        classes, class_to_idx = self.find_classes(self.root)\n        samples, samples_per_classes = self.make_dataset(self.root, class_to_idx, at_most_samples_per_class)\n\n        self.classes = classes\n        self.class_to_idx = class_to_idx\n        self.samples = samples\n        self.samples_per_classes = samples_per_classes\n        self.mlb = MultiLabelBinarizer(classes=list(class_to_idx.values()))\n        self.mlb.fit([])\n        \n        self.device = 'cuda' if torch.cuda.is_available() else 'cpu'\n        self.transform = self.transform.to(self.device)\n        \n\n    @staticmethod\n    def make_dataset(\n        directory: Union[str, Path],\n        class_to_idx: Dict[str, int],\n        at_most_samples_per_class: int,\n    ) -> List[Tuple[str, int]]:\n        if class_to_idx is None:\n            raise ValueError(\"The class_to_idx parameter cannot be None.\")\n\n        instances = []\n        samples_per_classes = [0]*len(class_to_idx.keys())\n        for target_class in sorted(class_to_idx.keys()):\n            class_index = class_to_idx[target_class]\n            target_dir = os.path.join(directory, target_class)\n            if not os.path.isdir(target_dir):\n                continue\n            \n            fnames = sorted(os.listdir(target_dir))\n            if at_most_samples_per_class and len(fnames) >= at_most_samples_per_class:\n                fnames = fnames[:at_most_samples_per_class]\n\n            samples_per_classes[class_to_idx[target_class]] = len(fnames)\n                \n            for fname in fnames:\n                path = os.path.join(target_dir, fname)\n                if '.ogg' in fname:\n                    item = path, class_index\n                    instances.append(item)\n\n        return instances, samples_per_classes\n\n    def find_classes(self, directory: Union[str, Path]) -> Tuple[List[str], Dict[str, int]]:\n        classes = ['nocall', *sorted(entry.name for entry in os.scandir(directory) if entry.is_dir())]\n        class_to_idx = {cls_name: i for i, cls_name in enumerate(classes)}\n        return classes, class_to_idx\n\n    def __getitem__(self, index: int) -> Tuple[Any, Any]:\n        path, target_idx = self.samples[index]\n        audio = lr.load(path, sr=sr)[0]\n        # preprocess sample\n        # sample = filter_audio(sample, 0.7)\n        samples = cut_into_slices(audio, 5, 0.0)\n        sample = samples[np.random.choice(len(samples))]\n        sample = lr.util.fix_length(sample, size=max_sample_dur*sr)\n\n        # transform with device\n        if self.transform is not None:\n            sample = torch.tensor(sample).to(self.device)\n            sample = self.transform(sample)\n            sample = (sample - sample.mean()) / (sample.std() + 1e-6)\n            sample = (sample -  sample.min()) / (sample.max() - sample.min())\n            \n            sample = torch.stack([sample, sample, sample], 0)\n\n        # sample = sample.unsqueeze(0)\n        \n        min_t, max_t = 0.0025, 0.995\n        target = self.mlb.transform([[target_idx]])[0]# * (max_t-min_t) + min_t\n            \n        return sample, target\n\n    def __len__(self) -> int:\n        return len(self.samples)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:40:58.885331Z","iopub.execute_input":"2024-11-19T13:40:58.885691Z","iopub.status.idle":"2024-11-19T13:40:58.900866Z","shell.execute_reply.started":"2024-11-19T13:40:58.88566Z","shell.execute_reply":"2024-11-19T13:40:58.899707Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def mono_to_color(mono, eps=1e-6, mean=None, std=None):\n    mean = mean or mono.mean()\n    std = std or mono.std()\n    mono = (mono - mean) / (std + eps)\n    \n    _min, _max = mono.min(), mono.max()\n\n    if (_max - _min) > eps:\n        color = np.clip(mono, _min, _max)\n        color = 255 * (color - _min) / (_max - _min)\n        color = color.astype(np.uint8)\n    else:\n        color = np.zeros_like(color, dtype=np.uint8)\n\n    return color","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T12:26:17.950016Z","iopub.execute_input":"2024-11-19T12:26:17.950369Z","iopub.status.idle":"2024-11-19T12:26:17.956484Z","shell.execute_reply.started":"2024-11-19T12:26:17.950339Z","shell.execute_reply":"2024-11-19T12:26:17.955566Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndef test_melspec(filter_sample=False):\n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    print(device)\n    sample = audio\n    # preprocess sample\n    if (filter_sample):\n        sample = filter_audio(sample, 0.6)\n    print(f'sample duration = {len(sample) / sr} s')\n    sample = lr.util.fix_length(sample, size=max_sample_dur*sr)\n    sample = torch.tensor(sample).to(device)\n    \n    n_ffts = [256, 512, 1024, 2048, 4096]\n    fig, axs = plt.subplots(nrows=len(n_ffts), ncols=1, figsize=(10, 3*len(n_ffts)))\n    \n    for i, n_fft in enumerate(n_ffts):\n        melspec_eveluator = torchaudio.transforms.MelSpectrogram(\n            sample_rate=sr,\n            n_fft=n_fft,\n            n_mels=128, # птички высокочастотные -- дальше мы это увидим\n            hop_length=n_fft//4,\n            normalized=True,\n        ).to(device)\n        amp_to_db = torchaudio.transforms.AmplitudeToDB().to(device)\n        \n        melspec = amp_to_db(melspec_eveluator(sample)).cpu().detach().numpy()\n        melspec = (melspec - melspec.mean()) / (melspec.std() + 1e-6)\n        melspec = (melspec -  melspec.min()) / (melspec.max() - melspec.min())\n        # melspec = mono_to_color(melspec)\n        # print(np.min(melspec), np.max(melspec))\n        # melspec = lr.power_to_db(melspec)\n\n        axs[i].set_title(f'n_fft = {n_fft}')\n        axs[i].set_ylabel('mel')\n        axs[i].set_xlabel('window num')\n        axs[i].imshow(melspec, origin=\"lower\", aspect=\"auto\", interpolation=\"nearest\")\n    fig.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T12:39:10.897632Z","iopub.execute_input":"2024-11-19T12:39:10.897971Z","iopub.status.idle":"2024-11-19T12:39:10.906467Z","shell.execute_reply.started":"2024-11-19T12:39:10.897943Z","shell.execute_reply":"2024-11-19T12:39:10.905516Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_melspec(filter_sample=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T12:39:12.477685Z","iopub.execute_input":"2024-11-19T12:39:12.478523Z","iopub.status.idle":"2024-11-19T12:39:13.661851Z","shell.execute_reply.started":"2024-11-19T12:39:12.478456Z","shell.execute_reply":"2024-11-19T12:39:13.660976Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"А так выглядит мел-спек фильтрованной записи","metadata":{}},{"cell_type":"code","source":"test_melspec(filter_sample=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T12:39:17.59104Z","iopub.execute_input":"2024-11-19T12:39:17.591403Z","iopub.status.idle":"2024-11-19T12:39:18.878631Z","shell.execute_reply.started":"2024-11-19T12:39:17.591371Z","shell.execute_reply":"2024-11-19T12:39:18.877815Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Шумненько. И там, и там. Смысла использовать фильтрацию мало -- да, мы всё же немного уменьшаем количество бесполезных признаков, но слишком затратно по времени. \n\nВидим, что оптимально n_fft=1024","metadata":{}},{"cell_type":"code","source":"n_fft = 1024","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T12:04:08.269445Z","iopub.execute_input":"2024-11-19T12:04:08.270273Z","iopub.status.idle":"2024-11-19T12:04:08.273894Z","shell.execute_reply.started":"2024-11-19T12:04:08.270243Z","shell.execute_reply":"2024-11-19T12:04:08.273029Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torchaudio\nimport torch.nn as nn\n\nmelspec_eveluator = nn.Sequential(\n    torchaudio.transforms.MelSpectrogram(\n        sample_rate=sr,\n        n_fft=n_fft,\n        n_mels=128, # птички высокочастотные -- как уже увидели\n        hop_length=n_fft//4,\n        normalized=True,\n    ),\n    torchaudio.transforms.AmplitudeToDB()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:41:18.423113Z","iopub.execute_input":"2024-11-19T13:41:18.423479Z","iopub.status.idle":"2024-11-19T13:41:18.434361Z","shell.execute_reply.started":"2024-11-19T13:41:18.42345Z","shell.execute_reply":"2024-11-19T13:41:18.433676Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_val_short_set = DatasetShort(\n    root=train_short_dir,\n    transform=melspec_eveluator,\n    at_most_samples_per_class=100\n) \n\nclasses = train_val_short_set.classes\ncls_to_idx = train_val_short_set.class_to_idx\nsamples_per_classes = train_val_short_set.samples_per_classes","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:41:22.420889Z","iopub.execute_input":"2024-11-19T13:41:22.421543Z","iopub.status.idle":"2024-11-19T13:41:23.029443Z","shell.execute_reply.started":"2024-11-19T13:41:22.421487Z","shell.execute_reply":"2024-11-19T13:41:23.028715Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"melspec, label = train_val_short_set[0]\nprint(melspec.shape, melspec.min(), melspec.max())\nprint(label.shape, label.min(), label.max())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:41:37.204007Z","iopub.execute_input":"2024-11-19T13:41:37.204384Z","iopub.status.idle":"2024-11-19T13:41:37.284509Z","shell.execute_reply.started":"2024-11-19T13:41:37.204343Z","shell.execute_reply":"2024-11-19T13:41:37.283554Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(melspec.cpu().detach().numpy()[0], origin=\"lower\", aspect=\"auto\", interpolation=\"nearest\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:41:39.038678Z","iopub.execute_input":"2024-11-19T13:41:39.039018Z","iopub.status.idle":"2024-11-19T13:41:39.293683Z","shell.execute_reply.started":"2024-11-19T13:41:39.038988Z","shell.execute_reply":"2024-11-19T13:41:39.292798Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.utils.class_weight import compute_class_weight\n\nnums = [[cls]*num for cls, num in enumerate(samples_per_classes)]\n\ny = [0]\nfor num in nums:\n    y = [*y, *num]\nweights = compute_class_weight(\n    'balanced', \n    classes=list(range(len(classes))), \n    y=y\n)\nweights","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:25:07.833483Z","iopub.execute_input":"2024-11-19T13:25:07.833853Z","iopub.status.idle":"2024-11-19T13:25:07.872057Z","shell.execute_reply.started":"2024-11-19T13:25:07.833823Z","shell.execute_reply":"2024-11-19T13:25:07.871322Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mlb = MultiLabelBinarizer(classes=list(cls_to_idx.values()))\nmlb.fit([])\n# mlb.transform([[0]])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:25:11.674818Z","iopub.execute_input":"2024-11-19T13:25:11.675595Z","iopub.status.idle":"2024-11-19T13:25:11.686469Z","shell.execute_reply.started":"2024-11-19T13:25:11.67556Z","shell.execute_reply":"2024-11-19T13:25:11.685547Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import random_split\n\ntrain_short_set, val_short_set = random_split(\n    train_val_short_set,\n    [0.8, 0.2],\n)\n\nprint(len(train_short_set), len(val_short_set))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:41:48.245574Z","iopub.execute_input":"2024-11-19T13:41:48.246144Z","iopub.status.idle":"2024-11-19T13:41:48.260173Z","shell.execute_reply.started":"2024-11-19T13:41:48.246098Z","shell.execute_reply":"2024-11-19T13:41:48.258475Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch_size = 32","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:42:01.272965Z","iopub.execute_input":"2024-11-19T13:42:01.273267Z","iopub.status.idle":"2024-11-19T13:42:01.277601Z","shell.execute_reply.started":"2024-11-19T13:42:01.273241Z","shell.execute_reply":"2024-11-19T13:42:01.276576Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\ntrain_loader = DataLoader(\n    train_val_short_set,\n    batch_size=batch_size,\n    shuffle=True,\n    num_workers=0\n)\n\ntrain_short_loader = DataLoader(\n    train_short_set,\n    batch_size=batch_size,\n    shuffle=True,\n    num_workers=0\n)\nvalid_short_loader = DataLoader(\n    val_short_set,\n    batch_size=batch_size,\n    shuffle=True,\n    num_workers=0\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:42:02.42679Z","iopub.execute_input":"2024-11-19T13:42:02.427446Z","iopub.status.idle":"2024-11-19T13:42:02.432203Z","shell.execute_reply.started":"2024-11-19T13:42:02.427413Z","shell.execute_reply":"2024-11-19T13:42:02.431338Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10 mins data","metadata":{}},{"cell_type":"code","source":"train_10mins_dir = home_dir / 'train_soundscapes'\ntest_10mins_dir = home_dir / 'test_soundscapes'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T12:05:59.229164Z","iopub.execute_input":"2024-11-19T12:05:59.229506Z","iopub.status.idle":"2024-11-19T12:05:59.233626Z","shell.execute_reply.started":"2024-11-19T12:05:59.229466Z","shell.execute_reply":"2024-11-19T12:05:59.232715Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(*os.listdir(train_10mins_dir), sep='\\n', end='\\n\\n')\nprint(*os.listdir(test_10mins_dir), sep='\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T12:06:02.937975Z","iopub.execute_input":"2024-11-19T12:06:02.93831Z","iopub.status.idle":"2024-11-19T12:06:02.962886Z","shell.execute_reply.started":"2024-11-19T12:06:02.938275Z","shell.execute_reply":"2024-11-19T12:06:02.962028Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nfrom pandas import DataFrame\n\ndf_10mins_train = pd.read_csv(home_dir / 'train_soundscape_labels.csv')\ndf_10mins_train","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:25:32.666706Z","iopub.execute_input":"2024-11-19T13:25:32.667041Z","iopub.status.idle":"2024-11-19T13:25:32.684723Z","shell.execute_reply.started":"2024-11-19T13:25:32.667011Z","shell.execute_reply":"2024-11-19T13:25:32.683843Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Нарежем наши 10 минутные записи на куски по 5 секунд, как этого трубуют тренировочные данные.\n\nАнализировать мы будем эти куски окнами с перекрытием, поэтому есть доп параметры для размера окна в секундах и перекрытия в долях.","metadata":{}},{"cell_type":"markdown","source":"Ячейки в столбце `birds` -- птицы, которые пели в текущем 5 секундном промежутке, разделённые пробелом. \n\nЭто multilabel классификация -- используем [MultiLabelBinarizer](https://scikit-learn.org/dev/modules/generated/sklearn.preprocessing.MultiLabelBinarizer.html)","metadata":{}},{"cell_type":"code","source":"import torch\n# from torch.utils.data import Dataset\nfrom torchvision.datasets.vision import VisionDataset\nfrom sklearn.preprocessing import MultiLabelBinarizer\nimport os\nimport pandas as pd\n\ndef to_tag(name: str):\n    return '_'.join(name.split('_')[:2])\n\ndef get_tag_to_audio(dir: Path):\n    names = os.listdir(dir)\n    \n    audios = [lr.load(dir / name, sr=sr)[0] for name in names if '.ogg' in name]\n    audios = [cut_into_slices(audio, 5.0, 0) for audio in audios]\n\n    tags = [to_tag(name) for name in names]\n    tag_to_audio = dict(zip(tags, audios))\n    return tag_to_audio\n\nclass DatasetCsv10mins(VisionDataset):\n    def __init__(self, csv_file, ogg_dir, classes, transform=None):\n        super().__init__(ogg_dir, transform=transform)\n        self.df = pd.read_csv(csv_file)\n        self.classes = classes\n        self.cls_to_idx = dict(zip(self.classes, list(range(len(classes)))))\n        \n        self.tag_to_audio = get_tag_to_audio(ogg_dir)\n        \n        self.mlb = MultiLabelBinarizer(classes=list(self.cls_to_idx.values()))\n        self.mlb.fit([])\n\n        self.device = 'cuda' if torch.cuda.is_available() else 'cpu'\n        self.transform = self.transform.to(self.device)\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        item = self.df.loc[idx]\n\n        row_id = item.row_id\n        sample = self.tag_to_audio[to_tag(row_id)][item.seconds//5 - 1]\n        if self.transform is not None:\n            sample = torch.tensor(sample).to(self.device)\n            sample = self.transform(sample)\n            sample = (sample - sample.mean()) / (sample.std() + 1e-6)\n            sample = (sample -  sample.min()) / (sample.max() - sample.min())\n            sample = torch.stack([sample, sample, sample], 0)\n            \n        # sample = sample.unsqueeze(0)\n        \n        clss = [self.cls_to_idx[cls] for cls in item.birds.split()]\n        min_t, max_t = 0.0025, 0.995\n        target = self.mlb.transform([clss])[0] # * (max_t-min_t) + min_t\n\n        return sample, target\n\nclass DatasetCsv10minsTest(DatasetCsv10mins):\n    # def __init__(self, csv_file, ogg_dir, classes):\n    #     super().__init__(csv_file, ogg_dir, classes)\n    def __getitem__(self, idx):\n        item = self.df.loc[idx]\n\n        row_id = item.row_id\n        \n        sample = self.tag_to_audio[to_tag(row_id)][item.seconds//5 - 1]\n        if self.transform is not None:\n            sample = torch.tensor(sample).to(self.device)\n            sample = self.transform(sample)\n            sample = (sample - sample.mean()) / (sample.std() + 1e-6)\n            sample = (sample -  sample.min()) / (sample.max() - sample.min())\n            sample = torch.stack([sample, sample, sample], 0)\n\n        # sample = sample.unsqueeze(0)\n\n        clss = [self.cls_to_idx[cls] for cls in item.birds.split()]\n        min_t, max_t = 0.0025, 0.995\n        target = self.mlb.transform([clss])[0]# * (max_t-min_t) + min_t\n\n        return row_id, sample, target\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:43:00.756208Z","iopub.execute_input":"2024-11-19T13:43:00.757118Z","iopub.status.idle":"2024-11-19T13:43:00.769886Z","shell.execute_reply.started":"2024-11-19T13:43:00.757081Z","shell.execute_reply":"2024-11-19T13:43:00.769108Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_val_10mins_set = DatasetCsv10mins(\n    home_dir / 'train_soundscape_labels.csv', \n    train_10mins_dir, \n    classes,\n    transform=melspec_eveluator,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:43:03.196674Z","iopub.execute_input":"2024-11-19T13:43:03.197003Z","iopub.status.idle":"2024-11-19T13:43:16.30026Z","shell.execute_reply.started":"2024-11-19T13:43:03.196975Z","shell.execute_reply":"2024-11-19T13:43:16.299546Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"melspec, target = train_val_10mins_set[20]\nprint(melspec.shape, melspec.min(), melspec.max())\nprint(target.shape, target.min(), target.max())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:43:19.917535Z","iopub.execute_input":"2024-11-19T13:43:19.917895Z","iopub.status.idle":"2024-11-19T13:43:19.9289Z","shell.execute_reply.started":"2024-11-19T13:43:19.917863Z","shell.execute_reply":"2024-11-19T13:43:19.927772Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(melspec.cpu().detach().numpy()[0], origin=\"lower\", aspect=\"auto\", interpolation=\"nearest\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:43:24.674248Z","iopub.execute_input":"2024-11-19T13:43:24.674598Z","iopub.status.idle":"2024-11-19T13:43:24.942481Z","shell.execute_reply.started":"2024-11-19T13:43:24.674568Z","shell.execute_reply":"2024-11-19T13:43:24.941658Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_10mins_set = DatasetCsv10minsTest(\n    home_dir / 'train_soundscape_labels.csv', \n    train_10mins_dir, \n    classes,\n    transform=melspec_eveluator,\n)\n\nrow_id, sample, target = test_10mins_set[0]\nprint(row_id)\nprint(sample.shape)\nprint(target.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:43:29.007837Z","iopub.execute_input":"2024-11-19T13:43:29.008142Z","iopub.status.idle":"2024-11-19T13:43:42.216875Z","shell.execute_reply.started":"2024-11-19T13:43:29.008117Z","shell.execute_reply":"2024-11-19T13:43:42.215815Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import random_split\n\ntrain_10mins_set, val_10mins_set = random_split(\n    train_val_10mins_set,\n    [0.8, 0.2],\n)\n\nprint(len(train_10mins_set), len(val_10mins_set))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:43:42.218618Z","iopub.execute_input":"2024-11-19T13:43:42.218924Z","iopub.status.idle":"2024-11-19T13:43:42.22477Z","shell.execute_reply.started":"2024-11-19T13:43:42.218893Z","shell.execute_reply":"2024-11-19T13:43:42.223716Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\nvalid_loader = DataLoader(\n    train_val_10mins_set,\n    batch_size=batch_size,\n    shuffle=True,\n    num_workers=0\n)\n\ntrain_10mins_loader = DataLoader(\n    train_10mins_set,\n    batch_size=batch_size,\n    shuffle=True,\n    num_workers=0\n)\nvalid_10mins_loader = DataLoader(\n    val_10mins_set,\n    batch_size=batch_size,\n    shuffle=True,\n    num_workers=0\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:43:42.226006Z","iopub.execute_input":"2024-11-19T13:43:42.226376Z","iopub.status.idle":"2024-11-19T13:43:42.235229Z","shell.execute_reply.started":"2024-11-19T13:43:42.226333Z","shell.execute_reply":"2024-11-19T13:43:42.234384Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Machine Learning Model","metadata":{}},{"cell_type":"markdown","source":"## EfficintNet\n","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torchvision.ops import Conv2dNormActivation\nfrom torchvision import models\n\nclass ParallelLstm(nn.Module):\n    def __init__(self, num_parallel, window, hidden, num_layers=1):\n        super().__init__()\n        self.device = 'cuda' if torch.cuda.is_available() else 'cpu'\n        self.num_parallel = num_parallel\n        self.lstms = [nn.LSTM(window, hidden, num_layers=num_layers, batch_first=True) for _ in range(self.num_parallel)]\n\n    def forward(self, x):\n        # x.shape = [B, num_parallel, window, S]\n        #           [num_parallel, (B, S, window)]\n        output = []\n        for lstm, inp in zip(self.lstms, x.permute([1, 0, 3, 2])):\n            lstm = lstm.to(self.device)\n            output.append( lstm(inp)[0][:,-1,:] )\n        return torch.cat( output, dim=1 )\n\n\nclass EfficientNet(nn.Module):\n    def __init__(self, cls_num):\n        super().__init__()\n        # self.efficientnet = torch.hub.load('NVIDIA/DeepLearningExamples:torchhub', 'nvidia_efficientnet_b0', pretrained=True)\n        self.efficientnet = models.efficientnet_b2(weights=models.EfficientNet_B2_Weights.IMAGENET1K_V1)\n        # self.efficientnet.stem = Conv2dNormActivation(\n        #     1, 32, kernel_size=3, stride=2, activation_layer=nn.SiLU\n        # )\n        self.efficientnet.classifier = nn.Sequential(\n            # nn.AdaptiveAvgPool2d(1),\n            # nn.Flatten(1),\n            nn.Dropout(p=0.3, inplace=True),\n            nn.Linear(1408, cls_num),\n        )\n        \n        # self.memory = ParallelLstm(1280, 4, 1, 4)\n        # self.classifier = nn.Sequential(\n        #     nn.Flatten(1),\n        #     nn.Dropout(p=0.3, inplace=True),\n        #     nn.Linear(1280, cls_num),\n        # )\n\n    def forward(self, x):\n        x = self.efficientnet(x)\n        # x = self.memory(x)\n        # x = self.classifier(x)\n        return x\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T15:16:36.45015Z","iopub.execute_input":"2024-11-19T15:16:36.450719Z","iopub.status.idle":"2024-11-19T15:16:36.459006Z","shell.execute_reply.started":"2024-11-19T15:16:36.450685Z","shell.execute_reply":"2024-11-19T15:16:36.457985Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# print(EfficientNet(len(classes)).efficientnet.stem, end='\\n\\n')\nprint(EfficientNet(len(classes)).efficientnet.classifier, end='\\n\\n')\n# print(EfficientNet(398).efficientnet._named_members, end='\\n\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T15:16:40.301618Z","iopub.execute_input":"2024-11-19T15:16:40.302265Z","iopub.status.idle":"2024-11-19T15:16:40.524163Z","shell.execute_reply.started":"2024-11-19T15:16:40.302233Z","shell.execute_reply":"2024-11-19T15:16:40.523189Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Сколько требуется памяти для входного окна в 1с?","metadata":{}},{"cell_type":"code","source":"from torchinfo import summary\nsummary(EfficientNet(len(classes)), input_size=(batch_size, 3, 128, 626), device='cuda', dtypes=[torch.float32])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T15:16:43.151801Z","iopub.execute_input":"2024-11-19T15:16:43.152487Z","iopub.status.idle":"2024-11-19T15:16:43.56241Z","shell.execute_reply.started":"2024-11-19T15:16:43.152455Z","shell.execute_reply":"2024-11-19T15:16:43.561547Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Resnet(nn.Module):\n    def __init__(self, cls_num):\n        super().__init__()\n        self.resnet50 = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2) # \n        # self.resnet50.conv1 = nn.Conv2d(1, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n        self.resnet50.fc = nn.Linear(in_features=2048, out_features=cls_num, bias=True)\n        \n        # self.memory = ParallelLstm(1280, 4, 1, 4)\n        # self.classifier = nn.Sequential(\n        #     nn.Flatten(1),\n        #     nn.Dropout(p=0.3, inplace=True),\n        #     nn.Linear(1280, cls_num),\n        # )\n\n    def forward(self, x):\n        x = self.resnet50(x)\n        # x = self.memory(x)\n        # x = self.classifier(x)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:44:10.966482Z","iopub.execute_input":"2024-11-19T13:44:10.966878Z","iopub.status.idle":"2024-11-19T13:44:10.972368Z","shell.execute_reply.started":"2024-11-19T13:44:10.966845Z","shell.execute_reply":"2024-11-19T13:44:10.971417Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# print(Resnet(len(classes)).resnet50._named_members, end='\\n\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T23:30:46.391562Z","iopub.execute_input":"2024-11-18T23:30:46.391923Z","iopub.status.idle":"2024-11-18T23:30:46.396038Z","shell.execute_reply.started":"2024-11-18T23:30:46.391891Z","shell.execute_reply":"2024-11-18T23:30:46.395025Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchinfo import summary\nsummary(Resnet(len(classes)), input_size=(batch_size, 3, 128, 626), device='cuda', dtypes=[torch.float32])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:44:15.474745Z","iopub.execute_input":"2024-11-19T13:44:15.475403Z","iopub.status.idle":"2024-11-19T13:44:16.656179Z","shell.execute_reply.started":"2024-11-19T13:44:15.475371Z","shell.execute_reply":"2024-11-19T13:44:16.655286Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"markdown","source":"## Trainer","metadata":{}},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\n# mlb.transform([[0,1], [2]])\ntarget = torch.tensor(mlb.transform([[0,1], [2]]), dtype=torch.float32, requires_grad=True)\noutput = torch.tensor(mlb.transform([[1,3], [2, 3]]), dtype=torch.float32, requires_grad=True)\nloss = criterion(output, target)\nloss.backward()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T23:32:05.743279Z","iopub.execute_input":"2024-11-18T23:32:05.744154Z","iopub.status.idle":"2024-11-18T23:32:05.76477Z","shell.execute_reply.started":"2024-11-18T23:32:05.744122Z","shell.execute_reply":"2024-11-18T23:32:05.763906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nfrom tqdm import tqdm\nimport numpy as np\nimport os\n\nimport cv2\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset\n\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import ConfusionMatrixDisplay, classification_report, f1_score\nimport math\n\nimport gc\n\nclass ConfusionMatrix():\n    def __init__(self, labels):\n        self.labels = np.array(labels) # idx 0: nocall\n        self.cls_num = len(self.labels)\n        # row == True : col == Pred\n        self.matrix = np.array([[0]*self.cls_num for _ in range(self.cls_num)])\n\n        self.true = []\n        self.pred = []\n\n\n    def __call__(self, output, target):\n        return self.calc_confmat(output, target)\n\n\n    def calc_confmat(self, outputs, targets):\n        # multilabel classification:\n        # outputs.shape = [Batch_N, N]\n        # targets.shape = [Batch_N, N]\n        if torch.is_tensor(outputs) or torch.is_tensor(targets):\n            outputs = outputs.cpu().detach().numpy()\n            targets = targets.cpu().detach().numpy()\n            \n        output_classes = []\n        for output in outputs:\n            for i, out in enumerate(output):\n                output[i] = 1 if out >= 0.5 else 0\n            output_classes.append(output.astype(np.int32))\n\n        target_classes = []\n        for target in targets:\n            for i, tar in enumerate(target):\n                target[i] = 1 if tar >= 0.5 else 0\n            target_classes.append(target.astype(np.int32))\n\n        # target_classes = targets\n        self.matrix[target_classes, output_classes] += 1\n        \n        # for further multilabel 'classification_report()'\n        self.true = [*self.true, *target_classes]\n        self.pred = [*self.pred, *output_classes]\n\n    def f1_score(self):\n        return f1_score(self.true, self.pred, average='weighted', zero_division=0.0,)\n        \n    def classification_report(self, output_dict=False):\n        cls_report = classification_report(\n            self.true,\n            self.pred,\n            target_names=self.labels,\n            output_dict=output_dict,\n            zero_division=0.0,\n        )\n        return cls_report\n\n\n    def plot(self):\n        figsize = (8, 8)\n        fig, ax = plt.subplots(figsize=figsize)\n        disp = ConfusionMatrixDisplay(self.matrix, display_labels=self.labels)\n        disp.plot(ax=ax)\n        plt.show()\n\n\nclass EarlyStopper:\n    def __init__(self, model: nn.Module, weights_name, scripted_name, save_dir, patience=10, delta=3):\n        self.model = model\n        self.weights_name = save_dir / weights_name\n        self.scripted_name = save_dir / scripted_name\n\n        self.patience = patience\n        self.delta = delta\n\n        self.counter = 0\n\n        self.best_loss = float('inf')\n        self.best_f1 = 0.0\n\n\n    def __call__(self, valid_loss, valid_f1=0.0):\n        match self.early_stop(valid_loss, valid_f1):\n            case 0:\n                print(f'\\nEarly stopped: best_loss = {self.best_loss:.3} | best_f1 = {self.best_f1:.3}')\n                print(f'Model weights saved with name  \"{self.weights_name}\"')\n                print(f'Scripted model saved with name \"{self.scripted_name}\"')\n                return True\n            case 2:\n                self.save_model()\n                return False\n            case _:\n                return False\n\n\n    def save_model(self):\n        torch.save(self.model, self.weights_name)\n        try:\n            model_scripted = torch.jit.script(self.model)\n            model_scripted.save(self.scripted_name)\n        except RuntimeError: # конвертированные модели не скриптуются\n            pass\n\n\n    def early_stop(self, valid_loss, valid_f1):\n        \"\"\"\n        0: stop\n        1: continue\n        2: save best\n        \"\"\"\n        valid_loss = float(f'{valid_loss:.{self.delta}}')\n        valid_f1 = float(f'{valid_f1:.{self.delta}}')\n        if valid_f1 > self.best_f1:\n            self.best_f1 = valid_f1\n            if valid_loss < self.best_loss:\n                self.best_loss = valid_loss\n            self.counter = 0\n            return 2\n        elif valid_loss < self.best_loss:\n            self.best_loss = valid_loss\n            self.counter = 0\n            return 1\n        else:\n            self.counter += 1\n            if self.counter >= self.patience:\n                return 0\n\n        return 1\n    \n\nclass Trainer():\n    '''\n    Classification task trainer. Should be reimplemented through a base class\n    '''\n    def __init__(self, model: nn.Module, train_loader, valid_loader, test_loader, name: str, save_dir: Path, weights=None, labels=[]):\n        self.device = 'cuda:0' if torch.cuda.is_available() else 'cpu'\n\n        self.model = model.to(self.device)\n\n        self.weights_name = f'{name}_weights.pt'\n        self.scripted_name = f'{name}_scripted.pt'\n\n        self.train_loader, self.valid_loader, self.test_loader = train_loader, valid_loader, test_loader\n\n        if weights is not None:\n          weights = torch.tensor(weights, dtype=torch.float32).to(self.device)\n        # self.class_criterion = nn.CrossEntropyLoss(weight=weights)\n        self.class_criterion = nn.BCEWithLogitsLoss(weight=weights)\n        \n        self.regression_criterion = nn.MSELoss()\n\n        self.train_losses = []\n        self.train_f1s = []\n\n        self.valid_losses = []\n        self.valid_f1s = []\n\n        self.early_stopper = EarlyStopper(\n                self.model,\n                self.weights_name,\n                self.scripted_name,\n                save_dir,\n                patience=5,\n                delta=3,\n        )\n\n        self.labels = labels\n        if len(self.labels) == 0:\n            # model.eval()\n            self.labels = list(range(len(train_loader.dataset[0][1])))\n\n        self.confmat_test: ConfusionMatrix\n\n    def get_best_model(self):\n        if os.path.exists(self.scripted_name):\n            print(f'loaded best from disk {self.scripted_name}')\n            return torch.jit.load(self.scripted_name)\n        return self.model\n\n\n    def train(self, epoch_num, lr=0.001, load_from_disk=False):\n        gc.collect()\n        torch.cuda.empty_cache()\n        \n        self.optimizer = torch.optim.Adam(\n                self.model.parameters(),\n                lr=lr,\n                # betas=(0.999, 0.9999),\n                # weight_decay=0.01,\n        )\n        self.scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(self.optimizer, eta_min=1e-5, T_max=epoch_num)\n        self.early_stopper.counter = 0\n\n        if load_from_disk:\n            self.model = self.get_best_model()\n        # else:\n        #     if os.path.exists(self.weights_name): os.remove(self.weights_name)\n        #     if os.path.exists(self.scripted_name): os.remove(self.scripted_name)\n\n        self.model.to(self.device)\n        for epoch in range(epoch_num):\n            print(f'\\nEpoch: {epoch+1} / {epoch_num}')\n\n            train_loss, train_f1 = self.train_step(self.train_loader)\n            self.train_losses.append(train_loss)\n            self.train_f1s.append(train_f1)\n\n            valid_loss, valid_f1 = self.valid_step(self.valid_loader)\n            self.valid_losses.append(valid_loss)\n            self.valid_f1s.append(valid_f1)\n\n            self.scheduler.step()\n\n            if self.early_stopper(valid_loss, valid_f1):\n                break\n                \n        gc.collect()\n        torch.cuda.empty_cache()\n\n\n    def report(self, max_classes_not_to_reduce=20):\n        print('\\nTest-Epoch')\n        self.model = self.get_best_model()\n        self.model.to(self.device)\n        confmat = self.valid_step(self.valid_loader, return_confmat=True)\n\n        report = confmat.classification_report(output_dict=False)\n        if confmat.cls_num > max_classes_not_to_reduce:\n            print(f'\\n\\nlen(classes) is bigger than {max_classes_not_to_reduce}, so output will be reduced')\n            report_list = report.split('\\n')\n            report = '\\n'.join([*report_list[:2], *report_list[-3:]])\n        print('\\n', report)\n        return confmat\n\n\n    def train_step(self, train_loader):\n        self.model.train()\n        run_loss = 0.0\n        confmat = ConfusionMatrix(labels=self.labels)\n\n        for img, target in (pbar := tqdm(train_loader)):\n            self.optimizer.zero_grad()\n            target = target.to(self.device)\n            batch_size = len(target)\n\n            img = img.to(self.device)\n            output = self.model(img)\n\n            loss = self.class_criterion(output, target.to(torch.float32))\n            \n            if math.isnan(loss):\n                print(img)\n                print(output)\n                print(target)\n                np.save(work_dir / 'img.npy', img.cpu().detach().numpy())\n\n            loss.backward()\n            self.optimizer.step()\n\n            run_loss += loss.item()\n            confmat(output, target)\n            f1 = confmat.f1_score()\n            pbar.set_postfix({'loss ': f' {loss:.3}', 'f1 ': f' {f1:.3}'})\n\n        train_loss = run_loss / len(train_loader)\n        train_f1 = confmat.classification_report(output_dict=True)['weighted avg']['f1-score']\n        print(f'  train loss:\\t {train_loss:.3}\\t | train f1:\\t {train_f1:.3}')\n\n        return train_loss, train_f1\n\n\n    def valid_step(self, valid_loader, return_confmat=False) -> (tuple | ConfusionMatrix):\n        self.model.eval()\n        run_loss = 0.0\n        confmat = ConfusionMatrix(labels=self.labels)\n\n        for img, target in (pbar := tqdm(valid_loader)):\n            target = target.to(self.device)\n            batch_size = len(target)\n            \n            img = img.to(self.device)\n            output = self.model(img)\n\n            loss = self.class_criterion(output, target.to(torch.float32))\n            \n            run_loss += loss.item()\n            confmat(output, target)\n            f1 = confmat.f1_score()\n            pbar.set_postfix({'loss ': f' {loss:.3}', 'f1 ': f' {f1:.3}'})\n\n        valid_loss = run_loss / len(valid_loader)\n        valid_f1 = confmat.classification_report(output_dict=True)['weighted avg']['f1-score']\n        print(f'  valid loss:\\t {valid_loss:.3}\\t | valid f1:\\t {valid_f1:.3}')\n\n        if return_confmat:\n            return confmat\n        return valid_loss, valid_f1\n    \n    \n    def predict(self, valid_loader, return_confmat=False) -> np.ndarray:\n        self.model.eval()\n        preds = []\n\n        for inputs, target in tqdm(valid_loader):\n            inputs, target = inputs.to(self.device), target.to(self.device)\n            outputs = self.model(inputs).detach().numpy()\n            pred = np.argmax(outputs, axis=1 if len(outputs.shape) > 1 else 0)\n            preds = [*preds, *pred]\n\n        return np.array(preds)\n    \n\n    def plot(self):\n        epochs = list(range(1, len(self.train_losses)+1))\n\n        fig, axes = plt.subplots(nrows=1, ncols=2, figsize=(8, 4))\n\n        axes[0].plot(epochs, self.train_losses, '--b',label='train')\n        axes[0].plot(epochs, self.valid_losses, 'r',label='valid')\n        axes[0].set(xlabel='epoch num', ylabel='loss', title='Loss')\n        axes[0].grid()\n        axes[0].legend()\n\n        axes[1].plot(epochs, self.train_f1s, '--b', label='train')\n        axes[1].plot(epochs, self.valid_f1s, 'r', label='valid')\n        axes[1].set(xlabel='epoch num', ylabel='f1', title='F1-score')\n        axes[1].grid()\n        axes[1].legend()\n\n        fig.tight_layout()\n        plt.show() ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T15:48:32.857585Z","iopub.execute_input":"2024-11-19T15:48:32.857985Z","iopub.status.idle":"2024-11-19T15:48:32.895108Z","shell.execute_reply.started":"2024-11-19T15:48:32.857951Z","shell.execute_reply":"2024-11-19T15:48:32.894211Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Process","metadata":{}},{"cell_type":"code","source":"model = EfficientNet(len(classes))\n# model = Resnet(len(classes))\n\ntrainer = Trainer(\n    model=model,\n    train_loader=train_loader,\n    valid_loader=valid_loader,\n    test_loader=None,\n    name='efficientnet', # efficientnet resnet\n    save_dir=Path('/kaggle/working'),\n    # weights=weights,\n    labels=classes,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T15:48:35.359654Z","iopub.execute_input":"2024-11-19T15:48:35.360075Z","iopub.status.idle":"2024-11-19T15:48:35.621832Z","shell.execute_reply.started":"2024-11-19T15:48:35.360036Z","shell.execute_reply":"2024-11-19T15:48:35.620864Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer.train(1, 0.0008, load_from_disk=False)\n# trainer.plot()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T15:53:38.690264Z","iopub.execute_input":"2024-11-19T15:53:38.691312Z","iopub.status.idle":"2024-11-19T15:53:38.69564Z","shell.execute_reply.started":"2024-11-19T15:53:38.691274Z","shell.execute_reply":"2024-11-19T15:53:38.694503Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img = np.load(work_dir / 'img.npy')\nfig, ax = plt.subplots(figsize=(10, 3))\nax.imshow(melspec.cpu().detach().numpy()[0], origin=\"lower\", aspect=\"auto\", interpolation=\"nearest\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T15:56:34.468232Z","iopub.execute_input":"2024-11-19T15:56:34.468601Z","iopub.status.idle":"2024-11-19T15:56:34.725842Z","shell.execute_reply.started":"2024-11-19T15:56:34.468568Z","shell.execute_reply":"2024-11-19T15:56:34.724967Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"out = model(torch.tensor(img).to('cuda'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T15:58:18.920976Z","iopub.execute_input":"2024-11-19T15:58:18.921319Z","iopub.status.idle":"2024-11-19T15:58:18.966265Z","shell.execute_reply.started":"2024-11-19T15:58:18.921288Z","shell.execute_reply":"2024-11-19T15:58:18.965374Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T15:58:41.75182Z","iopub.execute_input":"2024-11-19T15:58:41.752154Z","iopub.status.idle":"2024-11-19T15:58:41.759586Z","shell.execute_reply.started":"2024-11-19T15:58:41.752125Z","shell.execute_reply":"2024-11-19T15:58:41.758814Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer.report(400)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T13:15:39.22778Z","iopub.execute_input":"2024-11-19T13:15:39.228628Z","iopub.status.idle":"2024-11-19T13:15:51.284276Z","shell.execute_reply.started":"2024-11-19T13:15:39.228587Z","shell.execute_reply":"2024-11-19T13:15:51.28319Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer.train(5, 0.001, load_from_disk=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T01:49:26.709051Z","iopub.execute_input":"2024-11-19T01:49:26.709348Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = EfficientNet(len(classes))\n\ntrainer = Trainer(\n    model=model,\n    train_loader=train_10mins_loader,\n    valid_loader=valid_10mins_loader,\n    test_loader=None,\n    name='efficientnet',\n    save_dir=Path('/kaggle/working'),\n    # weights=weights,\n    labels=classes,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T22:16:08.717806Z","iopub.execute_input":"2024-11-18T22:16:08.718464Z","iopub.status.idle":"2024-11-18T22:16:09.851113Z","shell.execute_reply.started":"2024-11-18T22:16:08.718433Z","shell.execute_reply":"2024-11-18T22:16:09.850414Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer.train(10, 0.5, load_from_disk=False)\n# trainer.plot()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T22:17:39.738881Z","iopub.execute_input":"2024-11-18T22:17:39.739644Z","iopub.status.idle":"2024-11-18T22:20:40.456884Z","shell.execute_reply.started":"2024-11-18T22:17:39.739613Z","shell.execute_reply":"2024-11-18T22:20:40.45586Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer.report(400)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T22:22:11.641937Z","iopub.execute_input":"2024-11-18T22:22:11.642326Z","iopub.status.idle":"2024-11-18T22:22:13.334458Z","shell.execute_reply.started":"2024-11-18T22:22:11.642292Z","shell.execute_reply":"2024-11-18T22:22:13.333597Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nfrom xgboost import XGBClassifier\nimport warnings\n\nX_train, X_val, y_train, y_val = train_test_split(\n    train_val_short_set.df.drop(columns=['cls', 'row_id', '0', '1']),\n    train_val_short_set.df.cls-1,\n    test_size=0.2,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T00:23:24.300322Z","iopub.execute_input":"2024-11-18T00:23:24.300886Z","iopub.status.idle":"2024-11-18T00:23:24.318647Z","shell.execute_reply.started":"2024-11-18T00:23:24.300809Z","shell.execute_reply":"2024-11-18T00:23:24.31715Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import accuracy_score\n\nwith warnings.catch_warnings():\n    warnings.simplefilter('ignore')\n\n    xgb = XGBClassifier()\n    xgb.fit(X_train, y_train)\n    preds = xgb.predict(X_val)\naccuracy_score(y_val, preds)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T00:23:27.717718Z","iopub.execute_input":"2024-11-18T00:23:27.718553Z","iopub.status.idle":"2024-11-18T00:26:52.765907Z","shell.execute_reply.started":"2024-11-18T00:23:27.718472Z","shell.execute_reply":"2024-11-18T00:26:52.763911Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"accuracy_score(y_val.values.astype(np.int32), preds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T00:32:58.269762Z","iopub.execute_input":"2024-11-18T00:32:58.271136Z","iopub.status.idle":"2024-11-18T00:32:58.282618Z","shell.execute_reply.started":"2024-11-18T00:32:58.271079Z","shell.execute_reply":"2024-11-18T00:32:58.281054Z"}},"outputs":[],"execution_count":null}]}