{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":"none","dataSources":[{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"},{"sourceId":8108072,"sourceType":"datasetVersion","datasetId":4789213},{"sourceId":8273787,"sourceType":"datasetVersion","datasetId":4903383,"isSourceIdPinned":true}],"dockerImageVersionId":30684,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Import Packages","metadata":{}},{"cell_type":"code","source":"import re\nimport os\nimport gc\nimport sys\nimport cv2\nimport math\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\nimport librosa\nfrom scipy import signal as sci_signal\n\nimport torch\nfrom torch import nn\nimport timm","metadata":{"execution":{"iopub.status.busy":"2024-05-01T17:36:51.010223Z","iopub.execute_input":"2024-05-01T17:36:51.010637Z","iopub.status.idle":"2024-05-01T17:37:01.652586Z","shell.execute_reply.started":"2024-05-01T17:36:51.010606Z","shell.execute_reply":"2024-05-01T17:37:01.651453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install /kaggle/input/onnxruntime/humanfriendly-10.0-py2.py3-none-any.whl --no-index --find-links /kaggle/input/onnxruntime\n!pip install /kaggle/input/onnxruntime/coloredlogs-15.0.1-py2.py3-none-any.whl --no-index --find-links /kaggle/input/onnxruntime\n!pip install /kaggle/input/onnxruntime/onnxruntime-1.17.3-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl --no-index --find-links /kaggle/input/onnxruntime","metadata":{"execution":{"iopub.status.busy":"2024-05-01T17:37:01.654427Z","iopub.execute_input":"2024-05-01T17:37:01.654762Z","iopub.status.idle":"2024-05-01T17:37:47.101804Z","shell.execute_reply.started":"2024-05-01T17:37:01.654734Z","shell.execute_reply":"2024-05-01T17:37:47.100276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# labels\nlabel_list = sorted(os.listdir(os.path.join('/kaggle/input/birdclef-2024', 'train_audio')))\nlabel_id_list = list(range(len(label_list)))\nlabel2id = dict(zip(label_list, label_id_list))\nid2label = dict(zip(label_id_list, label_list))","metadata":{"execution":{"iopub.status.busy":"2024-05-01T17:37:47.104103Z","iopub.execute_input":"2024-05-01T17:37:47.104523Z","iopub.status.idle":"2024-05-01T17:37:47.121589Z","shell.execute_reply.started":"2024-05-01T17:37:47.104487Z","shell.execute_reply":"2024-05-01T17:37:47.120480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Network","metadata":{}},{"cell_type":"code","source":"from torchaudio.transforms import AmplitudeToDB, MelSpectrogram\nclass NormalizeMelSpec(nn.Module):\n    def __init__(self, eps=1e-6):\n        super().__init__()\n        self.eps = eps\n\n    def forward(self, X):\n        mean = X.mean((1, 2), keepdim=True)\n        std = X.std((1, 2), keepdim=True)\n        Xstd = (X - mean) / (std + self.eps)\n        norm_min, norm_max = \\\n            Xstd.min(-1)[0].min(-1)[0], Xstd.max(-1)[0].max(-1)[0]\n        fix_ind = (norm_max - norm_min) > self.eps * torch.ones_like(\n            (norm_max - norm_min)\n        )\n        V = torch.zeros_like(Xstd)\n        if fix_ind.sum():\n            V_fix = Xstd[fix_ind]\n            norm_max_fix = norm_max[fix_ind, None, None]\n            norm_min_fix = norm_min[fix_ind, None, None]\n            V_fix = torch.max(\n                torch.min(V_fix, norm_max_fix),\n                norm_min_fix,\n            )\n            V_fix = (V_fix - norm_min_fix) / (norm_max_fix - norm_min_fix)\n            V[fix_ind] = V_fix\n        return V\n    \nclass Transform(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.logmelspec_extractor = nn.Sequential(\n                MelSpectrogram(\n                    sample_rate=32000,\n                    n_mels=128,\n                    f_min=20,\n                    f_max=16000,\n                    n_fft=2048,\n                    hop_length=512,\n                    normalized=True,\n                ),\n                AmplitudeToDB(top_db=80.0),\n                NormalizeMelSpec(),\n            )\n\n    def forward(self, x):\n        x = self.logmelspec_extractor(x)[:, None]\n        return x\n    \nclass Net(nn.Module):\n    def __init__(self, backbone):\n        super().__init__()\n        self.model = timm.create_model(backbone, num_classes=182, pretrained=False, in_chans=1)\n\n    def forward(self, x):\n        logits = self.model(x)\n        probs = torch.nn.Sigmoid()(logits)\n        return probs","metadata":{"execution":{"iopub.status.busy":"2024-05-01T18:55:29.694698Z","iopub.execute_input":"2024-05-01T18:55:29.695615Z","iopub.status.idle":"2024-05-01T18:55:29.720037Z","shell.execute_reply.started":"2024-05-01T18:55:29.695576Z","shell.execute_reply":"2024-05-01T18:55:29.719145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference & Submision","metadata":{}},{"cell_type":"code","source":"ONNX_FOLDER = '/kaggle/working/onnx'","metadata":{"execution":{"iopub.status.busy":"2024-05-01T17:37:47.846730Z","iopub.execute_input":"2024-05-01T17:37:47.847036Z","iopub.status.idle":"2024-05-01T17:37:47.850727Z","shell.execute_reply.started":"2024-05-01T17:37:47.847010Z","shell.execute_reply":"2024-05-01T17:37:47.849947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nos.makedirs(ONNX_FOLDER, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-01T17:37:47.852008Z","iopub.execute_input":"2024-05-01T17:37:47.852342Z","iopub.status.idle":"2024-05-01T17:37:47.860975Z","shell.execute_reply.started":"2024-05-01T17:37:47.852314Z","shell.execute_reply":"2024-05-01T17:37:47.859962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ckpt_list = ['/kaggle/input/bird-2024/fold_0_exp_4.pth', '/kaggle/input/bird-2024/fold_1_exp_4.pth', '/kaggle/input/bird-2024/fold_3_exp_4.pth', '/kaggle/input/bird-2024/fold_4_exp_4.pth']\ninput_tensor = torch.randn(3, 1, 128, 313)  # input shape\ninput_names = ['x']\noutput_names = ['output']","metadata":{"execution":{"iopub.status.busy":"2024-05-01T17:37:47.862111Z","iopub.execute_input":"2024-05-01T17:37:47.862469Z","iopub.status.idle":"2024-05-01T17:37:47.884057Z","shell.execute_reply.started":"2024-05-01T17:37:47.862441Z","shell.execute_reply":"2024-05-01T17:37:47.883076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"onnx_paths = []\nfor ckpt in ckpt_list:\n    bird_model = Net('convnext_femto_ols.d1_in1k')\n    pretrained_dict = torch.load(ckpt, map_location=torch.device('cpu'))\n    \n    model_dict = bird_model.state_dict()\n    new_pretrained_dict = {}\n    for k, v in pretrained_dict.items():\n        if k in model_dict and v.size() == model_dict[k].size():\n              new_pretrained_dict[k] = v\n        else:\n            print(\"Don't load layer: \", k)\n    \n    bird_model.load_state_dict(new_pretrained_dict, strict=True)\n    filename = os.path.basename(ckpt).replace(\".pth\", \".onnx\")\n    onnx_path = os.path.join(ONNX_FOLDER, filename)\n    torch.onnx.export(bird_model, input_tensor, onnx_path, verbose=False, input_names=input_names, output_names=output_names, dynamic_axes={'x' : {0 : 'batch_size'},    # variable lenght axes\n                                    'output' : {0 : 'batch_size'}}, opset_version=17)\n    onnx_paths.append(onnx_path)","metadata":{"execution":{"iopub.status.busy":"2024-05-01T17:51:04.781818Z","iopub.execute_input":"2024-05-01T17:51:04.782229Z","iopub.status.idle":"2024-05-01T17:51:07.672428Z","shell.execute_reply.started":"2024-05-01T17:51:04.782195Z","shell.execute_reply":"2024-05-01T17:51:07.671121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Main Loop","metadata":{}},{"cell_type":"code","source":"import time\nimport multiprocessing as mp\nfrom joblib import Parallel, delayed\n\nfrom multiprocessing import Process, Queue, Manager, Lock","metadata":{"execution":{"iopub.status.busy":"2024-05-01T17:51:12.139892Z","iopub.execute_input":"2024-05-01T17:51:12.140596Z","iopub.status.idle":"2024-05-01T17:51:12.214500Z","shell.execute_reply.started":"2024-05-01T17:51:12.140560Z","shell.execute_reply":"2024-05-01T17:51:12.213646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_audio(filepath, sr=32000, normalize=True):\n    audio, orig_sr = librosa.load(filepath, sr=sr)\n#     if sr!=orig_sr:\n#         audio = librosa.resample(y, orig_sr, sr)\n#     audio = audio.astype('float32').ravel()\n#     print(audio.shape)\n    \n    return audio\n","metadata":{"execution":{"iopub.status.busy":"2024-05-01T17:51:13.813755Z","iopub.execute_input":"2024-05-01T17:51:13.814136Z","iopub.status.idle":"2024-05-01T17:51:13.820237Z","shell.execute_reply.started":"2024-05-01T17:51:13.814101Z","shell.execute_reply":"2024-05-01T17:51:13.818887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_frame(wave, chunk_size=32000*5, pad_value=0):\n    num_chunks = -(-len(wave) // chunk_size)  # Ceiling division\n\n    # Split the array into chunks\n    chunks = torch.chunk(wave, num_chunks)\n\n    # Pad the last chunk if necessary\n    last_chunk_size = chunks[-1].shape[0]\n    if last_chunk_size < chunk_size:\n        padding_needed = (chunk_size - last_chunk_size) % chunk_size\n        chunks[-1] = torch.nn.functional.pad(chunks[-1], (0, padding_needed), value=0)\n\n    return chunks","metadata":{"execution":{"iopub.status.busy":"2024-05-01T17:51:17.044789Z","iopub.execute_input":"2024-05-01T17:51:17.045902Z","iopub.status.idle":"2024-05-01T17:51:17.052053Z","shell.execute_reply.started":"2024-05-01T17:51:17.045862Z","shell.execute_reply":"2024-05-01T17:51:17.051140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import onnx\nimport onnxruntime as ort\nimport time\ntorch.set_num_threads(1)\ndef predict_worker(q_result, onnx_paths, test_paths):\n\n    transform_model = Transform()\n    output_names = ['output']\n    input_names = ['x']\n    onnx_models = []\n\n    for onnx_path in onnx_paths:\n        onnx_model = onnx.load(onnx_path)\n        onnx_model_graph = onnx_model.graph\n        sess_opt = ort.SessionOptions()\n        onnx_session = ort.InferenceSession(onnx_model.SerializeToString(), sess_opt)\n        onnx_models.append(onnx_session)\n    t1 = time.time()\n    total_1 = 0\n    total_2 = 0\n    total_3 = 0\n    total_4 = 0\n    for path in test_paths:\n        t5 = time.time()\n        audio = load_audio(path)\n        t6 = time.time()\n#         print(\"HEHE: \", t6 - t5)\n        total_1 += t6 - t5\n        \n        audio = torch.from_numpy(audio).float()\n        t7 = time.time()\n        total_2 += t7 - t6\n        chunks = make_frame(audio)\n\n        chunk_preds = np.zeros(shape=(len(chunks), 182), dtype=np.float32)\n        t8 = time.time()\n        total_3 += t8 - t7\n                \n        for i in range(0, len(chunks), 3):\n#             if len(chunks) - i > 3:\n#                 batch_0 = torch.zeros((3, 160000))\n#             else:\n#                 batch_0 = torch.zeros((len(chunks) - i, 160000))\n#             for j in range(3):\n#                 if i + j < len(chunks):\n#                     batch_0[j] = chunks[i+j]\n            if len(chunks) - i > 3:\n                batch = torch.zeros((3, 1, 128, 313))\n            else:\n                batch = torch.zeros((len(chunks) - i, 1, 128, 313))\n#             batch = torch.zeros((batch_0.shape[0], 1, 128, 313))\n            for j in range(3):\n                if i + j < len(chunks):        \n                    batch[j] = transform_model(chunks[i+j].unsqueeze(0))\n#             for j in range(3):\n#                 if i + j < len(chunks): \n#                     print(torch.sum(batch_0[j]))\n#             print(batch_0.shape)\n#             batch = transform_model(batch_0)\n#             for j in range(batch.shape[0]):\n#                 print(torch.sum(batch[j]))\n            batch = batch.numpy()\n#             print(batch.shape)\n            \n            \n            for onnx_session in onnx_models:\n                pred = onnx_session.run(output_names, {input_names[0]: batch})[0]\n                \n                for j in range(pred.shape[0]):\n                    chunk_preds[i + j] += pred[j]\n            \n            for j in range(pred.shape[0]):    \n                chunk_preds[i + j] /= len(onnx_models)\n        for i in range(0, len(chunks), 3):\n            if i + 1 < len(chunks):\n                vote_pred = (chunk_preds[i] + chunk_preds[i + 1] + chunk_preds[i + 2])/3\n                chunk_preds[i] = vote_pred\n                chunk_preds[i + 1] = vote_pred\n                chunk_preds[i + 2] = vote_pred\n            \n        filename = os.path.basename(path).replace(\".ogg\", '')\n        rec_ids = [f'{filename}_{(frame_id+1)*5}' for frame_id in range(len(chunks))]\n        for i in range(len(chunks)):\n#             print(i, np.max(chunk_preds[i]))\n            q_result.put({\"row_id\": rec_ids[i], \"pred\": chunk_preds[i]}, block=True, timeout=None)\n    t2 = time.time()\n    print(\"Infer: \", t2 - t1)\n    print(\"Preprocess time: \", total_1, total_2, total_3, total_4)\n    for onnx_session in onnx_models:\n        del onnx_session\n    del transform_model\n    q_result.put(None, block=True, timeout=None)\n    print(\"Shutdown process\")\n","metadata":{"execution":{"iopub.status.busy":"2024-05-01T19:11:53.296655Z","iopub.execute_input":"2024-05-01T19:11:53.297037Z","iopub.status.idle":"2024-05-01T19:11:53.321419Z","shell.execute_reply.started":"2024-05-01T19:11:53.297003Z","shell.execute_reply":"2024-05-01T19:11:53.320055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_audio_dir = '/kaggle/input/birdclef-2024/test_soundscapes/'\ntest_paths = [test_audio_dir+f for f in sorted(os.listdir(test_audio_dir))]\nif len(test_paths) == 1:\n    test_audio_dir = '/kaggle/input/birdclef-2024/unlabeled_soundscapes/'\n    test_paths = [test_audio_dir+f for f in sorted(os.listdir(test_audio_dir))][:2]\nnum_test = len(test_paths)\nprint(num_test)","metadata":{"execution":{"iopub.status.busy":"2024-05-01T17:51:35.518153Z","iopub.execute_input":"2024-05-01T17:51:35.518831Z","iopub.status.idle":"2024-05-01T17:51:35.727819Z","shell.execute_reply.started":"2024-05-01T17:51:35.518798Z","shell.execute_reply":"2024-05-01T17:51:35.726998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_pred(q_result_0, q_result_1):\n    predictions = []\n    row_ids = []\n    while True:\n        r = q_result_0.get()\n        if r == None:\n            break\n        row_id = r[\"row_id\"]\n        pred = r[\"pred\"]\n        predictions.append(pred)\n        row_ids.append(row_id)\n    while True:\n        r = q_result_1.get()\n        if r == None:\n            break\n        row_id = r[\"row_id\"]\n        pred = r[\"pred\"]\n        predictions.append(pred)\n        row_ids.append(row_id)    \n\n    sub_pred = pd.DataFrame(predictions, columns=label_list)\n    sub_id = pd.DataFrame({'row_id': row_ids})\n\n    sub = pd.concat([sub_id, sub_pred], axis=1)\n\n    sub.to_csv('submission.csv',index=False)\n    print(f'Submissionn shape: {sub.shape}')\n    sub.head(5)","metadata":{"execution":{"iopub.status.busy":"2024-05-01T17:51:36.764992Z","iopub.execute_input":"2024-05-01T17:51:36.765374Z","iopub.status.idle":"2024-05-01T17:51:36.775256Z","shell.execute_reply.started":"2024-05-01T17:51:36.765345Z","shell.execute_reply":"2024-05-01T17:51:36.774090Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"q_result_0 = Queue(maxsize=4000)\nq_result_1 = Queue(maxsize=4000)\n       \nprocess_0 = Process(target=predict_worker, args=(q_result_0, onnx_paths, test_paths[:num_test//2], ))\nprocess_1 = Process(target=predict_worker, args=(q_result_1, onnx_paths, test_paths[num_test//2:], ))\nsave_process = Process(target=save_pred, args=(q_result_0, q_result_1, ))\n\nsave_process.start()\nprocess_0.start()\nprocess_1.start()\n\nprocess_0.join()\nprocess_1.join()\nsave_process.join()","metadata":{"execution":{"iopub.status.busy":"2024-05-01T19:12:01.758985Z","iopub.execute_input":"2024-05-01T19:12:01.762349Z","iopub.status.idle":"2024-05-01T19:12:12.132205Z","shell.execute_reply.started":"2024-05-01T19:12:01.762294Z","shell.execute_reply":"2024-05-01T19:12:12.130171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -rf /kaggle/working/onnx","metadata":{"execution":{"iopub.status.busy":"2024-05-01T17:37:49.512855Z","iopub.status.idle":"2024-05-01T17:37:49.513370Z","shell.execute_reply.started":"2024-05-01T17:37:49.513115Z","shell.execute_reply":"2024-05-01T17:37:49.513136Z"},"trusted":true},"execution_count":null,"outputs":[]}]}