{"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":"# PyTorch-Simple-Starter-Using only 21 classes","metadata":{}},{"cell_type":"markdown","source":"## Library","metadata":{}},{"cell_type":"code","source":"import os\nimport json\nimport tqdm\nimport random\nimport shutil\nimport pandas as pd\nimport numpy as np\nimport torch\nimport torch.optim as optim\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\nimport torchaudio\nimport torchaudio.transforms as T\nfrom torchvision.models.resnet import ResNet, BasicBlock\nimport seaborn as sns\nimport matplotlib.pyplot as plt","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-05-25T03:26:21.516291Z","iopub.execute_input":"2022-05-25T03:26:21.516591Z","iopub.status.idle":"2022-05-25T03:26:24.268321Z","shell.execute_reply.started":"2022-05-25T03:26:21.516508Z","shell.execute_reply":"2022-05-25T03:26:24.267194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(f\"Using {device} device\")","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:24.275666Z","iopub.execute_input":"2022-05-25T03:26:24.278843Z","iopub.status.idle":"2022-05-25T03:26:24.356648Z","shell.execute_reply.started":"2022-05-25T03:26:24.278797Z","shell.execute_reply":"2022-05-25T03:26:24.355744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Loading","metadata":{}},{"cell_type":"code","source":"root_path = \"../input/birdclef-2022/\"\ninput_path = root_path + '/train_audio/'\nout_path = \"./train/\"\n\ntry:\n    os.mkdir(out_path)\nexcept FileExistsError:\n    pass\n\n\ntrain_meta = pd.read_csv(root_path + 'train_metadata.csv')\n\nwith open(root_path + '/scored_birds.json') as sbfile:\n    scored_birds = json.load(sbfile)","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:24.369899Z","iopub.execute_input":"2022-05-25T03:26:24.370612Z","iopub.status.idle":"2022-05-25T03:26:24.502085Z","shell.execute_reply.started":"2022-05-25T03:26:24.370567Z","shell.execute_reply":"2022-05-25T03:26:24.501367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### only 21 classes","metadata":{}},{"cell_type":"code","source":"train_meta = train_meta[train_meta['primary_label'].isin(scored_birds)]\nbird_label = train_meta[\"primary_label\"].unique()\nprint(bird_label)","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:24.503495Z","iopub.execute_input":"2022-05-25T03:26:24.503969Z","iopub.status.idle":"2022-05-25T03:26:24.528427Z","shell.execute_reply.started":"2022-05-25T03:26:24.503933Z","shell.execute_reply":"2022-05-25T03:26:24.527674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils","metadata":{}},{"cell_type":"code","source":"sample_rate = 32000\nn_fft = 4096\nwin_length = None\nhop_length = 512\nn_mels = 256\nmin_sec_proc = sample_rate*5\nf_min = 1000\nf_max = 16000 \n\nmel_spectrogram = T.MelSpectrogram(\n    sample_rate=sample_rate,\n    n_fft=n_fft,\n    win_length=win_length,\n    hop_length=hop_length,\n    center=True,\n    f_min=f_min,\n    f_max=f_max,\n    pad_mode=\"reflect\",\n    power=2.0,\n    norm='slaney',\n    onesided=True,\n    n_mels=n_mels,\n    mel_scale=\"htk\",\n)","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:24.532843Z","iopub.execute_input":"2022-05-25T03:26:24.533243Z","iopub.status.idle":"2022-05-25T03:26:24.654114Z","shell.execute_reply.started":"2022-05-25T03:26:24.533206Z","shell.execute_reply":"2022-05-25T03:26:24.653375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def torch_fix_seed(seed=42):\n    # Python random\n    random.seed(seed)\n    # Numpy\n    np.random.seed(seed)\n    # Pytorch\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.use_deterministic_algorithms = True\n    \n    \ndef normalize_std(spec):\n    return (spec- torch.mean(spec))/torch.std(spec)\n    \n    \ndef audio_to_mel_label(filepath, min_sec_proc, mode='train', data_index=0, label_list=[], bird_label=[], label_file=[], mel_list=[]):\n    if mode == 'train':\n        label_file_all = np.zeros(bird_label.shape)\n        for label_file_temp in label_file:\n            label_file_all += (label_file_temp == bird_label)\n        label_file_all = np.clip(label_file_all, 0, 1)\n    \n    waveform, sample_rate_file = torchaudio.load(filepath=filepath)\n    len_wav = waveform.shape[1]\n    waveform = waveform[0,:].reshape(1, len_wav) # stereo->mono mono->mono\n    if mode == 'train':\n        if len_wav < min_sec_proc:\n            for _ in range(round(min_sec_proc/len_wav)):\n                waveform = torch.cat((waveform,waveform[:,0:len_wav]),1)\n            len_wav = min_sec_proc\n            waveform = waveform[:,0:len_wav]\n    elif mode == 'test':\n        if not len_wav < min_sec_proc*12:\n            waveform = torch.cat((waveform,waveform[:,0:len_wav]),1)\n            len_wav = min_sec_proc*12\n            waveform = waveform[:,0:len_wav]\n\n    for index in range(int(len_wav/min_sec_proc)):\n        log_melspec = torch.log10(mel_spectrogram(waveform[0, index*min_sec_proc:index*min_sec_proc+min_sec_proc]).unsqueeze(0)+1e-10)\n        log_melspec = normalize_std(log_melspec)\n        if mode == 'train':\n            torch.save(log_melspec, out_path + str(data_index) + '.pt')\n            label_list.append(label_file_all)\n            data_index += 1\n        elif mode == 'test':\n            mel_list.append(log_melspec)\n            \n    if mode == 'train':\n        return data_index\n    elif mode == 'test':\n        return mel_list\n    \n\ndef load_tensor(path, file_name):\n    return torch.load(path + str(file_name) + '.pt')\n\n\ndef get_X_y(path, idx, label_list):\n    batch_X = torch.stack([load_tensor(path, x.item()) for x in idx])\n    batch_y = torch.stack([label_list[x.item()] for x in idx])\n    return batch_X, batch_y\n\n        \ndef plor_history(history):\n    plt.figure(figsize=(10, 10)) \n    plt.plot(history[:,0], history[:,1], label='loss')\n    plt.plot(history[:,0], history[:,2], label='val_loss')\n    plt.xlabel('epoch')\n    plt.ylabel('loss')\n    plt.legend()\n    \ndef show_heatmap(df):\n    # make heatmap\n    heatmap = []\n\n    for idx, row in df.iterrows():\n        _, _, species, sec = row.row_id.split(\"_\")\n        if sec == \"5\":\n            sec = \"05\"\n        true_or_false = row.target\n        heatmap.append([species,sec,true_or_false])\n    \n    heatmap = pd.DataFrame(heatmap, columns=[\"species\", \"sec\", \"True_or_False\"])\n\n    # show heamap\n    fig,ax = plt.subplots(figsize=(10,5))\n    cmap = sns.color_palette(\"Blues\")\n    heatmap = heatmap.pivot(\"species\", \"sec\", \"True_or_False\")\n    sns.heatmap(heatmap,ax=ax,linecolor='k',lw=1,cmap=cmap)\n    plt.title(\"Prediction result in soundscape_453028782.ogg\")\n    plt.show()    \n    \ntorch_fix_seed()","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:24.655327Z","iopub.execute_input":"2022-05-25T03:26:24.655583Z","iopub.status.idle":"2022-05-25T03:26:24.678696Z","shell.execute_reply.started":"2022-05-25T03:26:24.65555Z","shell.execute_reply":"2022-05-25T03:26:24.677992Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Save melspectrogram","metadata":{}},{"cell_type":"code","source":"data_index = 0\nlabel_list = []\nfor pri_label, secon_label, f_name in zip(tqdm.notebook.tqdm(train_meta['primary_label']), train_meta['secondary_labels'],train_meta['filename']):\n    data_index = audio_to_mel_label(input_path+f_name, min_sec_proc,'train', data_index, label_list, bird_label, [pri_label] + eval(secon_label))\n\ntorch.save(np.stack(label_list), out_path + 'label_list.pt')\nlabel_list = torch.from_numpy(np.stack(label_list)).clone()","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:24.680223Z","iopub.execute_input":"2022-05-25T03:26:24.680542Z","iopub.status.idle":"2022-05-25T03:26:37.236393Z","shell.execute_reply.started":"2022-05-25T03:26:24.680505Z","shell.execute_reply":"2022-05-25T03:26:37.234808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"n_output = len(bird_label)\n\nout_sigmoid = nn.Sigmoid()\n\nclass ResNetBird(ResNet):\n    def __init__(self):\n        super().__init__(BasicBlock, [1, 4, 6, 3], num_classes=n_output)\n\n        self.conv1 = nn.Conv2d(1, 64, kernel_size=7, stride=1, padding=3, bias=False)\n\n        \nnet = ResNetBird().to(device)","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:37.237555Z","iopub.status.idle":"2022-05-25T03:26:37.238142Z","shell.execute_reply.started":"2022-05-25T03:26:37.237879Z","shell.execute_reply":"2022-05-25T03:26:37.237905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms = torch.nn.Sequential(\n    torchaudio.transforms.FrequencyMasking(freq_mask_param=15),\n    torchaudio.transforms.TimeMasking(time_mask_param=20),\n)\nComputeDeltas = torchaudio.transforms.ComputeDeltas(win_length= 5)\n\ndef time_shift(melspec):\n    for i in range(melspec.shape[0]):\n        ii = int(np.random.randint(melspec.shape[3]/4, melspec.shape[3]/(4/3), (1)))\n        melspec[i,:,:,:] = torch.cat((melspec[i,:,:,ii:-1], melspec[i,:,:,0:ii+1]),2)\n    \n    return melspec","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:37.239303Z","iopub.status.idle":"2022-05-25T03:26:37.239845Z","shell.execute_reply.started":"2022-05-25T03:26:37.239618Z","shell.execute_reply":"2022-05-25T03:26:37.239643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data split","metadata":{}},{"cell_type":"code","source":"train_idx = np.arange(0, label_list.shape[0])\n\ndata_len = train_idx.shape[0]\ntrain_idx, val_idx = torch.utils.data.random_split(train_idx, [int(data_len*0.8), data_len-int(data_len*0.8)])","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:37.240876Z","iopub.status.idle":"2022-05-25T03:26:37.24142Z","shell.execute_reply.started":"2022-05-25T03:26:37.241188Z","shell.execute_reply":"2022-05-25T03:26:37.241213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train loop","metadata":{}},{"cell_type":"code","source":"num_epochs = 50\nlr = 0.001\nbatch_size = 32\n\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = optim.Adam(net.parameters(), lr=lr)\nhistory = np.zeros((0, 3))\n\ntrain_loader = DataLoader(train_idx, batch_size=batch_size, shuffle=True)\nval_loader = DataLoader(val_idx, batch_size=batch_size, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:37.242458Z","iopub.status.idle":"2022-05-25T03:26:37.24299Z","shell.execute_reply.started":"2022-05-25T03:26:37.242764Z","shell.execute_reply":"2022-05-25T03:26:37.242789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(num_epochs):\n    train_loss, val_loss = 0, 0\n    n_train, n_val = 0, 0\n\n    net.train()\n    for idx in train_loader:\n        inputs, labels = get_X_y(out_path, idx, label_list)\n        \n        n_train += len(labels)\n        \n        inputs = inputs.to(device)\n        labels = labels.to(device)\n\n        optimizer.zero_grad()\n        outputs = net(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        train_loss += loss.item()\n\n    net.eval()\n    with torch.no_grad():\n        for idx in val_loader:\n            inputs_val, labels_val = get_X_y(out_path, idx, label_list)\n            n_val += len(labels)\n\n            inputs_val = inputs_val.to(device)\n            labels_val = labels_val.to(device)\n\n            outputs_val = net(inputs_val)\n\n            loss_val = criterion(outputs_val, labels_val)\n\n            val_loss += loss_val.item()\n    \n\n    train_loss = train_loss * batch_size / n_train\n    val_loss = val_loss * batch_size / n_val\n    print (f'Epoch [{(epoch+1)}/{num_epochs}], loss: {train_loss:.5f}, val_loss: {val_loss:.5f}')\n    item = np.array([epoch+1, train_loss, val_loss])\n    history = np.vstack((history, item))\n\ntorch.save(net.state_dict(), 'model.pt')\nplor_history(history)","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:37.244061Z","iopub.status.idle":"2022-05-25T03:26:37.244599Z","shell.execute_reply.started":"2022-05-25T03:26:37.244368Z","shell.execute_reply":"2022-05-25T03:26:37.244401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"test_audio_dir = '../input/birdclef-2022/test_soundscapes/'\nfile_list = [f.split('.')[0] for f in sorted(os.listdir(test_audio_dir))]\n\nprint('Number of test soundscapes:', len(file_list))","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:37.245639Z","iopub.status.idle":"2022-05-25T03:26:37.246175Z","shell.execute_reply.started":"2022-05-25T03:26:37.245939Z","shell.execute_reply":"2022-05-25T03:26:37.245964Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred = {'row_id': [], 'target': []}\nbinary_th = 0.30\nnet.eval()\n\nfor afile in file_list:\n    \n    path = test_audio_dir + afile + '.ogg'\n    \n    chunks = [[] for i in range(12)]\n    \n    mel_list_test = []\n    mel_list_test = audio_to_mel_label(path, min_sec_proc,'test', mel_list=mel_list_test)\n    mel_list_test = torch.stack(mel_list_test).to(device)\n    \n    outputs = net(mel_list_test)\n    \n    outputs_test = out_sigmoid(outputs)\n    \n    for idx, i in enumerate(range(len(chunks))):        \n        chunk_end_time = (i + 1) * 5\n        for bird in scored_birds:\n            \n            try:\n                score = outputs_test[idx][np.where(bird_label==bird)]\n            except IndexError:\n                score = 0\n            \n            row_id = afile + '_' + bird + '_' + str(chunk_end_time)\n            \n            pred['row_id'].append(row_id)\n            pred['target'].append(True if score > binary_th else False)","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:37.247215Z","iopub.status.idle":"2022-05-25T03:26:37.247749Z","shell.execute_reply.started":"2022-05-25T03:26:37.247526Z","shell.execute_reply":"2022-05-25T03:26:37.24755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"results = pd.DataFrame(pred, columns = ['row_id', 'target'])\n\nprint(results['target']) \n    \nresults.to_csv(\"submission.csv\", index=False)    ","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:37.248788Z","iopub.status.idle":"2022-05-25T03:26:37.249323Z","shell.execute_reply.started":"2022-05-25T03:26:37.249104Z","shell.execute_reply":"2022-05-25T03:26:37.249128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_heatmap(results)","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:37.250354Z","iopub.status.idle":"2022-05-25T03:26:37.250884Z","shell.execute_reply.started":"2022-05-25T03:26:37.250663Z","shell.execute_reply":"2022-05-25T03:26:37.250686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"OUTPUT_DATA_DELETE = True\n\nif OUTPUT_DATA_DELETE == True:\n    shutil.rmtree(out_path)\n    os.remove('model.pt')","metadata":{"execution":{"iopub.status.busy":"2022-05-25T03:26:37.251917Z","iopub.status.idle":"2022-05-25T03:26:37.252471Z","shell.execute_reply.started":"2022-05-25T03:26:37.252238Z","shell.execute_reply":"2022-05-25T03:26:37.252262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## References\n[BirdCLEF2022:PyTorch-ResNet34-Starter](https://www.kaggle.com/code/myso1987/birdclef2022-pytorch-resnet34-starter-lb-0-50)\n\n[How to submit to BirdCLEF 2022](https://www.kaggle.com/stefankahl/how-to-submit-to-birdclef-2022)\n\n[[Birdclef2022] Soundscape Visualisations](https://www.kaggle.com/code/shinmurashinmura/birdclef2022-soundscape-visualisations)\n","metadata":{}}]}