{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import numpy as np\nimport os\nimport pandas as pd\nimport librosa\nimport librosa.display\nimport matplotlib.pyplot as plt\nimport random\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nfrom tqdm.auto import tqdm\nfrom pathlib import Path","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/birdsong-recognition\"\nINPUT_DIR = \"/kaggle/input/\"\nSR = 32000\nBATCH_SIZE = 16\nDEVICE = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"BIRD_CODE = {\n    'aldfly': 0, 'ameavo': 1, 'amebit': 2, 'amecro': 3, 'amegfi': 4,\n    'amekes': 5, 'amepip': 6, 'amered': 7, 'amerob': 8, 'amewig': 9,\n    'amewoo': 10, 'amtspa': 11, 'annhum': 12, 'astfly': 13, 'baisan': 14,\n    'baleag': 15, 'balori': 16, 'banswa': 17, 'barswa': 18, 'bawwar': 19,\n    'belkin1': 20, 'belspa2': 21, 'bewwre': 22, 'bkbcuc': 23, 'bkbmag1': 24,\n    'bkbwar': 25, 'bkcchi': 26, 'bkchum': 27, 'bkhgro': 28, 'bkpwar': 29,\n    'bktspa': 30, 'blkpho': 31, 'blugrb1': 32, 'blujay': 33, 'bnhcow': 34,\n    'boboli': 35, 'bongul': 36, 'brdowl': 37, 'brebla': 38, 'brespa': 39,\n    'brncre': 40, 'brnthr': 41, 'brthum': 42, 'brwhaw': 43, 'btbwar': 44,\n    'btnwar': 45, 'btywar': 46, 'buffle': 47, 'buggna': 48, 'buhvir': 49,\n    'bulori': 50, 'bushti': 51, 'buwtea': 52, 'buwwar': 53, 'cacwre': 54,\n    'calgul': 55, 'calqua': 56, 'camwar': 57, 'cangoo': 58, 'canwar': 59,\n    'canwre': 60, 'carwre': 61, 'casfin': 62, 'caster1': 63, 'casvir': 64,\n    'cedwax': 65, 'chispa': 66, 'chiswi': 67, 'chswar': 68, 'chukar': 69,\n    'clanut': 70, 'cliswa': 71, 'comgol': 72, 'comgra': 73, 'comloo': 74,\n    'commer': 75, 'comnig': 76, 'comrav': 77, 'comred': 78, 'comter': 79,\n    'comyel': 80, 'coohaw': 81, 'coshum': 82, 'cowscj1': 83, 'daejun': 84,\n    'doccor': 85, 'dowwoo': 86, 'dusfly': 87, 'eargre': 88, 'easblu': 89,\n    'easkin': 90, 'easmea': 91, 'easpho': 92, 'eastow': 93, 'eawpew': 94,\n    'eucdov': 95, 'eursta': 96, 'evegro': 97, 'fiespa': 98, 'fiscro': 99,\n    'foxspa': 100, 'gadwal': 101, 'gcrfin': 102, 'gnttow': 103, 'gnwtea': 104,\n    'gockin': 105, 'gocspa': 106, 'goleag': 107, 'grbher3': 108, 'grcfly': 109,\n    'greegr': 110, 'greroa': 111, 'greyel': 112, 'grhowl': 113, 'grnher': 114,\n    'grtgra': 115, 'grycat': 116, 'gryfly': 117, 'haiwoo': 118, 'hamfly': 119,\n    'hergul': 120, 'herthr': 121, 'hoomer': 122, 'hoowar': 123, 'horgre': 124,\n    'horlar': 125, 'houfin': 126, 'houspa': 127, 'houwre': 128, 'indbun': 129,\n    'juntit1': 130, 'killde': 131, 'labwoo': 132, 'larspa': 133, 'lazbun': 134,\n    'leabit': 135, 'leafly': 136, 'leasan': 137, 'lecthr': 138, 'lesgol': 139,\n    'lesnig': 140, 'lesyel': 141, 'lewwoo': 142, 'linspa': 143, 'lobcur': 144,\n    'lobdow': 145, 'logshr': 146, 'lotduc': 147, 'louwat': 148, 'macwar': 149,\n    'magwar': 150, 'mallar3': 151, 'marwre': 152, 'merlin': 153, 'moublu': 154,\n    'mouchi': 155, 'moudov': 156, 'norcar': 157, 'norfli': 158, 'norhar2': 159,\n    'normoc': 160, 'norpar': 161, 'norpin': 162, 'norsho': 163, 'norwat': 164,\n    'nrwswa': 165, 'nutwoo': 166, 'olsfly': 167, 'orcwar': 168, 'osprey': 169,\n    'ovenbi1': 170, 'palwar': 171, 'pasfly': 172, 'pecsan': 173, 'perfal': 174,\n    'phaino': 175, 'pibgre': 176, 'pilwoo': 177, 'pingro': 178, 'pinjay': 179,\n    'pinsis': 180, 'pinwar': 181, 'plsvir': 182, 'prawar': 183, 'purfin': 184,\n    'pygnut': 185, 'rebmer': 186, 'rebnut': 187, 'rebsap': 188, 'rebwoo': 189,\n    'redcro': 190, 'redhea': 191, 'reevir1': 192, 'renpha': 193, 'reshaw': 194,\n    'rethaw': 195, 'rewbla': 196, 'ribgul': 197, 'rinduc': 198, 'robgro': 199,\n    'rocpig': 200, 'rocwre': 201, 'rthhum': 202, 'ruckin': 203, 'rudduc': 204,\n    'rufgro': 205, 'rufhum': 206, 'rusbla': 207, 'sagspa1': 208, 'sagthr': 209,\n    'savspa': 210, 'saypho': 211, 'scatan': 212, 'scoori': 213, 'semplo': 214,\n    'semsan': 215, 'sheowl': 216, 'shshaw': 217, 'snobun': 218, 'snogoo': 219,\n    'solsan': 220, 'sonspa': 221, 'sora': 222, 'sposan': 223, 'spotow': 224,\n    'stejay': 225, 'swahaw': 226, 'swaspa': 227, 'swathr': 228, 'treswa': 229,\n    'truswa': 230, 'tuftit': 231, 'tunswa': 232, 'veery': 233, 'vesspa': 234,\n    'vigswa': 235, 'warvir': 236, 'wesblu': 237, 'wesgre': 238, 'weskin': 239,\n    'wesmea': 240, 'wessan': 241, 'westan': 242, 'wewpew': 243, 'whbnut': 244,\n    'whcspa': 245, 'whfibi': 246, 'whtspa': 247, 'whtswi': 248, 'wilfly': 249,\n    'wilsni1': 250, 'wiltur': 251, 'winwre3': 252, 'wlswar': 253, 'wooduc': 254,\n    'wooscj2': 255, 'woothr': 256, 'y00475': 257, 'yebfly': 258, 'yebsap': 259,\n    'yehbla': 260, 'yelwar': 261, 'yerwar': 262, 'yetvir': 263\n}\n\nINV_BIRD_CODE = {v: k for k, v in BIRD_CODE.items()}","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class SimpleCNNModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_channels=1, out_channels=8, kernel_size=5)\n        self.maxpool1 = nn.MaxPool2d(2)\n        self.conv2 = nn.Conv2d(in_channels=8, out_channels=16, kernel_size=5)\n        self.maxpool2 = nn.MaxPool2d(2)\n        self.conv3 = nn.Conv2d(in_channels=16, out_channels=32, kernel_size=3)\n        self.maxpool3 = nn.MaxPool2d(2)\n        self.conv4 = nn.Conv2d(in_channels=32, out_channels=64, kernel_size=3)\n        self.maxpool4 = nn.MaxPool2d(2)\n\n        self.fc1 = nn.Linear(in_features=64*5*17, out_features=400)\n        self.out = nn.Linear(in_features=400, out_features=264)\n\n    def forward(self, t):\n        out = self.maxpool1(self.conv1(t))\n        out = F.relu(out)\n\n        out = self.maxpool2(self.conv2(out))\n        out = F.relu(out)\n        \n        out = self.maxpool3(self.conv3(out))\n        out = F.relu(out)\n        \n        out = self.maxpool4(self.conv4(out))\n        out = F.relu(out)\n\n        out = out.reshape(-1, 64*5*17)\n        out = self.fc1(out)\n        return self.out(out)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = SimpleCNNModel()\ncheckpoint = torch.load(\"/kaggle/input/simplecnnmodel/best_model.pt\")\nmodel.load_state_dict(checkpoint['model_state_dict'])\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=4e-5)\nmodel.to(DEVICE)\n# model.load_state_dict(torch.load(\"/kaggle/input/birdcallcnnck/checkpoint.pt\"))\n# model.to(DEVICE)\nmodel.eval()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"TEST = Path(\"/kaggle/input/birdsong-recognition/test_audio\").exists()\n\nif TEST:\n    TEST_DF_DIR = \"/kaggle/input/birdsong-recognition/\"\nelse:\n    TEST_DF_DIR = \"/kaggle/input/birdcall-check/\"\n    \ntest_df = pd.read_csv(f\"{TEST_DF_DIR}test.csv\")\ntest_audio = TEST_DF_DIR + \"test_audio/\"\ntest_df.sample(5)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_melspectrogram(y):\n    mel_spec = librosa.feature.melspectrogram(y=y, sr=SR)\n    mel_spec = librosa.power_to_db(mel_spec, ref=np.max)\n    return mel_spec","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class BirdSoundDatasetTest(Dataset):\n    def __init__(self, df, clip):\n        self.df = df\n        self.clip = clip\n    \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        sample = self.df.iloc[idx]\n        site = sample.site\n        row_id = sample.row_id\n        \n        if site == \"site_3\":\n            y = self.clip.astype(np.float32)\n            len_y = len(y)\n            start = 0\n            end = SR * 5\n            images = []\n            while len_y > start:\n                y_batch = y[start:end].astype(np.float32)\n                if len(y_batch) != (SR * 5):\n                    break\n                start = end\n                end = end + SR * 5\n                \n                image = get_melspectrogram(y_batch)\n                image = np.resize(image, (128, 313))\n                images.append(image)\n            return images, row_id, site\n        else:\n            end_seconds = int(sample.seconds)\n            start_seconds = int(end_seconds - 5)\n            \n            start_index = SR * start_seconds\n            end_index = SR * end_seconds\n            \n            y = self.clip[start_index:end_index].astype(np.float32)\n\n            image = get_melspectrogram(y)\n            image = np.resize(image, (128, 313))\n\n            return image, row_id, site","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"prediction_dict = {}\nprediction_threshold = 0.6\n\nfor audio_id in test_df.audio_id.unique():\n    y, sr = librosa.load(test_audio + (audio_id + \".mp3\"),\n                        sr=SR)\n    new_test_df = test_df.query(\n        f\"audio_id == '{audio_id}'\"\n    ).reset_index(drop=True)\n\n    dataset = BirdSoundDatasetTest(new_test_df, y)\n    loader = DataLoader(dataset, batch_size=1, shuffle=False)\n\n    model.eval()\n\n    for image, row_id, site in loader:\n        site = site[0]\n        row_id = row_id[0]\n        if site in {\"site_1\", \"site_2\"}:\n            image = image.to(DEVICE).unsqueeze(dim=0)\n\n            with torch.no_grad():\n                prediction = model(image)\n                \n                prediction = F.softmax(prediction, dim=1).cpu().numpy()\n                \n                events = prediction >= prediction_threshold\n                labels = np.argwhere(events).tolist()\n                \n                if len(labels) == 0:\n                    prediction_dict[row_id] = \"nocall\"\n                else:\n                    labels = [x[1] for x in labels]\n\n                    birds = set()\n                    for bird in labels:\n                        birds.add(INV_BIRD_CODE[bird])\n                    prediction_dict[row_id] = \" \".join(birds)\n        else:\n            birds = set()\n            for img in image:\n                img = img.to(DEVICE).unsqueeze(dim=0)\n                with torch.no_grad():\n                    prediction = model(img)\n                    prediction = F.softmax(prediction, dim=1).cpu().numpy()\n                    \n                    events = prediction >= prediction_threshold\n                    labels = np.argwhere(events).tolist()\n                    \n                    if len(labels) == 0:\n                        pass\n                    else:\n                        labels = [x[1] for x in labels]\n                        \n                        for bird in labels:\n                            birds.add(INV_BIRD_CODE[bird])\n            if len(birds) == 0:\n                prediction_dict[row_id] = \"nocall\"\n            else:\n                prediction_dict[row_id] = \" \".join(birds)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"row_id = list(prediction_dict.keys())\nbirds = list(prediction_dict.values())\n\nprediction_df = pd.DataFrame({\n    \"row_id\": row_id,\n    \"birds\": birds\n})","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"prediction_df.tail(25)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"prediction_df.to_csv(\"submission.csv\", index=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}