{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"!pip install torch-audiomentations\n!pip install efficientnet-pytorch","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2021-07-01T18:20:48.679522Z","iopub.execute_input":"2021-07-01T18:20:48.679838Z","iopub.status.idle":"2021-07-01T18:21:05.405707Z","shell.execute_reply.started":"2021-07-01T18:20:48.679766Z","shell.execute_reply":"2021-07-01T18:21:05.404784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Warnings\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Python standar library\nimport os\nimport pickle\nimport glob\nimport copy\nimport tqdm\nimport random\n\n# DS tools\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\n\n# Scikit-learn\nfrom sklearn.metrics import f1_score\nfrom sklearn.model_selection import train_test_split\n\n# Common pytorch\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\n# EfficientNet\nfrom efficientnet_pytorch import EfficientNet\nfrom efficientnet_pytorch.utils import Conv2dStaticSamePadding\n\n# Torchaudio\nimport torchaudio\nimport torchaudio.transforms as T\nfrom torch_audiomentations import *","metadata":{"execution":{"iopub.status.busy":"2021-07-01T18:21:05.409054Z","iopub.execute_input":"2021-07-01T18:21:05.409331Z","iopub.status.idle":"2021-07-01T18:21:08.917498Z","shell.execute_reply.started":"2021-07-01T18:21:05.409302Z","shell.execute_reply":"2021-07-01T18:21:08.916569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Constants","metadata":{}},{"cell_type":"code","source":"# Randomness\nSEED = 451\n\n# Paths\nDATA_DIR = '/kaggle/input/itmo-acoustic-event-detection-2021/'\nTRAIN_DIR = os.path.join(DATA_DIR, 'audio_train', 'train')\nTEST_DIR = os.path.join(DATA_DIR, 'audio_test', 'test')\n\n# Features \nSAMPLE_RATE = 16000\nFFT_SIZE = 1024\nWIN_LEN = 512\nHOP_LEN = WIN_LEN // 2\nN_MELS = 128\n\n# Dataset\nBATCH_SIZE = 64\nNUM_CLASSES = 41\n\n# Models parameters\nEMB_SIZE = 256\nLEARNING_RATE = 1e-3\nN_EPOCHS = 20\nDEVICE = 'cuda:0'","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:00:03.021195Z","iopub.execute_input":"2021-07-01T06:00:03.021526Z","iopub.status.idle":"2021-07-01T06:00:03.02751Z","shell.execute_reply.started":"2021-07-01T06:00:03.021496Z","shell.execute_reply":"2021-07-01T06:00:03.026325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Setup randomness","metadata":{}},{"cell_type":"code","source":"random.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:00:03.625979Z","iopub.execute_input":"2021-07-01T06:00:03.626372Z","iopub.status.idle":"2021-07-01T06:00:03.633032Z","shell.execute_reply.started":"2021-07-01T06:00:03.626339Z","shell.execute_reply":"2021-07-01T06:00:03.632196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data preprocessing","metadata":{}},{"cell_type":"code","source":"sample_submission_df = pd.read_csv(os.path.join(DATA_DIR, 'sample_submission.csv'))\ntrain_df = pd.read_csv(os.path.join(DATA_DIR, 'train.csv'))\n\nprint(\"Train dir len:\", len(os.listdir(TRAIN_DIR)))\nprint(\"Test dir len:\", len(os.listdir(TEST_DIR)))\nprint(\"Sample submission df shape:\", sample_submission_df.shape)\nprint(\"Train df shape:\", train_df.shape)","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:00:03.934503Z","iopub.execute_input":"2021-07-01T06:00:03.934818Z","iopub.status.idle":"2021-07-01T06:00:03.966084Z","shell.execute_reply.started":"2021-07-01T06:00:03.934789Z","shell.execute_reply":"2021-07-01T06:00:03.965062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Unique labels:\", len(train_df.label.unique()))\n\n# Dummy\nsample_submission_df['label_encoded'] = sample_submission_df['label']\n\nclasses_dict = {cl: i for i, cl in enumerate(train_df.label.unique())}\nidx_to_class_dict = {val: key for key, val in classes_dict.items()}\ntrain_df['label_encoded'] = train_df['label'].apply(lambda x: classes_dict[x])\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:00:04.021352Z","iopub.execute_input":"2021-07-01T06:00:04.02164Z","iopub.status.idle":"2021-07-01T06:00:04.051861Z","shell.execute_reply.started":"2021-07-01T06:00:04.021613Z","shell.execute_reply":"2021-07-01T06:00:04.050994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Divide train into train and eval subsets with saved labels distributions\ntrain_df, eval_df = train_test_split(train_df, train_size=0.85, stratify=train_df['label_encoded'])\nprint(\"Train shape: %s\\nEval shape: %s\" % (train_df.shape, eval_df.shape))","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:00:04.161479Z","iopub.execute_input":"2021-07-01T06:00:04.161778Z","iopub.status.idle":"2021-07-01T06:00:04.178604Z","shell.execute_reply.started":"2021-07-01T06:00:04.16175Z","shell.execute_reply":"2021-07-01T06:00:04.177482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"def find_most_loud_period(wav, segment_size):\n    \"\"\" Find loudest segment in audio \"\"\"\n    \n    result = None\n    if wav.shape[1] < segment_size:\n        num_repeats = np.ceil(segment_size / wav.shape[1]).astype(int)\n        wav = wav.repeat((1, num_repeats))[:, :segment_size]\n        result = wav\n    elif wav.shape[1] == segment_size:\n        result = wav\n    else:\n        loud_idx = wav.argmax(axis=1)\n        if loud_idx < segment_size // 2:\n            result = wav[:, :segment_size]\n        elif loud_idx > wav.shape[1] - segment_size // 2:\n            result = wav[:, -segment_size:]\n        else:\n            result = wav[:, loud_idx - segment_size // 2 : loud_idx + segment_size // 2]\n    return result\n\n\ndef get_class_weights(y):\n    \"\"\" Compute class weights (least frequent class -> bigger weight) \"\"\"\n    \n    _, counts = np.unique(y, return_counts=True)\n    weights = counts / y.shape[0]\n    weights = weights.min() / weights\n    weights = np.array([weights[int(y_)] for y_ in y])\n    return weights\n\n\ndef apply_aug_transform(wav, sample_rate):\n    \"\"\" Apply random transformations to input wav \"\"\"\n    \n    apply_augmentation = Compose(transforms=[\n        AddColoredNoise(min_snr_in_db=-10, max_snr_in_db=10, p=0.3),\n        Gain(min_gain_in_db=-15.0, max_gain_in_db=5.0, p=0.3),\n        Shift(min_shift=-0.1, max_shift=0.1, shift_unit=\"seconds\", p=0.3),\n        \n    ], shuffle=True)\n    wav = apply_augmentation(wav, sample_rate=sample_rate)\n    return wav\n\n\ndef init_weights(module):\n    \"\"\" Weights and biases initialization \"\"\"\n    \n    if type(module) == nn.Linear or type(module) == nn.Conv2d:\n        torch.nn.init.xavier_normal_(module.weight)\n    if type(module) == nn.Linear and module.bias is not None:\n        module.bias.data.fill_(0.01)\n                \n\ndef plot_history(history, title, ylabel):\n    plt.figure(figsize=(10, 6))\n    plt.title(title)\n    plt.ylabel(ylabel)\n    plt.xlabel(\"Epoch\")\n    plt.plot(history['train_loss'], label='train')\n    plt.plot(history['eval_loss'], label='eval')\n    plt.xticks(np.arange(len(history['train_loss'])), np.arange(len(history['train_loss'])))\n    plt.legend()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:00:04.438046Z","iopub.execute_input":"2021-07-01T06:00:04.438396Z","iopub.status.idle":"2021-07-01T06:00:04.452928Z","shell.execute_reply.started":"2021-07-01T06:00:04.438368Z","shell.execute_reply":"2021-07-01T06:00:04.451795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AED_Dataset(Dataset):\n    \"\"\" Torch dataset-wrapper class for extracted features \"\"\"\n    \n    def __init__(self, path_to_wav_dir, x: np.ndarray, y: np.ndarray, augment_factor=1):\n        self.path_to_wav_dir = path_to_wav_dir\n        self.x = self.x = np.array(list(x) * augment_factor)\n        self.y = np.array(list(y) * augment_factor)\n        self.augment_factor = augment_factor\n    \n    def __getitem__(self, idx):\n        wav, sr = torchaudio.load(\n            os.path.join(self.path_to_wav_dir, self.x[idx]), normalize=True)\n        if idx % self.augment_factor != 0:\n            wav = apply_aug_transform(wav.unsqueeze(0), sr).squeeze(0)\n        \n        # Find loudest timestamp and get 2 seconds around it\n        wav = find_most_loud_period(wav, segment_size=sr * 3)\n        return self.x[idx], wav, self.y[idx]\n    \n    def __len__(self):\n        return len(self.y)","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:00:04.709682Z","iopub.execute_input":"2021-07-01T06:00:04.710014Z","iopub.status.idle":"2021-07-01T06:00:04.717566Z","shell.execute_reply.started":"2021-07-01T06:00:04.709978Z","shell.execute_reply":"2021-07-01T06:00:04.716431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = AED_Dataset(\n    TRAIN_DIR, \n    train_df['fname'], \n    train_df['label_encoded'],\n    augment_factor=10,\n)\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=2)\n\ntest_dataset = AED_Dataset(\n    TEST_DIR, \n    sample_submission_df['fname'], \n    sample_submission_df['label_encoded'],\n)\ntest_loader = DataLoader(test_dataset, batch_size=128)\n\neval_dataset = AED_Dataset(\n    TRAIN_DIR, \n    eval_df['fname'], \n    eval_df['label_encoded']\n)\neval_loader = DataLoader(eval_dataset, batch_size=128)","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:00:04.815301Z","iopub.execute_input":"2021-07-01T06:00:04.815625Z","iopub.status.idle":"2021-07-01T06:00:04.845757Z","shell.execute_reply.started":"2021-07-01T06:00:04.815597Z","shell.execute_reply":"2021-07-01T06:00:04.844725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Models implementation","metadata":{}},{"cell_type":"code","source":"class AMSoftmaxLoss(nn.Module):\n    def __init__(self, in_features, out_features, s=30.0, m=0.4):\n        \"\"\" AM Softmax Loss \"\"\"\n        \n        super(AMSoftmaxLoss, self).__init__()\n        self.s = s\n        self.m = m\n        self.in_features = in_features\n        self.out_features = out_features\n        self.fc = nn.Linear(in_features, out_features, bias=False)\n\n    def forward(self, x, labels):        \n        assert len(x) == len(labels)\n        assert torch.min(labels) >= 0\n        assert torch.max(labels) < self.out_features\n        \n        for W in self.fc.parameters():\n            W = F.normalize(W, dim=1)\n\n        x = F.normalize(x, dim=1)\n\n        wf = self.fc(x)\n        numerator = self.s * (torch.diagonal(wf.transpose(0, 1)[labels]) - self.m)\n        excl = torch.cat([torch.cat((wf[i, :y], wf[i, y+1:])).unsqueeze(0) for i, y in enumerate(labels)], dim=0)\n        denominator = torch.exp(numerator) + torch.sum(torch.exp(self.s * excl), dim=1)\n        L = numerator - torch.log(denominator)\n        return -torch.mean(L)","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:00:05.223229Z","iopub.execute_input":"2021-07-01T06:00:05.22355Z","iopub.status.idle":"2021-07-01T06:00:05.233883Z","shell.execute_reply.started":"2021-07-01T06:00:05.22352Z","shell.execute_reply":"2021-07-01T06:00:05.232954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EfficientNet_AED_Extractor(nn.Module):\n    def __init__(self, efficient_net_version, n_classes, sample_rate, n_fft, win_length, hop_length, n_mels):\n        super(EfficientNet_AED_Extractor, self).__init__()\n        self.ms = T.MelSpectrogram(sample_rate=sample_rate, n_fft=n_fft, \n                                   win_length=win_length, hop_length=hop_length, \n                                   n_mels=n_mels, normalized=True)\n        self.to_db = T.AmplitudeToDB()\n        \n        assert efficient_net_version is not None\n        self.efficientnet_model = EfficientNet.from_pretrained(\n            'efficientnet-%s' % efficient_net_version\n        )\n        self.efficientnet_model._conv_stem = Conv2dStaticSamePadding(\n            1, 32, kernel_size=(3, 3), stride=(2, 2), bias=False, image_size=(128, 188)\n        )\n        self.fc = nn.Linear(1000, EMB_SIZE, bias=True)\n        self.dropout = nn.Dropout(0.3)\n        self.am_softmax_loss = AMSoftmaxLoss(EMB_SIZE, n_classes, s=10.0, m=0.5)\n        \n    def forward(self, x, labels, return_type='loss'):\n        x = self.ms(x)\n        x = self.to_db(x)\n        x = F.relu(self.efficientnet_model(x))\n        x = self.dropout(x)\n        x = F.relu(self.fc(x))\n        if return_type == 'emb':\n            return x\n        loss = self.am_softmax_loss(x, labels)\n        if return_type == 'loss':\n            return loss\n        elif return_type == 'both':\n            return x, loss","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:02:08.028622Z","iopub.execute_input":"2021-07-01T06:02:08.029054Z","iopub.status.idle":"2021-07-01T06:02:08.042631Z","shell.execute_reply.started":"2021-07-01T06:02:08.029014Z","shell.execute_reply":"2021-07-01T06:02:08.040333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EfficientNet_AED_Classifier(nn.Module):\n    def __init__(self, n_classes, extractor):\n        super(EfficientNet_AED_Classifier, self).__init__()\n        self.extractor = extractor\n        for param in self.extractor.parameters():\n            param.requires_grad = False\n            \n        self.dropout1 = nn.Dropout(0.3)\n        self.fc1 = nn.Linear(EMB_SIZE, 256, bias=True)\n        self.dropout2 = nn.Dropout(0.2)\n        self.fc2 = nn.Linear(256, n_classes, bias=True)\n        \n    def forward(self, x, extract_emb=False):\n        x = self.extractor(x, None, 'emb')\n        if extract_emb:\n            return x\n        \n        x = self.dropout1(x)\n        x = F.relu(self.fc1(x))\n        x = self.dropout2(x)\n        # x = F.softmax(self.fc2(x))\n        x = self.fc2(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:02:08.437594Z","iopub.execute_input":"2021-07-01T06:02:08.437911Z","iopub.status.idle":"2021-07-01T06:02:08.445294Z","shell.execute_reply.started":"2021-07-01T06:02:08.437882Z","shell.execute_reply":"2021-07-01T06:02:08.444188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# NN models training","metadata":{}},{"cell_type":"code","source":"# Inference functions\n\ndef extract_embeddings(model, data_loader, device=DEVICE):\n    \"\"\" Extract embeddings from dataloader \"\"\"\n    \n    with torch.no_grad():\n        model.eval()\n        model.to(device)\n        fnames, labels = [], torch.Tensor([])\n        embeddings = torch.Tensor([]).to(device)\n        for fname, X, y in tqdm.tqdm(data_loader):\n            X = X.to(device)\n            emb = model(X, extract_emb=True)\n            fnames += list(fname)\n            labels = torch.cat((labels, y), 0)\n            embeddings = torch.cat((embeddings, emb.squeeze()), 0)\n\n        embeddings = embeddings.detach().cpu().numpy()\n        embeddings_info = [{'fname': fname, 'emb': emb, 'label': label} \n                           for fname, emb, label in zip(fnames, embeddings, labels)]\n    return embeddings_info\n\n\ndef predict_classifier(model, data_loader, return_probs=False, device=DEVICE):\n    \"\"\" Get predictions from classifier model \"\"\"\n    \n    model.eval()\n    model.to(device)\n    preds, trues = torch.Tensor([]).to(device), torch.Tensor([])\n    for _, X, y in tqdm.tqdm(data_loader):\n        X = X.to(device)\n        pred = F.softmax(model(X))\n        # pred = model(X)\n        \n        trues = torch.cat((trues, y), 0)\n        if return_probs:\n            preds = torch.cat((preds, pred.squeeze()), 0)\n        else:\n            preds = torch.cat((preds, pred.argmax(axis=1).squeeze()), 0)\n    \n    trues = trues.detach().cpu().numpy()\n    preds = preds.detach().cpu().numpy()\n    return trues, preds","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:02:08.780212Z","iopub.execute_input":"2021-07-01T06:02:08.780555Z","iopub.status.idle":"2021-07-01T06:02:08.791517Z","shell.execute_reply.started":"2021-07-01T06:02:08.780526Z","shell.execute_reply":"2021-07-01T06:02:08.790656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_one_epoch_extractor(current_epoch, model, train_loader, eval_loader, log_every, device):\n    \"\"\" Training loop for extractor model \"\"\"\n    \n    global best_loss\n    \n    model.train()\n    train_loss = []\n    pbar = tqdm.tqdm(enumerate(train_loader), total=len(train_loader))\n    for i, (_, batch_x, batch_y) in pbar:\n        batch_x = batch_x.to(device)\n        batch_y = batch_y.to(device)\n        \n        optimizer.zero_grad()\n        loss = model(batch_x, batch_y, return_type='loss')\n        loss.backward()\n        optimizer.step()\n        \n        train_loss.append(loss.item())\n        if i != 0 and i % log_every == 0:\n            pbar.set_description(\"Train Loss: %s\" % np.mean(train_loss))\n\n    # Save loss info\n    model.eval()\n    eval_loss = []\n    for _, batch_x, batch_y in eval_loader:\n        batch_x = batch_x.to(device)\n        batch_y = batch_y.to(device)\n        loss = model(batch_x, batch_y, return_type='loss')\n        eval_loss.append(loss.item())\n    eval_loss = np.mean(eval_loss)\n    train_loss = np.mean(train_loss)\n    \n    if eval_loss < best_loss:\n        print(\"New best eval loss. Saving copy of the model.\")\n        best_loss = eval_loss\n        save_dict = {\n            'epoch': current_epoch,\n            'loss': eval_loss,\n            'model': model.state_dict(),\n            'optimizer': optimizer.state_dict(),\n        }\n        torch.save(save_dict, 'aed_extractor.pth')\n        \n    print(\"Eval loss: %s\\n\" % eval_loss)\n    scheduler.step(eval_loss)\n    \n    return train_loss, eval_loss\n\n\ndef train_one_epoch_classifier(current_epoch, model, train_loader, eval_loader, log_every, device):\n    \"\"\" Train loop for classifier model \"\"\"\n    \n    global best_f1\n    model.train()\n    train_loss = []\n    pbar = tqdm.tqdm(enumerate(train_loader), total=len(train_loader))\n    for i, (_, batch_x, batch_y) in pbar:\n        batch_x = batch_x.to(device)\n        batch_y = batch_y.to(device)\n        \n        optimizer.zero_grad()\n        batch_out = model(batch_x)\n        loss = criterion(batch_out, batch_y)\n        loss.backward()\n        optimizer.step()\n        \n        train_loss.append(loss.item())\n        if i != 0 and i % log_every == 0:\n            pbar.set_description(\"Train Loss: %s\" % np.mean(train_loss))\n            \n    y_true, y_pred = predict_classifier(model, eval_loader, device=device)\n    f1 = f1_score(y_true, y_pred, average='macro')\n    scheduler.step(f1)\n    print(\"Eval F1-score: %s\" % f1)\n    \n    if f1 > best_f1:\n        print(\"New best f1-score. Saving copy of the model.\")\n        save_dict = {\n            'epoch': current_epoch,\n            'f1_score': f1,\n            'model': model.state_dict(),\n            'optimizer': optimizer.state_dict(),\n            'scheduler': scheduler.state_dict(),\n        }\n        torch.save(save_dict, 'aed_classifier.pth')\n        best_f1 = f1\n    \n    return np.mean(train_loss), f1\n        \n\ndef train(model, train_loader, eval_loader, epoch_trainer, n_epochs, device):\n    \"\"\" Training process \"\"\"\n    \n    model.to(device)\n    train_losses, eval_losses = [], []\n    for i in range(n_epochs):\n        print(\"Training: %s/%s epochs.\" % (i + 1, n_epochs))\n        train_loss, eval_loss = epoch_trainer(i, model, train_loader, eval_loader, log_every=2, device=device)\n        train_losses.append(train_loss)\n        eval_losses.append(eval_loss)\n        \n    history = {'train_loss': train_losses, 'eval_loss': eval_losses}\n    return model, history","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:02:09.068079Z","iopub.execute_input":"2021-07-01T06:02:09.068453Z","iopub.status.idle":"2021-07-01T06:02:09.099027Z","shell.execute_reply.started":"2021-07-01T06:02:09.068421Z","shell.execute_reply":"2021-07-01T06:02:09.097354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Extractor training\n*Skip this section if you're using pretrained model*","metadata":{}},{"cell_type":"code","source":"aed_extractor = EfficientNet_AED_Extractor(\n    efficient_net_version='b0',\n    n_classes=41,\n    sample_rate=SAMPLE_RATE,\n    n_mels=N_MELS,\n    win_length=WIN_LEN,\n    hop_length=HOP_LEN,\n    n_fft=FFT_SIZE,\n)\naed_extractor = aed_extractor.apply(init_weights)\n\n# Optimizer params\noptimizer = optim.SGD(aed_extractor.parameters(), lr=1e-3, momentum=0.96)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=1, factor=0.75, mode='min')\n\n# Model chosing parameter\nbest_loss = 10.0\n\n# Model training\naed_extractor, history_ext = train(\n    aed_extractor, \n    train_loader, \n    eval_loader, \n    epoch_trainer=train_one_epoch_extractor,\n    n_epochs=10, \n    device=DEVICE\n)","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:02:10.28861Z","iopub.execute_input":"2021-07-01T06:02:10.288936Z","iopub.status.idle":"2021-07-01T06:14:34.615638Z","shell.execute_reply.started":"2021-07-01T06:02:10.288904Z","shell.execute_reply":"2021-07-01T06:14:34.612634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history(\n    history_ext, \n    title=\"Loss functions after training (embedding extractor)\",\n    ylabel=\"AMSoftmaxLoss\"\n)","metadata":{"execution":{"iopub.status.busy":"2021-07-01T06:14:34.616638Z","iopub.status.idle":"2021-07-01T06:14:34.617002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Classifier training\n*Skip this section if you're using pretrained model*","metadata":{}},{"cell_type":"code","source":"# Load extractor best checkpoint\naed_extractor = EfficientNet_AED_Extractor('b0', 41, SAMPLE_RATE, n_mels=N_MELS, win_length=WIN_LEN, hop_length=HOP_LEN, n_fft=FFT_SIZE)\naed_extractor.load_state_dict(torch.load('/kaggle/working/aed_extractor.pth')['model'])\n# aed_extractor.load_state_dict(torch.load('/kaggle/input/efficientnet-b0-amsoftmax-aed-classifier/aed_extractor.pth')['model'])\naed_classifier = EfficientNet_AED_Classifier(n_classes=41, extractor=aed_extractor)\n\n# Optimizer params\noptimizer = optim.RMSprop(aed_classifier.parameters(), lr=LEARNING_RATE, momentum=0.1)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=1, factor=0.8, mode='max')\n\n# Train parameters for classifier model\ncriterion = nn.CrossEntropyLoss()\n\n# Model chosing parameter\nbest_f1 = 0\n\n# Model training\naed_classifier, history_clf = train(\n    aed_classifier, \n    train_loader, \n    eval_loader, \n    epoch_trainer=train_one_epoch_classifier,\n    n_epochs=20, \n    device=DEVICE\n)","metadata":{"execution":{"iopub.status.busy":"2021-06-30T00:42:43.669467Z","iopub.execute_input":"2021-06-30T00:42:43.669788Z","iopub.status.idle":"2021-06-30T00:42:46.357501Z","shell.execute_reply.started":"2021-06-30T00:42:43.669757Z","shell.execute_reply":"2021-06-30T00:42:46.354694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax = plt.subplots(2, 1, figsize=(10, 10))\nax[0].set_title(\"Train loss\")\nax[0].set_xticks(np.arange(len(history_clf['train_loss'])))\nax[0].set_ylabel(\"AMSoftmaxLoss\")\nax[0].plot(history_clf['train_loss'], marker='o')\n\nax[1].set_title(\"Eval f1_score\")\nax[1].set_xlabel(\"Epoch\")\nax[1].set_ylabel(\"f1_score\")\nax[1].set_xticks(np.arange(len(history_clf['eval_loss'])))\nax[1].set_ylim(0.78, 0.83)\nax[1].plot(history_clf['eval_loss'], marker='o');","metadata":{"execution":{"iopub.status.busy":"2021-06-30T00:42:26.677495Z","iopub.status.idle":"2021-06-30T00:42:26.678064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Infer on eval and test sets","metadata":{}},{"cell_type":"code","source":"# Load extractor best checkpoint\nextractor = EfficientNet_AED_Extractor(\n    efficient_net_version='b0',\n    n_classes=41,\n    sample_rate=SAMPLE_RATE,\n    n_mels=N_MELS,\n    win_length=WIN_LEN,\n    hop_length=HOP_LEN,\n    n_fft=FFT_SIZE,\n)\n# extractor.load_state_dict(torch.load('/kaggle/working/aed_extractor.pth')['model'])\naed_classifier = EfficientNet_AED_Classifier(NUM_CLASSES, extractor)\naed_classifier.load_state_dict(torch.load('/kaggle/working/aed_classifier.pth', map_location=DEVICE)['model'])\naed_classifier.to(DEVICE);","metadata":{"execution":{"iopub.status.busy":"2021-06-29T23:50:26.846789Z","iopub.execute_input":"2021-06-29T23:50:26.847145Z","iopub.status.idle":"2021-06-29T23:50:27.063935Z","shell.execute_reply.started":"2021-06-29T23:50:26.847093Z","shell.execute_reply":"2021-06-29T23:50:27.063086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get probs via NN clasifier\n\nnn_eval_pred_probs = predict_classifier(aed_classifier, eval_loader, return_probs=True, device=DEVICE)\nnn_eval_pred_probs = nn_eval_pred_probs[1]","metadata":{"execution":{"iopub.status.busy":"2021-06-30T00:42:26.679248Z","iopub.status.idle":"2021-06-30T00:42:26.679865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"eval_df['nn_pred'] = nn_eval_pred_probs.argmax(axis=1)\neval_df.head()","metadata":{"execution":{"iopub.status.busy":"2021-06-30T00:42:26.681048Z","iopub.status.idle":"2021-06-30T00:42:26.68169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Eval NN\neval_y = eval_df['label_encoded'].values\nprint(\"F1 macro: %s\" % f1_score(eval_y, nn_eval_pred_probs.argmax(axis=1), average='macro'))","metadata":{"execution":{"iopub.status.busy":"2021-06-30T00:42:26.682864Z","iopub.status.idle":"2021-06-30T00:42:26.683473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make submission.csv","metadata":{}},{"cell_type":"code","source":"nn_test_pred_probs = predict_classifier(aed_classifier, test_loader, return_probs=True, device=DEVICE)\nnn_test_pred_probs = nn_test_pred_probs[1]","metadata":{"execution":{"iopub.status.busy":"2021-06-30T00:42:26.684596Z","iopub.status.idle":"2021-06-30T00:42:26.685194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_submission_df['label'] = nn_test_pred_probs.argmax(axis=1)\nsample_submission_df['label'] = sample_submission_df['label'].apply(lambda x: idx_to_class_dict[x])\nif 'label_encoded' in sample_submission_df.columns:\n    sample_submission_df = sample_submission_df.drop(columns=['label_encoded'])\nsample_submission_df.to_csv('submision.csv', index=None)\nsample_submission_df.head()","metadata":{"execution":{"iopub.status.busy":"2021-06-30T00:42:26.686315Z","iopub.status.idle":"2021-06-30T00:42:26.687075Z"},"trusted":true},"execution_count":null,"outputs":[]}]}