{"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":"code","source":"import json\nimport os\nimport random\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torchaudio\nimport torchaudio.transforms as T\nfrom torchvision.models.resnet import ResNet, BasicBlock","metadata":{"execution":{"iopub.status.busy":"2022-04-20T12:40:55.66529Z","iopub.execute_input":"2022-04-20T12:40:55.665743Z","iopub.status.idle":"2022-04-20T12:40:57.30597Z","shell.execute_reply.started":"2022-04-20T12:40:55.665667Z","shell.execute_reply":"2022-04-20T12:40:57.305129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check device\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(f\"Using {device} device\")\n\n# Prepare paths\nroot_path = \"../input/birdclef-2022/\"\ninput_path = root_path + '/train_audio/'\n\n# Read dataset and labels\ntrain_meta = pd.read_csv(root_path + 'train_metadata.csv')\nwith open(root_path + '/scored_birds.json') as sbfile:\n    scored_birds = json.load(sbfile)\n\n# bird_label = train_meta[\"primary_label\"].unique()\nbird_label = np.asarray(scored_birds)\n\n# Preprocessing data\nsample_rate = 48000\nn_fft = 2048\nwin_length = None\nhop_length = 1024\nn_mels = 128\nmin_sec_proc = sample_rate*5\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    pad_mode=\"reflect\",\n    power=1.0,\n    norm='slaney',\n    onesided=True,\n    n_mels=n_mels,\n    mel_scale=\"slaney\",\n)","metadata":{"execution":{"iopub.status.busy":"2022-04-20T12:40:57.30812Z","iopub.execute_input":"2022-04-20T12:40:57.308397Z","iopub.status.idle":"2022-04-20T12:40:57.567809Z","shell.execute_reply.started":"2022-04-20T12:40:57.308342Z","shell.execute_reply":"2022-04-20T12:40:57.566708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Set pseudo randomize\ndef 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\ntorch_fix_seed()","metadata":{"execution":{"iopub.status.busy":"2022-04-20T12:40:57.569189Z","iopub.execute_input":"2022-04-20T12:40:57.569457Z","iopub.status.idle":"2022-04-20T12:40:57.57727Z","shell.execute_reply.started":"2022-04-20T12:40:57.569418Z","shell.execute_reply":"2022-04-20T12:40:57.576507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create spectrogramm for audio files\ndef audio_to_mel_label(filepath,\n                       min_sec_proc,\n                       mode='train',\n                       data_index=0,\n                       label_list=[],\n                       bird_label=[],\n                       label_file=[],\n                       mel_list=[]):\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 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(\n            mel_spectrogram(waveform[0, index * min_sec_proc:index * min_sec_proc + min_sec_proc]).reshape(1, 128,\n\n  235) + 1e-10)\n        log_melspec = (log_melspec - torch.mean(log_melspec)) / torch.std(log_melspec)\n\n        mel_list.append(log_melspec)\n\n    return mel_list\n\n# class ResNetBird(ResNet):\n#     def __init__(self):\n#         super().__init__(BasicBlock, [5, 8, 6, 3], num_classes=21)\n\n#         self.conv1 = nn.Conv2d(1, 64, kernel_size=7, stride=1, padding=3, bias=False)\n\n\n# net = ResNetBird().to(device)\n\nimport torchvision\nfrom torchvision import datasets, models, transforms\n\nnet = models.resnet18(pretrained=False)\nnet.conv1 = nn.Conv2d(1, 64, kernel_size=7, stride=1, padding=3, bias=False)\nnet.fc = nn.Linear(net.fc.in_features, 21)\nnet.to(device)\n\nnet.load_state_dict(torch.load('../input/birdsclassification/model.pt'))\nout_sigmoid = nn.Sigmoid()\n\ntest_audio_dir = root_path + '/test_soundscapes/'\nfile_list = [f.split('.')[0] for f in sorted(os.listdir(test_audio_dir))]\n\npred = {'row_id': [], 'target': []}\nbinary_th = 0.001\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            print(score)\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)\n\nresults = pd.DataFrame(pred, columns=['row_id', 'target'])\n\nprint(results)\n\nresults.to_csv(\"./submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-04-20T12:40:57.57919Z","iopub.execute_input":"2022-04-20T12:40:57.579639Z","iopub.status.idle":"2022-04-20T12:41:07.155989Z","shell.execute_reply.started":"2022-04-20T12:40:57.579584Z","shell.execute_reply":"2022-04-20T12:41:07.155129Z"},"trusted":true},"execution_count":null,"outputs":[]}]}