{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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":"gpu","dataSources":[{"sourceId":73154,"databundleVersionId":8047843,"sourceType":"competition"},{"sourceId":169133106,"sourceType":"kernelVersion"}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# path to your train/test/meta folders\nDATA_PATH = '../'\n\n# names of valuable files/folders\ntrain_meta_fname = '/kaggle/input/itmo-acoustic-event-detection-2024/train.csv'\ntest_meta_fname = '/kaggle/input/itmo-acoustic-event-detection-2024/sample_submission.csv'\ntrain_data_folder = '/kaggle/input/itmo-acoustic-event-detection-2024/audio_train/train'\ntest_data_folder = '/kaggle/input/itmo-acoustic-event-detection-2024/audio_test/test'","metadata":{"execution":{"iopub.status.busy":"2024-05-09T13:33:57.155304Z","iopub.execute_input":"2024-05-09T13:33:57.15566Z","iopub.status.idle":"2024-05-09T13:33:57.166914Z","shell.execute_reply.started":"2024-05-09T13:33:57.155632Z","shell.execute_reply":"2024-05-09T13:33:57.165966Z"},"trusted":true},"outputs":[],"execution_count":1},{"cell_type":"code","source":"!pip install efficientnet_pytorch","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport os\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nimport torchaudio\nimport torchvision\nfrom torchaudio import transforms\n# from efficientnet_pytorch import EfficientNet\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import f1_score\nfrom tqdm import tqdm\n# Load model directly\nfrom transformers import AutoImageProcessor, AutoModelForImageClassification","metadata":{"execution":{"iopub.status.busy":"2024-04-25T12:59:12.383352Z","iopub.execute_input":"2024-04-25T12:59:12.384195Z","iopub.status.idle":"2024-04-25T12:59:29.829152Z","shell.execute_reply.started":"2024-04-25T12:59:12.384166Z","shell.execute_reply":"2024-04-25T12:59:29.828141Z"},"trusted":true},"outputs":[{"name":"stderr","text":"2024-04-25 12:59:21.778756: E external/local_xla/xla/stream_executor/cuda/cuda_dnn.cc:9261] Unable to register cuDNN factory: Attempting to register factory for plugin cuDNN when one has already been registered\n2024-04-25 12:59:21.778908: E external/local_xla/xla/stream_executor/cuda/cuda_fft.cc:607] Unable to register cuFFT factory: Attempting to register factory for plugin cuFFT when one has already been registered\n2024-04-25 12:59:21.892750: E external/local_xla/xla/stream_executor/cuda/cuda_blas.cc:1515] Unable to register cuBLAS factory: Attempting to register factory for plugin cuBLAS when one has already been registered\n","output_type":"stream"}],"execution_count":3},{"cell_type":"code","source":"# set seeds\nimport random\nimport numpy as np\n\nrandom.seed(42)\nnp.random.seed(42)\ntorch.manual_seed(42)\ntorch.cuda.manual_seed(42)\ntorch.backends.cudnn.deterministic = True","metadata":{"execution":{"iopub.status.busy":"2024-04-25T12:59:29.840063Z","iopub.execute_input":"2024-04-25T12:59:29.840424Z","iopub.status.idle":"2024-04-25T12:59:29.860749Z","shell.execute_reply.started":"2024-04-25T12:59:29.840378Z","shell.execute_reply":"2024-04-25T12:59:29.859996Z"},"trusted":true},"outputs":[],"execution_count":5},{"cell_type":"code","source":"df_train = pd.read_csv(os.path.join(DATA_PATH, train_meta_fname))\ndf_test = pd.read_csv(os.path.join(DATA_PATH, test_meta_fname))\ndf_train.head(2)","metadata":{"execution":{"iopub.status.busy":"2024-04-25T12:59:29.862711Z","iopub.execute_input":"2024-04-25T12:59:29.862968Z","iopub.status.idle":"2024-04-25T12:59:29.914715Z","shell.execute_reply.started":"2024-04-25T12:59:29.862946Z","shell.execute_reply":"2024-04-25T12:59:29.913802Z"},"trusted":true},"outputs":[{"execution_count":6,"output_type":"execute_result","data":{"text/plain":"                      fname            label\n0  8bcbcc394ba64fe85ed4.wav  Finger_snapping\n1  00d77b917e241afa06f1.wav           Squeak","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>fname</th>\n      <th>label</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>0</th>\n      <td>8bcbcc394ba64fe85ed4.wav</td>\n      <td>Finger_snapping</td>\n    </tr>\n    <tr>\n      <th>1</th>\n      <td>00d77b917e241afa06f1.wav</td>\n      <td>Squeak</td>\n    </tr>\n  </tbody>\n</table>\n</div>"},"metadata":{}}],"execution_count":6},{"cell_type":"code","source":"n_classes = df_train.label.nunique()\nprint(n_classes)\nclasses_dict = {cl:i for i,cl in enumerate(df_train.label.unique())}\ndf_train['label_encoded'] = df_train.label.map(classes_dict)\ndf_train.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-25T12:59:29.915701Z","iopub.execute_input":"2024-04-25T12:59:29.915976Z","iopub.status.idle":"2024-04-25T12:59:29.935106Z","shell.execute_reply.started":"2024-04-25T12:59:29.915954Z","shell.execute_reply":"2024-04-25T12:59:29.934211Z"},"trusted":true},"outputs":[{"name":"stdout","text":"41\n","output_type":"stream"},{"execution_count":7,"output_type":"execute_result","data":{"text/plain":"                      fname            label  label_encoded\n0  8bcbcc394ba64fe85ed4.wav  Finger_snapping              0\n1  00d77b917e241afa06f1.wav           Squeak              1\n2  17bb93b73b8e79234cb3.wav   Electric_piano              2\n3  7d5c7a40a936136da55e.wav        Harmonica              3\n4  17e0ee7565a33d6c2326.wav       Snare_drum              4","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>fname</th>\n      <th>label</th>\n      <th>label_encoded</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>0</th>\n      <td>8bcbcc394ba64fe85ed4.wav</td>\n      <td>Finger_snapping</td>\n      <td>0</td>\n    </tr>\n    <tr>\n      <th>1</th>\n      <td>00d77b917e241afa06f1.wav</td>\n      <td>Squeak</td>\n      <td>1</td>\n    </tr>\n    <tr>\n      <th>2</th>\n      <td>17bb93b73b8e79234cb3.wav</td>\n      <td>Electric_piano</td>\n      <td>2</td>\n    </tr>\n    <tr>\n      <th>3</th>\n      <td>7d5c7a40a936136da55e.wav</td>\n      <td>Harmonica</td>\n      <td>3</td>\n    </tr>\n    <tr>\n      <th>4</th>\n      <td>17e0ee7565a33d6c2326.wav</td>\n      <td>Snare_drum</td>\n      <td>4</td>\n    </tr>\n  </tbody>\n</table>\n</div>"},"metadata":{}}],"execution_count":7},{"cell_type":"code","source":"label2id = classes_dict\nid2label = inv_map = {v: k for k, v in label2id.items()}","metadata":{"execution":{"iopub.status.busy":"2024-04-25T12:59:29.936917Z","iopub.execute_input":"2024-04-25T12:59:29.93722Z","iopub.status.idle":"2024-04-25T12:59:29.949559Z","shell.execute_reply.started":"2024-04-25T12:59:29.937189Z","shell.execute_reply":"2024-04-25T12:59:29.948577Z"},"trusted":true},"outputs":[],"execution_count":8},{"cell_type":"code","source":"# https://github.com/lukemelas/EfficientNet-PyTorch\nfrom transformers import ResNetConfig, ResNetModel, ResNetForImageClassification\n\nclass BaseLineModel(nn.Module):\n    \n    def __init__(self, sample_rate=16000, n_classes=41):\n        super().__init__()\n        self.ms = torchaudio.transforms.MelSpectrogram(sample_rate, n_fft=2048, hop_length = 80, win_length = 800, n_mels=128, normalized=True)\n#         self.ms = torchaudio.transforms.Spectrogram(n_fft=448)\n        \n#         encoder_layer = nn.TransformerEncoderLayer(d_model=64, nhead=4, batch_first=True)\n#         self.t_encoder = nn.TransformerEncoder(encoder_layer, num_layers=2)\n        \n#         self.bn1 = nn.BatchNorm2d(1)\n        \n#         self.cnn1 = nn.Conv2d(in_channels=1, out_channels=10, kernel_size=3, padding=1)\n#         self.cnn3 = nn.Conv2d(in_channels=10, out_channels=3, kernel_size=3, padding=1)\n        \n#         self.features = EfficientNet.from_pretrained('efficientnet-b4')\n#         self.features = AutoModelForImageClassification.from_pretrained(\"microsoft/swin-tiny-patch4-window7-224\")\n\n        self.features = AutoModelForImageClassification.from_pretrained(\"microsoft/resnet-18\", num_channels=1, num_labels=41, ignore_mismatched_sizes=True)\n#         config = ResNetConfig(num_channels=1, num_labels=41)\n        \n#         self.features = ResNetForImageClassification(config)\n    # use it as features\n#         for param in self.features.parameters():\n#             param.requires_grad = False\n            \n#         self.lin1 = nn.Linear(1000, 333)\n        \n#         self.lin2 = nn.Linear(333, 111)\n                \n#         self.lin3 = nn.Linear(111, n_classes)\n        \n    def forward(self, x):\n        x = self.ms(x)\n        x = self.features(x).logits\n#         x = self.bn1(x)\n\n#         x = x.squeeze().swapaxes(1,2)\n#         x = self.t_encoder(x).unsqueeze(1)\n#         x = x.swapaxes(1,2).unsqueeze(1)\n\n#         s = self.t_encoder(x)\n                \n#         x = F.relu(self.cnn1(x))\n#         x = F.relu(self.cnn3(x))\n        \n#         x = self.features(x).logits\n#         x = x.view(x.shape[0], -1)\n#         x = F.relu(x)\n\n#         x = F.relu(self.lin1(x))\n#         x = F.relu(self.lin2(x))\n#         x = self.lin3(x)\n        return x\n    \n    def inference(self, x):\n        x = self.forward(x)\n        x = F.softmax(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-04-24T20:12:42.652812Z","iopub.execute_input":"2024-04-24T20:12:42.653206Z","iopub.status.idle":"2024-04-24T20:12:42.985902Z","shell.execute_reply.started":"2024-04-24T20:12:42.653177Z","shell.execute_reply":"2024-04-24T20:12:42.984851Z"},"trusted":true},"outputs":[],"execution_count":11},{"cell_type":"code","source":"def sample_or_pad(waveform, wav_len=64000):\n    m, n = waveform.shape\n    if n < wav_len:\n        padded_wav = torch.zeros(1, wav_len)\n        padded_wav[:, :n] = waveform\n        return padded_wav\n    elif n > wav_len:\n        offset = np.random.randint(0, n - wav_len)\n        sampled_wav = waveform[:, offset:offset+wav_len]\n        return sampled_wav\n    else:\n        return waveform\n        \nclass EventDetectionDataset(Dataset):\n    def __init__(self, data_path, x, y=None):\n        self.x = x\n        self.y = y\n        self.data_path = data_path\n        self.ms = torchaudio.transforms.MelSpectrogram(16000, n_fft=2048, hop_length = 80, win_length = 800, n_mels=128, normalized=True)\n    \n    def __len__(self):\n        return len(self.x)\n\n    def __getitem__(self, idx):\n        path2wav = os.path.join(self.data_path, self.x[idx])\n        waveform, sample_rate = torchaudio.load(path2wav, normalize=True)\n        waveform = sample_or_pad(waveform)\n        if self.y is not None:\n            image = self.ms(waveform)\n            return {'pixel_values':image, 'label':self.y[idx]}\n#             return waveform, self.y[idx]\n        return waveform","metadata":{"execution":{"iopub.status.busy":"2024-04-24T20:12:49.499448Z","iopub.execute_input":"2024-04-24T20:12:49.500154Z","iopub.status.idle":"2024-04-24T20:12:49.510361Z","shell.execute_reply.started":"2024-04-24T20:12:49.500117Z","shell.execute_reply":"2024-04-24T20:12:49.509444Z"},"trusted":true},"outputs":[],"execution_count":13},{"cell_type":"code","source":"batch_size=64\nX_train, X_val, y_train, y_val = train_test_split(df_train.fname.values, df_train.label_encoded.values, stratify=df_train.label_encoded.values, \n                                                  test_size=0.2, random_state=42)\ntrain_loader = DataLoader(\n                        EventDetectionDataset(os.path.join(DATA_PATH, train_data_folder), X_train, y_train),\n                        batch_size=batch_size\n                )\nval_loader = DataLoader(\n                        EventDetectionDataset(os.path.join(DATA_PATH, train_data_folder), X_val, y_val),\n                        batch_size=batch_size\n                )\ntest_loader = DataLoader(\n                        EventDetectionDataset(os.path.join(DATA_PATH, test_data_folder), df_test.fname.values, None),\n                        batch_size=batch_size, shuffle=False\n                )","metadata":{"execution":{"iopub.status.busy":"2024-04-24T16:58:25.687441Z","iopub.execute_input":"2024-04-24T16:58:25.688322Z","iopub.status.idle":"2024-04-24T16:58:25.707542Z","shell.execute_reply.started":"2024-04-24T16:58:25.688288Z","shell.execute_reply":"2024-04-24T16:58:25.706547Z"},"trusted":true},"outputs":[],"execution_count":29},{"cell_type":"code","source":"def eval_model(model, eval_dataset):\n    model.eval()\n    forecast, true_labs = [], []\n    with torch.no_grad():\n        for wavs, labs in tqdm(eval_dataset):\n            wavs, labs = wavs.cuda(), labs.detach().numpy()\n            true_labs.append(labs)\n            outputs = model.inference(wavs)\n            \n            outputs = outputs.detach().cpu().numpy().argmax(axis=1)\n            forecast.append(outputs)\n    forecast = [x for sublist in forecast for x in sublist]\n    true_labs = [x for sublist in true_labs for x in sublist]\n    return f1_score(forecast, true_labs, average='macro')","metadata":{"execution":{"iopub.status.busy":"2024-04-24T16:47:53.061152Z","iopub.execute_input":"2024-04-24T16:47:53.061889Z","iopub.status.idle":"2024-04-24T16:47:53.068871Z","shell.execute_reply.started":"2024-04-24T16:47:53.061853Z","shell.execute_reply":"2024-04-24T16:47:53.067913Z"},"trusted":true},"outputs":[],"execution_count":13},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\nmodel = BaseLineModel()\nmodel = model.cuda()\nlr = 1e-2\n\n# model.features.init_weights()\n\noptimizer = torch.optim.Adam(model.parameters(), lr=lr)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model","metadata":{"execution":{"iopub.status.busy":"2024-04-24T17:11:39.174808Z","iopub.execute_input":"2024-04-24T17:11:39.175201Z","iopub.status.idle":"2024-04-24T17:11:39.189424Z","shell.execute_reply.started":"2024-04-24T17:11:39.175172Z","shell.execute_reply":"2024-04-24T17:11:39.188552Z"},"trusted":true},"outputs":[{"execution_count":42,"output_type":"execute_result","data":{"text/plain":"ResNetForImageClassification(\n  (resnet): ResNetModel(\n    (embedder): ResNetEmbeddings(\n      (embedder): ResNetConvLayer(\n        (convolution): Conv2d(1, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n        (normalization): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n        (activation): ReLU()\n      )\n      (pooler): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)\n    )\n    (encoder): ResNetEncoder(\n      (stages): ModuleList(\n        (0): ResNetStage(\n          (layers): Sequential(\n            (0): ResNetBasicLayer(\n              (shortcut): Identity()\n              (layer): Sequential(\n                (0): ResNetConvLayer(\n                  (convolution): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): ReLU()\n                )\n                (1): ResNetConvLayer(\n                  (convolution): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): Identity()\n                )\n              )\n              (activation): ReLU()\n            )\n            (1): ResNetBasicLayer(\n              (shortcut): Identity()\n              (layer): Sequential(\n                (0): ResNetConvLayer(\n                  (convolution): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): ReLU()\n                )\n                (1): ResNetConvLayer(\n                  (convolution): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): Identity()\n                )\n              )\n              (activation): ReLU()\n            )\n          )\n        )\n        (1): ResNetStage(\n          (layers): Sequential(\n            (0): ResNetBasicLayer(\n              (shortcut): ResNetShortCut(\n                (convolution): Conv2d(64, 128, kernel_size=(1, 1), stride=(2, 2), bias=False)\n                (normalization): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n              )\n              (layer): Sequential(\n                (0): ResNetConvLayer(\n                  (convolution): Conv2d(64, 128, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): ReLU()\n                )\n                (1): ResNetConvLayer(\n                  (convolution): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): Identity()\n                )\n              )\n              (activation): ReLU()\n            )\n            (1): ResNetBasicLayer(\n              (shortcut): Identity()\n              (layer): Sequential(\n                (0): ResNetConvLayer(\n                  (convolution): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): ReLU()\n                )\n                (1): ResNetConvLayer(\n                  (convolution): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): Identity()\n                )\n              )\n              (activation): ReLU()\n            )\n          )\n        )\n        (2): ResNetStage(\n          (layers): Sequential(\n            (0): ResNetBasicLayer(\n              (shortcut): ResNetShortCut(\n                (convolution): Conv2d(128, 256, kernel_size=(1, 1), stride=(2, 2), bias=False)\n                (normalization): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n              )\n              (layer): Sequential(\n                (0): ResNetConvLayer(\n                  (convolution): Conv2d(128, 256, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): ReLU()\n                )\n                (1): ResNetConvLayer(\n                  (convolution): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): Identity()\n                )\n              )\n              (activation): ReLU()\n            )\n            (1): ResNetBasicLayer(\n              (shortcut): Identity()\n              (layer): Sequential(\n                (0): ResNetConvLayer(\n                  (convolution): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): ReLU()\n                )\n                (1): ResNetConvLayer(\n                  (convolution): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): Identity()\n                )\n              )\n              (activation): ReLU()\n            )\n          )\n        )\n        (3): ResNetStage(\n          (layers): Sequential(\n            (0): ResNetBasicLayer(\n              (shortcut): ResNetShortCut(\n                (convolution): Conv2d(256, 512, kernel_size=(1, 1), stride=(2, 2), bias=False)\n                (normalization): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n              )\n              (layer): Sequential(\n                (0): ResNetConvLayer(\n                  (convolution): Conv2d(256, 512, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): ReLU()\n                )\n                (1): ResNetConvLayer(\n                  (convolution): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): Identity()\n                )\n              )\n              (activation): ReLU()\n            )\n            (1): ResNetBasicLayer(\n              (shortcut): Identity()\n              (layer): Sequential(\n                (0): ResNetConvLayer(\n                  (convolution): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): ReLU()\n                )\n                (1): ResNetConvLayer(\n                  (convolution): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n                  (normalization): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n                  (activation): Identity()\n                )\n              )\n              (activation): ReLU()\n            )\n          )\n        )\n      )\n    )\n    (pooler): AdaptiveAvgPool2d(output_size=(1, 1))\n  )\n  (classifier): Sequential(\n    (0): Flatten(start_dim=1, end_dim=-1)\n    (1): Linear(in_features=512, out_features=41, bias=True)\n  )\n)"},"metadata":{}}],"execution_count":42},{"cell_type":"code","source":"train_batch = next(iter(train_loader))","metadata":{"execution":{"iopub.status.busy":"2024-04-24T17:26:38.071282Z","iopub.execute_input":"2024-04-24T17:26:38.071665Z","iopub.status.idle":"2024-04-24T17:26:38.976243Z","shell.execute_reply.started":"2024-04-24T17:26:38.071637Z","shell.execute_reply":"2024-04-24T17:26:38.975229Z"},"trusted":true},"outputs":[],"execution_count":46},{"cell_type":"code","source":"train_batch['label'].dtype","metadata":{"execution":{"iopub.status.busy":"2024-04-24T17:26:39.758143Z","iopub.execute_input":"2024-04-24T17:26:39.758803Z","iopub.status.idle":"2024-04-24T17:26:39.76606Z","shell.execute_reply.started":"2024-04-24T17:26:39.758773Z","shell.execute_reply":"2024-04-24T17:26:39.764901Z"},"trusted":true},"outputs":[{"execution_count":47,"output_type":"execute_result","data":{"text/plain":"torch.int64"},"metadata":{}}],"execution_count":47},{"cell_type":"code","source":"model.forward(train_batch['pixel_values'].to('cuda')).logits.dtype","metadata":{"execution":{"iopub.status.busy":"2024-04-24T17:27:38.272976Z","iopub.execute_input":"2024-04-24T17:27:38.273673Z","iopub.status.idle":"2024-04-24T17:27:38.293692Z","shell.execute_reply.started":"2024-04-24T17:27:38.273638Z","shell.execute_reply":"2024-04-24T17:27:38.292636Z"},"trusted":true},"outputs":[{"execution_count":55,"output_type":"execute_result","data":{"text/plain":"torch.float32"},"metadata":{}}],"execution_count":55},{"cell_type":"code","source":"mel = model.ms(train_batch[0].to('cuda'))\n# efnet = model.features(cnn3)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mel.shape","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install evaluate","metadata":{"execution":{"iopub.status.busy":"2024-04-25T13:00:06.936228Z","iopub.execute_input":"2024-04-25T13:00:06.937184Z","iopub.status.idle":"2024-04-25T13:00:20.977713Z","shell.execute_reply.started":"2024-04-25T13:00:06.937143Z","shell.execute_reply":"2024-04-25T13:00:20.976542Z"},"trusted":true},"outputs":[{"name":"stdout","text":"Collecting evaluate\n  Downloading evaluate-0.4.1-py3-none-any.whl.metadata (9.4 kB)\nRequirement already satisfied: datasets>=2.0.0 in /opt/conda/lib/python3.10/site-packages (from evaluate) (2.18.0)\nRequirement already satisfied: numpy>=1.17 in /opt/conda/lib/python3.10/site-packages (from evaluate) (1.26.4)\nRequirement already satisfied: dill in /opt/conda/lib/python3.10/site-packages (from evaluate) (0.3.8)\nRequirement already satisfied: pandas in /opt/conda/lib/python3.10/site-packages (from evaluate) (2.1.4)\nRequirement already satisfied: requests>=2.19.0 in /opt/conda/lib/python3.10/site-packages (from evaluate) (2.31.0)\nRequirement already satisfied: tqdm>=4.62.1 in /opt/conda/lib/python3.10/site-packages (from evaluate) (4.66.1)\nRequirement already satisfied: xxhash in /opt/conda/lib/python3.10/site-packages (from evaluate) (3.4.1)\nRequirement already satisfied: multiprocess in /opt/conda/lib/python3.10/site-packages (from evaluate) (0.70.16)\nRequirement already satisfied: fsspec>=2021.05.0 in /opt/conda/lib/python3.10/site-packages (from fsspec[http]>=2021.05.0->evaluate) (2024.2.0)\nRequirement already satisfied: huggingface-hub>=0.7.0 in /opt/conda/lib/python3.10/site-packages (from evaluate) (0.22.2)\nRequirement already satisfied: packaging in /opt/conda/lib/python3.10/site-packages (from evaluate) (21.3)\nCollecting responses<0.19 (from evaluate)\n  Downloading responses-0.18.0-py3-none-any.whl.metadata (29 kB)\nRequirement already satisfied: filelock in /opt/conda/lib/python3.10/site-packages (from datasets>=2.0.0->evaluate) (3.13.1)\nRequirement already satisfied: pyarrow>=12.0.0 in /opt/conda/lib/python3.10/site-packages (from datasets>=2.0.0->evaluate) (15.0.2)\nRequirement already satisfied: pyarrow-hotfix in /opt/conda/lib/python3.10/site-packages (from datasets>=2.0.0->evaluate) (0.6)\nRequirement already satisfied: aiohttp in /opt/conda/lib/python3.10/site-packages (from datasets>=2.0.0->evaluate) (3.9.1)\nRequirement already satisfied: pyyaml>=5.1 in /opt/conda/lib/python3.10/site-packages (from datasets>=2.0.0->evaluate) (6.0.1)\nRequirement already satisfied: typing-extensions>=3.7.4.3 in /opt/conda/lib/python3.10/site-packages (from huggingface-hub>=0.7.0->evaluate) (4.9.0)\nRequirement already satisfied: pyparsing!=3.0.5,>=2.0.2 in /opt/conda/lib/python3.10/site-packages (from packaging->evaluate) (3.1.1)\nRequirement already satisfied: charset-normalizer<4,>=2 in /opt/conda/lib/python3.10/site-packages (from requests>=2.19.0->evaluate) (3.3.2)\nRequirement already satisfied: idna<4,>=2.5 in /opt/conda/lib/python3.10/site-packages (from requests>=2.19.0->evaluate) (3.6)\nRequirement already satisfied: urllib3<3,>=1.21.1 in /opt/conda/lib/python3.10/site-packages (from requests>=2.19.0->evaluate) (1.26.18)\nRequirement already satisfied: certifi>=2017.4.17 in /opt/conda/lib/python3.10/site-packages (from requests>=2.19.0->evaluate) (2024.2.2)\nRequirement already satisfied: python-dateutil>=2.8.2 in /opt/conda/lib/python3.10/site-packages (from pandas->evaluate) (2.9.0.post0)\nRequirement already satisfied: pytz>=2020.1 in /opt/conda/lib/python3.10/site-packages (from pandas->evaluate) (2023.3.post1)\nRequirement already satisfied: tzdata>=2022.1 in /opt/conda/lib/python3.10/site-packages (from pandas->evaluate) (2023.4)\nRequirement already satisfied: attrs>=17.3.0 in /opt/conda/lib/python3.10/site-packages (from aiohttp->datasets>=2.0.0->evaluate) (23.2.0)\nRequirement already satisfied: multidict<7.0,>=4.5 in /opt/conda/lib/python3.10/site-packages (from aiohttp->datasets>=2.0.0->evaluate) (6.0.4)\nRequirement already satisfied: yarl<2.0,>=1.0 in /opt/conda/lib/python3.10/site-packages (from aiohttp->datasets>=2.0.0->evaluate) (1.9.3)\nRequirement already satisfied: frozenlist>=1.1.1 in /opt/conda/lib/python3.10/site-packages (from aiohttp->datasets>=2.0.0->evaluate) (1.4.1)\nRequirement already satisfied: aiosignal>=1.1.2 in /opt/conda/lib/python3.10/site-packages (from aiohttp->datasets>=2.0.0->evaluate) (1.3.1)\nRequirement already satisfied: async-timeout<5.0,>=4.0 in /opt/conda/lib/python3.10/site-packages (from aiohttp->datasets>=2.0.0->evaluate) (4.0.3)\nRequirement already satisfied: six>=1.5 in /opt/conda/lib/python3.10/site-packages (from python-dateutil>=2.8.2->pandas->evaluate) (1.16.0)\nDownloading evaluate-0.4.1-py3-none-any.whl (84 kB)\n\u001b[2K   \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m84.1/84.1 kB\u001b[0m \u001b[31m4.3 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hDownloading responses-0.18.0-py3-none-any.whl (38 kB)\nInstalling collected packages: responses, evaluate\nSuccessfully installed evaluate-0.4.1 responses-0.18.0\n","output_type":"stream"}],"execution_count":10},{"cell_type":"code","source":"import evaluate\n\naccuracy = evaluate.load(\"accuracy\")\nf1_metric = evaluate.load(\"f1\")\n\ndef compute_metrics(eval_pred):\n    predictions, labels = eval_pred\n    predictions = np.argmax(predictions, axis=1)\n    acc = accuracy.compute(predictions=predictions, references=labels)['accuracy']\n    f1 = f1_metric.compute(predictions=predictions, references=labels,average='weighted')['f1']\n\n    return {\n      'accuracy': acc,\n      'f1': f1\n  }","metadata":{"execution":{"iopub.status.busy":"2024-04-25T13:00:20.979761Z","iopub.execute_input":"2024-04-25T13:00:20.980089Z","iopub.status.idle":"2024-04-25T13:00:23.185896Z","shell.execute_reply.started":"2024-04-25T13:00:20.980058Z","shell.execute_reply":"2024-04-25T13:00:23.185117Z"},"trusted":true},"outputs":[{"output_type":"display_data","data":{"text/plain":"Downloading builder script:   0%|          | 0.00/4.20k [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"889130d081b04b01abbf92e85cb45a0a"}},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"Downloading builder script:   0%|          | 0.00/6.77k [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"926cf05c58e748eb9762f6f48eeee714"}},"metadata":{}}],"execution_count":11},{"cell_type":"code","source":"model_checkpoint = \"MIT/ast-finetuned-audioset-10-10-0.4593\"\nbatch_size = 32\n\nfrom transformers import AutoFeatureExtractor\n\nfeature_extractor = AutoFeatureExtractor.from_pretrained(model_checkpoint)\nfeature_extractor","metadata":{"execution":{"iopub.status.busy":"2024-04-25T13:06:23.848074Z","iopub.execute_input":"2024-04-25T13:06:23.849232Z","iopub.status.idle":"2024-04-25T13:06:24.103144Z","shell.execute_reply.started":"2024-04-25T13:06:23.849199Z","shell.execute_reply":"2024-04-25T13:06:24.102277Z"},"trusted":true},"outputs":[{"output_type":"display_data","data":{"text/plain":"preprocessor_config.json:   0%|          | 0.00/297 [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"99cab4c19cb7485db504f037a9a81210"}},"metadata":{}},{"execution_count":12,"output_type":"execute_result","data":{"text/plain":"ASTFeatureExtractor {\n  \"do_normalize\": true,\n  \"feature_extractor_type\": \"ASTFeatureExtractor\",\n  \"feature_size\": 1,\n  \"max_length\": 1024,\n  \"mean\": -4.2677393,\n  \"num_mel_bins\": 128,\n  \"padding_side\": \"right\",\n  \"padding_value\": 0.0,\n  \"return_attention_mask\": false,\n  \"sampling_rate\": 16000,\n  \"std\": 4.5689974\n}"},"metadata":{}}],"execution_count":12},{"cell_type":"code","source":"def sample_or_pad(waveform, wav_len=48000):\n    m, n = waveform.shape\n    if n < wav_len:\n        padded_wav = torch.zeros(1, wav_len)\n        padded_wav[:, :n] = waveform\n        return padded_wav\n    elif n > wav_len:\n        offset = np.random.randint(0, n - wav_len)\n        sampled_wav = waveform[:, offset:offset+wav_len]\n        return sampled_wav\n    else:\n        return waveform\n        \nclass EventDetectionDataset(Dataset):\n    def __init__(self, data_path, x, y=None):\n        self.x = x\n        self.y = y\n        self.data_path = data_path\n        self.masking = torchaudio.transforms.FrequencyMasking(freq_mask_param=128)\n#         self.ms = torchaudio.transforms.MelSpectrogram(16000, n_fft=2048, hop_length = 80, win_length = 800, n_mels=128, normalized=True)\n    \n    def __len__(self):\n        return len(self.x)\n\n    def __getitem__(self, idx):\n        path2wav = os.path.join(self.data_path, self.x[idx])\n        waveform, sample_rate = torchaudio.load(path2wav, normalize=False)\n        waveform = sample_or_pad(waveform).flatten()\n        \n        mel_spec = feature_extractor(\n            waveform, \n            sampling_rate=sample_rate, \n            truncation=True, \n        )['input_values'][0]\n\n        if self.y is not None:\n            mel_spec = self.masking(torch.tensor(mel_spec))\n            return {'input_values':mel_spec, 'label':int(self.y[idx])}\n#             return waveform, self.y[idx]\n        return {\"input_values\":mel_spec}","metadata":{"execution":{"iopub.status.busy":"2024-04-25T13:37:32.733319Z","iopub.execute_input":"2024-04-25T13:37:32.734202Z","iopub.status.idle":"2024-04-25T13:37:32.744168Z","shell.execute_reply.started":"2024-04-25T13:37:32.734169Z","shell.execute_reply":"2024-04-25T13:37:32.743254Z"},"trusted":true},"outputs":[],"execution_count":51},{"cell_type":"code","source":"X_train, X_val, y_train, y_val = train_test_split(df_train.fname.values, df_train.label_encoded.values, stratify=df_train.label_encoded.values, \n                                                  test_size=0.2, random_state=42)\n\ntrain_dataset = EventDetectionDataset(os.path.join(DATA_PATH, train_data_folder), X_train, y_train)\nval_dataset = EventDetectionDataset(os.path.join(DATA_PATH, train_data_folder), X_val, y_val)\n\ntest_dataset = DataLoader(\n                        EventDetectionDataset(os.path.join(DATA_PATH, test_data_folder), df_test.fname.values, None),\n                        batch_size=batch_size, shuffle=False\n                )","metadata":{"execution":{"iopub.status.busy":"2024-04-25T13:37:34.129099Z","iopub.execute_input":"2024-04-25T13:37:34.12947Z","iopub.status.idle":"2024-04-25T13:37:34.142359Z","shell.execute_reply.started":"2024-04-25T13:37:34.129441Z","shell.execute_reply":"2024-04-25T13:37:34.141321Z"},"trusted":true},"outputs":[],"execution_count":52},{"cell_type":"code","source":"sample = next(iter(train_dataset))","metadata":{"execution":{"iopub.status.busy":"2024-04-25T13:37:34.306747Z","iopub.execute_input":"2024-04-25T13:37:34.307475Z","iopub.status.idle":"2024-04-25T13:37:34.319015Z","shell.execute_reply.started":"2024-04-25T13:37:34.307446Z","shell.execute_reply":"2024-04-25T13:37:34.318038Z"},"trusted":true},"outputs":[],"execution_count":53},{"cell_type":"code","source":"masking = torchaudio.transforms.FrequencyMasking(freq_mask_param=128)","metadata":{"execution":{"iopub.status.busy":"2024-04-25T13:36:04.453082Z","iopub.execute_input":"2024-04-25T13:36:04.453791Z","iopub.status.idle":"2024-04-25T13:36:04.458254Z","shell.execute_reply.started":"2024-04-25T13:36:04.453761Z","shell.execute_reply":"2024-04-25T13:36:04.457248Z"},"trusted":true},"outputs":[],"execution_count":43},{"cell_type":"code","source":"torch.sum(torch.tensor(sample['input_values']) != masking(torch.tensor(sample['input_values'])), axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-04-25T13:36:10.725697Z","iopub.execute_input":"2024-04-25T13:36:10.726594Z","iopub.status.idle":"2024-04-25T13:36:10.735999Z","shell.execute_reply.started":"2024-04-25T13:36:10.72656Z","shell.execute_reply":"2024-04-25T13:36:10.735004Z"},"trusted":true},"outputs":[{"execution_count":47,"output_type":"execute_result","data":{"text/plain":"tensor([100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100,\n        100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100,\n        100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100,\n        100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100,\n        100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100,\n        100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100,\n        100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100,\n        100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100,\n        100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100, 100,\n        100, 100])"},"metadata":{}}],"execution_count":47},{"cell_type":"code","source":"sample['input_values'].shap","metadata":{"execution":{"iopub.status.busy":"2024-04-25T13:07:02.818383Z","iopub.execute_input":"2024-04-25T13:07:02.81926Z","iopub.status.idle":"2024-04-25T13:07:02.825102Z","shell.execute_reply.started":"2024-04-25T13:07:02.819229Z","shell.execute_reply":"2024-04-25T13:07:02.824131Z"},"trusted":true},"outputs":[{"execution_count":16,"output_type":"execute_result","data":{"text/plain":"(1024, 128)"},"metadata":{}}],"execution_count":16},{"cell_type":"code","source":"sample['input_values'].type(torch.double)","metadata":{"execution":{"iopub.status.busy":"2024-04-24T20:47:56.47828Z","iopub.execute_input":"2024-04-24T20:47:56.478625Z","iopub.status.idle":"2024-04-24T20:47:56.48781Z","shell.execute_reply.started":"2024-04-24T20:47:56.478596Z","shell.execute_reply":"2024-04-24T20:47:56.486616Z"},"trusted":true},"outputs":[{"execution_count":119,"output_type":"execute_result","data":{"text/plain":"tensor([ 0.0018,  0.0006,  0.0031,  ..., -0.0010, -0.0010, -0.0010],\n       dtype=torch.float64)"},"metadata":{}}],"execution_count":119},{"cell_type":"code","source":"from transformers import AutoModelForAudioClassification, TrainingArguments, Trainer\n\nnum_labels = len(id2label)\nmodel = AutoModelForAudioClassification.from_pretrained(\n    model_checkpoint, \n    num_labels=num_labels,\n    label2id=label2id,\n    id2label=id2label,\n    ignore_mismatched_sizes=True\n)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-25T13:37:40.348423Z","iopub.execute_input":"2024-04-25T13:37:40.348797Z","iopub.status.idle":"2024-04-25T13:37:43.566848Z","shell.execute_reply.started":"2024-04-25T13:37:40.348769Z","shell.execute_reply":"2024-04-25T13:37:43.566067Z"},"trusted":true},"outputs":[{"output_type":"display_data","data":{"text/plain":"config.json:   0%|          | 0.00/26.8k [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"512703d093764bfea20d037fe2f30564"}},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"model.safetensors:   0%|          | 0.00/346M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"e3c95e90f7574f9eb3c86c3e0baa7dd8"}},"metadata":{}},{"name":"stderr","text":"Some weights of ASTForAudioClassification were not initialized from the model checkpoint at MIT/ast-finetuned-audioset-10-10-0.4593 and are newly initialized because the shapes did not match:\n- classifier.dense.bias: found shape torch.Size([527]) in the checkpoint and torch.Size([41]) in the model instantiated\n- classifier.dense.weight: found shape torch.Size([527, 768]) in the checkpoint and torch.Size([41, 768]) in the model instantiated\nYou should probably TRAIN this model on a down-stream task to be able to use it for predictions and inference.\n","output_type":"stream"}],"execution_count":54},{"cell_type":"code","source":"batch_size = 8\n\nmodel_name = model_checkpoint.split(\"/\")[-1]\n\nargs = TrainingArguments(\n    f\"{model_name}-finetuned-ks\",\n    evaluation_strategy = \"steps\",\n    save_strategy = \"steps\",\n    eval_steps = 50,\n    save_steps = 50,\n    save_total_limit=1,\n    learning_rate=3e-5,\n    per_device_train_batch_size=batch_size,\n    gradient_accumulation_steps=4,\n    per_device_eval_batch_size=batch_size,\n    num_train_epochs=5,\n    warmup_ratio=0.1,\n    logging_steps=10,\n    load_best_model_at_end=True,\n    metric_for_best_model=\"f1\"\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-25T13:38:05.428289Z","iopub.execute_input":"2024-04-25T13:38:05.429134Z","iopub.status.idle":"2024-04-25T13:38:05.50477Z","shell.execute_reply.started":"2024-04-25T13:38:05.429103Z","shell.execute_reply":"2024-04-25T13:38:05.504014Z"},"trusted":true},"outputs":[],"execution_count":55},{"cell_type":"code","source":"from transformers import DefaultDataCollator\ndata_collator = DefaultDataCollator()\n\ntrainer = Trainer(\n    model,\n    args,\n    train_dataset=train_dataset,\n    eval_dataset=val_dataset,\n    compute_metrics=compute_metrics,\n    data_collator=data_collator\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-25T13:38:07.726436Z","iopub.execute_input":"2024-04-25T13:38:07.727684Z","iopub.status.idle":"2024-04-25T13:38:08.567428Z","shell.execute_reply.started":"2024-04-25T13:38:07.727642Z","shell.execute_reply":"2024-04-25T13:38:08.566679Z"},"trusted":true},"outputs":[{"name":"stderr","text":"/opt/conda/lib/python3.10/site-packages/accelerate/accelerator.py:436: FutureWarning: Passing the following arguments to `Accelerator` is deprecated and will be removed in version 1.0 of Accelerate: dict_keys(['dispatch_batches', 'split_batches', 'even_batches', 'use_seedable_sampler']). Please pass an `accelerate.DataLoaderConfiguration` instead: \ndataloader_config = DataLoaderConfiguration(dispatch_batches=None, split_batches=False, even_batches=True, use_seedable_sampler=True)\n  warnings.warn(\n","output_type":"stream"}],"execution_count":56},{"cell_type":"code","source":"trainer.train()","metadata":{"execution":{"iopub.status.busy":"2024-04-25T13:38:08.838308Z","iopub.execute_input":"2024-04-25T13:38:08.838708Z"},"trusted":true},"outputs":[{"name":"stderr","text":"\u001b[34m\u001b[1mwandb\u001b[0m: Logging into wandb.ai. (Learn how to deploy a W&B server locally: https://wandb.me/wandb-server)\n\u001b[34m\u001b[1mwandb\u001b[0m: You can find your API key in your browser here: https://wandb.ai/authorize\n\u001b[34m\u001b[1mwandb\u001b[0m: Paste an API key from your profile and hit enter, or press ctrl+c to quit:","output_type":"stream"},{"output_type":"stream","name":"stdin","text":"  ········································\n"},{"name":"stderr","text":"\u001b[34m\u001b[1mwandb\u001b[0m: Appending key for api.wandb.ai to your netrc file: /root/.netrc\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"<IPython.core.display.HTML object>","text/html":"Tracking run with wandb version 0.16.6"},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"<IPython.core.display.HTML object>","text/html":"Run data is saved locally in <code>/kaggle/working/wandb/run-20240425_133832-re5wl238</code>"},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"<IPython.core.display.HTML object>","text/html":"Syncing run <strong><a href='https://wandb.ai/saltandmatches/huggingface/runs/re5wl238' target=\"_blank\">blooming-oath-4</a></strong> to <a href='https://wandb.ai/saltandmatches/huggingface' target=\"_blank\">Weights & Biases</a> (<a href='https://wandb.me/run' target=\"_blank\">docs</a>)<br/>"},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"<IPython.core.display.HTML object>","text/html":" View project at <a href='https://wandb.ai/saltandmatches/huggingface' target=\"_blank\">https://wandb.ai/saltandmatches/huggingface</a>"},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"<IPython.core.display.HTML object>","text/html":" View run at <a href='https://wandb.ai/saltandmatches/huggingface/runs/re5wl238' target=\"_blank\">https://wandb.ai/saltandmatches/huggingface/runs/re5wl238</a>"},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"<IPython.core.display.HTML object>","text/html":"\n    <div>\n      \n      <progress value='551' max='710' style='width:300px; height:20px; vertical-align: middle;'></progress>\n      [551/710 1:00:33 < 17:32, 0.15 it/s, Epoch 3.87/5]\n    </div>\n    <table border=\"1\" class=\"dataframe\">\n  <thead>\n <tr style=\"text-align: left;\">\n      <th>Step</th>\n      <th>Training Loss</th>\n      <th>Validation Loss</th>\n      <th>Accuracy</th>\n      <th>F1</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <td>50</td>\n      <td>2.299800</td>\n      <td>2.083433</td>\n      <td>0.535620</td>\n      <td>0.495681</td>\n    </tr>\n    <tr>\n      <td>100</td>\n      <td>1.132700</td>\n      <td>1.121706</td>\n      <td>0.709763</td>\n      <td>0.709092</td>\n    </tr>\n    <tr>\n      <td>150</td>\n      <td>0.758200</td>\n      <td>0.884359</td>\n      <td>0.779244</td>\n      <td>0.779432</td>\n    </tr>\n    <tr>\n      <td>200</td>\n      <td>0.553800</td>\n      <td>0.836777</td>\n      <td>0.785400</td>\n      <td>0.784905</td>\n    </tr>\n    <tr>\n      <td>250</td>\n      <td>0.553600</td>\n      <td>0.780665</td>\n      <td>0.791557</td>\n      <td>0.790886</td>\n    </tr>\n    <tr>\n      <td>300</td>\n      <td>0.358100</td>\n      <td>0.730963</td>\n      <td>0.805629</td>\n      <td>0.805413</td>\n    </tr>\n    <tr>\n      <td>350</td>\n      <td>0.405800</td>\n      <td>0.702029</td>\n      <td>0.815303</td>\n      <td>0.814444</td>\n    </tr>\n    <tr>\n      <td>400</td>\n      <td>0.413400</td>\n      <td>0.718459</td>\n      <td>0.806508</td>\n      <td>0.807530</td>\n    </tr>\n    <tr>\n      <td>450</td>\n      <td>0.281100</td>\n      <td>0.710345</td>\n      <td>0.818821</td>\n      <td>0.816584</td>\n    </tr>\n    <tr>\n      <td>500</td>\n      <td>0.335300</td>\n      <td>0.659199</td>\n      <td>0.827617</td>\n      <td>0.828165</td>\n    </tr>\n  </tbody>\n</table><p>\n    <div>\n      \n      <progress value='134' max='143' style='width:300px; height:20px; vertical-align: middle;'></progress>\n      [134/143 01:04 < 00:04, 2.08 it/s]\n    </div>\n    "},"metadata":{}},{"name":"stderr","text":"Some non-default generation parameters are set in the model config. These should go into a GenerationConfig file (https://huggingface.co/docs/transformers/generation_strategies#save-a-custom-decoding-strategy-with-your-model) instead. This warning will be raised to an exception in v4.41.\nNon-default generation parameters: {'max_length': 1024}\nSome non-default generation parameters are set in the model config. These should go into a GenerationConfig file (https://huggingface.co/docs/transformers/generation_strategies#save-a-custom-decoding-strategy-with-your-model) instead. This warning will be raised to an exception in v4.41.\nNon-default generation parameters: {'max_length': 1024}\nSome non-default generation parameters are set in the model config. These should go into a GenerationConfig file (https://huggingface.co/docs/transformers/generation_strategies#save-a-custom-decoding-strategy-with-your-model) instead. This warning will be raised to an exception in v4.41.\nNon-default generation parameters: {'max_length': 1024}\nSome non-default generation parameters are set in the model config. These should go into a GenerationConfig file (https://huggingface.co/docs/transformers/generation_strategies#save-a-custom-decoding-strategy-with-your-model) instead. This warning will be raised to an exception in v4.41.\nNon-default generation parameters: {'max_length': 1024}\nSome non-default generation parameters are set in the model config. These should go into a GenerationConfig file (https://huggingface.co/docs/transformers/generation_strategies#save-a-custom-decoding-strategy-with-your-model) instead. This warning will be raised to an exception in v4.41.\nNon-default generation parameters: {'max_length': 1024}\nSome non-default generation parameters are set in the model config. These should go into a GenerationConfig file (https://huggingface.co/docs/transformers/generation_strategies#save-a-custom-decoding-strategy-with-your-model) instead. This warning will be raised to an exception in v4.41.\nNon-default generation parameters: {'max_length': 1024}\nSome non-default generation parameters are set in the model config. These should go into a GenerationConfig file (https://huggingface.co/docs/transformers/generation_strategies#save-a-custom-decoding-strategy-with-your-model) instead. This warning will be raised to an exception in v4.41.\nNon-default generation parameters: {'max_length': 1024}\nSome non-default generation parameters are set in the model config. These should go into a GenerationConfig file (https://huggingface.co/docs/transformers/generation_strategies#save-a-custom-decoding-strategy-with-your-model) instead. This warning will be raised to an exception in v4.41.\nNon-default generation parameters: {'max_length': 1024}\nSome non-default generation parameters are set in the model config. These should go into a GenerationConfig file (https://huggingface.co/docs/transformers/generation_strategies#save-a-custom-decoding-strategy-with-your-model) instead. This warning will be raised to an exception in v4.41.\nNon-default generation parameters: {'max_length': 1024}\nSome non-default generation parameters are set in the model config. These should go into a GenerationConfig file (https://huggingface.co/docs/transformers/generation_strategies#save-a-custom-decoding-strategy-with-your-model) instead. This warning will be raised to an exception in v4.41.\nNon-default generation parameters: {'max_length': 1024}\n","output_type":"stream"}],"execution_count":null},{"cell_type":"code","source":"test_dataset = DataLoader(\n                        EventDetectionDataset(os.path.join(DATA_PATH, test_data_folder), df_test.fname.values, None),\n                        batch_size=8, shuffle=False\n                )","metadata":{"execution":{"iopub.status.busy":"2024-04-24T23:59:02.666407Z","iopub.execute_input":"2024-04-24T23:59:02.666801Z","iopub.status.idle":"2024-04-24T23:59:02.674374Z","shell.execute_reply.started":"2024-04-24T23:59:02.666771Z","shell.execute_reply":"2024-04-24T23:59:02.673122Z"},"trusted":true},"outputs":[],"execution_count":180},{"cell_type":"code","source":"sample = next(iter(test_dataset))","metadata":{"execution":{"iopub.status.busy":"2024-04-24T23:59:19.03085Z","iopub.execute_input":"2024-04-24T23:59:19.031215Z","iopub.status.idle":"2024-04-24T23:59:19.095086Z","shell.execute_reply.started":"2024-04-24T23:59:19.031186Z","shell.execute_reply":"2024-04-24T23:59:19.09381Z"},"trusted":true},"outputs":[],"execution_count":181},{"cell_type":"code","source":"sample.shape","metadata":{"execution":{"iopub.status.busy":"2024-04-24T23:59:29.362116Z","iopub.execute_input":"2024-04-24T23:59:29.362471Z","iopub.status.idle":"2024-04-24T23:59:29.370237Z","shell.execute_reply.started":"2024-04-24T23:59:29.362446Z","shell.execute_reply":"2024-04-24T23:59:29.369271Z"},"trusted":true},"outputs":[{"execution_count":183,"output_type":"execute_result","data":{"text/plain":"torch.Size([8, 1024, 128])"},"metadata":{}}],"execution_count":183},{"cell_type":"code","source":"from transformers import EarlyStoppingCallback\nfrom transformers import TrainingArguments, Trainer\nfrom transformers import DefaultDataCollator\n\ndata_collator = DefaultDataCollator()\nmodel = AutoModelForImageClassification.from_pretrained(\"microsoft/resnet-18\", num_channels=1, num_labels=41, label2id = label2id, id2label=id2label, problem_type=\"single_label_classification\", ignore_mismatched_sizes=True)\n\ntraining_args = TrainingArguments(\n    output_dir=f\"resnet18\",\n    learning_rate=2e-5,\n    num_train_epochs=100,\n    weight_decay=0.01,\n    evaluation_strategy=\"epoch\",\n    save_strategy=\"epoch\",\n    save_total_limit=1,\n    metric_for_best_model='f1',\n    load_best_model_at_end=True,\n    report_to=None,\n    per_device_train_batch_size=32,\n    per_device_eval_batch_size=16,\n)\n\ntrainer = Trainer(\n    model=model,\n    args=training_args,\n    train_dataset=EventDetectionDataset(os.path.join(DATA_PATH, train_data_folder), X_train, y_train),\n    eval_dataset=EventDetectionDataset(os.path.join(DATA_PATH, train_data_folder), X_val, y_val),\n    compute_metrics=compute_metrics,\n    data_collator=data_collator,\n    callbacks = [EarlyStoppingCallback(early_stopping_patience=5)],\n)\n\ntrainer.train()","metadata":{"execution":{"iopub.status.busy":"2024-04-24T20:47:28.745692Z","iopub.status.idle":"2024-04-24T20:47:28.746231Z","shell.execute_reply.started":"2024-04-24T20:47:28.745954Z","shell.execute_reply":"2024-04-24T20:47:28.745977Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"n_epoch = 100\nbest_f1 = 0\nfor epoch in range(n_epoch):\n    model.train()\n    for wavs, labs in tqdm(train_loader):\n        optimizer.zero_grad()\n        wavs, labs = wavs.cuda(), labs.cuda()\n        outputs = model(wavs)\n        loss = criterion(outputs, labs)\n        loss.backward()\n        optimizer.step()\n#     if epoch % 10 == 0:\n    f1 = eval_model(model, val_loader)\n    f1_train = eval_model(model, train_loader)\n#     f1_train = \"Null\"\n    print(f'epoch: {epoch}, f1_test: {f1}, f1_train: {f1_train}')\n    if f1 > best_f1:\n        best_f1 = f1\n        torch.save(model.state_dict(), '../baseline_fulldiv.pt')\n        \n    lr = lr * 0.95\n    for param_group in optimizer.param_groups:\n        param_group['lr'] = lr","metadata":{"execution":{"iopub.status.busy":"2024-04-24T20:47:28.748469Z","iopub.status.idle":"2024-04-24T20:47:28.748977Z","shell.execute_reply.started":"2024-04-24T20:47:28.748704Z","shell.execute_reply":"2024-04-24T20:47:28.748723Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions = trainer.predict(EventDetectionDataset(os.path.join(DATA_PATH, test_data_folder), df_test.fname.values, None))","metadata":{"execution":{"iopub.status.busy":"2024-04-25T00:13:41.932005Z","iopub.execute_input":"2024-04-25T00:13:41.932418Z","iopub.status.idle":"2024-04-25T00:17:28.505238Z","shell.execute_reply.started":"2024-04-25T00:13:41.932388Z","shell.execute_reply":"2024-04-25T00:17:28.504225Z"},"trusted":true},"outputs":[{"output_type":"display_data","data":{"text/plain":"<IPython.core.display.HTML object>","text/html":""},"metadata":{}}],"execution_count":198},{"cell_type":"code","source":"pred_ids = np.argmax(predictions.predictions, axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-04-25T00:17:28.507529Z","iopub.execute_input":"2024-04-25T00:17:28.507948Z","iopub.status.idle":"2024-04-25T00:17:28.514429Z","shell.execute_reply.started":"2024-04-25T00:17:28.507913Z","shell.execute_reply":"2024-04-25T00:17:28.513299Z"},"trusted":true},"outputs":[],"execution_count":199},{"cell_type":"code","source":"pred_ids","metadata":{"execution":{"iopub.status.busy":"2024-04-25T00:09:58.781207Z","iopub.execute_input":"2024-04-25T00:09:58.781582Z","iopub.status.idle":"2024-04-25T00:09:58.790217Z","shell.execute_reply.started":"2024-04-25T00:09:58.781554Z","shell.execute_reply":"2024-04-25T00:09:58.789169Z"},"trusted":true},"outputs":[{"execution_count":193,"output_type":"execute_result","data":{"text/plain":"array([12, 32, 13, ...,  9, 16, 25])"},"metadata":{}}],"execution_count":193},{"cell_type":"code","source":"pred_ids","metadata":{"execution":{"iopub.status.busy":"2024-04-25T00:17:28.515813Z","iopub.execute_input":"2024-04-25T00:17:28.516558Z","iopub.status.idle":"2024-04-25T00:17:28.527443Z","shell.execute_reply.started":"2024-04-25T00:17:28.51652Z","shell.execute_reply":"2024-04-25T00:17:28.526329Z"},"trusted":true},"outputs":[{"execution_count":200,"output_type":"execute_result","data":{"text/plain":"array([12,  2, 12, ...,  9, 16, 25])"},"metadata":{}}],"execution_count":200},{"cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"decoder = {classes_dict[cl]:cl for cl in classes_dict}\nforecast1 = pd.Series(pred_ids).map(decoder)","metadata":{"execution":{"iopub.status.busy":"2024-04-25T00:18:02.734328Z","iopub.execute_input":"2024-04-25T00:18:02.735241Z","iopub.status.idle":"2024-04-25T00:18:02.742646Z","shell.execute_reply.started":"2024-04-25T00:18:02.735207Z","shell.execute_reply":"2024-04-25T00:18:02.741475Z"},"trusted":true},"outputs":[],"execution_count":203},{"cell_type":"code","source":"forecast == forecast1","metadata":{"execution":{"iopub.status.busy":"2024-04-25T00:18:05.061218Z","iopub.execute_input":"2024-04-25T00:18:05.061967Z","iopub.status.idle":"2024-04-25T00:18:05.072916Z","shell.execute_reply.started":"2024-04-25T00:18:05.061932Z","shell.execute_reply":"2024-04-25T00:18:05.071895Z"},"trusted":true},"outputs":[{"execution_count":204,"output_type":"execute_result","data":{"text/plain":"0        True\n1       False\n2       False\n3        True\n4        True\n        ...  \n3785     True\n3786     True\n3787     True\n3788     True\n3789     True\nLength: 3790, dtype: bool"},"metadata":{}}],"execution_count":204},{"cell_type":"code","source":"forecast = forecast1","metadata":{"execution":{"iopub.status.busy":"2024-04-25T00:18:25.620147Z","iopub.execute_input":"2024-04-25T00:18:25.620546Z","iopub.status.idle":"2024-04-25T00:18:25.626331Z","shell.execute_reply.started":"2024-04-25T00:18:25.620515Z","shell.execute_reply":"2024-04-25T00:18:25.625185Z"},"trusted":true},"outputs":[],"execution_count":205},{"cell_type":"code","source":"df_test['label'] = forecast\ndf_test.to_csv(f'{model_name}.csv', index=None)","metadata":{"execution":{"iopub.status.busy":"2024-04-25T00:18:26.908423Z","iopub.execute_input":"2024-04-25T00:18:26.909165Z","iopub.status.idle":"2024-04-25T00:18:26.925908Z","shell.execute_reply.started":"2024-04-25T00:18:26.90913Z","shell.execute_reply":"2024-04-25T00:18:26.924886Z"},"trusted":true},"outputs":[],"execution_count":206},{"cell_type":"code","source":"# make a model\nmodel_name = 'baseline_fulldiv.pt'\nmodel = BaseLineModel().cuda()\nmodel.load_state_dict(torch.load(os.path.join('..', model_name)))\nmodel.eval()\nforecast = []\nwith torch.no_grad():\n    for wavs in tqdm(test_loader):\n        wavs = wavs.cuda()\n        outputs = model.inference(wavs)\n        outputs = outputs.detach().cpu().numpy().argmax(axis=1)\n        forecast.append(outputs)\nforecast = [x for sublist in forecast for x in sublist]\ndecoder = {classes_dict[cl]:cl for cl in classes_dict}\nforecast = pd.Series(forecast).map(decoder)\ndf_test['label'] = forecast\ndf_test.to_csv(f'{model_name}.csv', index=None)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null}]}