{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":25954,"databundleVersionId":2091745,"sourceType":"competition"},{"sourceId":7197586,"sourceType":"datasetVersion","datasetId":4162699}],"dockerImageVersionId":30626,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport librosa as lb\n\nimport torch\nfrom torch.utils.data import Dataset\nfrom tqdm.notebook import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-14T05:45:23.914989Z","iopub.execute_input":"2023-12-14T05:45:23.915444Z","iopub.status.idle":"2023-12-14T05:45:23.921779Z","shell.execute_reply.started":"2023-12-14T05:45:23.915413Z","shell.execute_reply":"2023-12-14T05:45:23.920518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FD = 32_000\nduration = 5\nthreshold = 0.33\nnum_labels = 397","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:45:23.923774Z","iopub.execute_input":"2023-12-14T05:45:23.924155Z","iopub.status.idle":"2023-12-14T05:45:23.935020Z","shell.execute_reply.started":"2023-12-14T05:45:23.924124Z","shell.execute_reply":"2023-12-14T05:45:23.933485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:45:23.937328Z","iopub.execute_input":"2023-12-14T05:45:23.937707Z","iopub.status.idle":"2023-12-14T05:45:23.947307Z","shell.execute_reply.started":"2023-12-14T05:45:23.937676Z","shell.execute_reply":"2023-12-14T05:45:23.946116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pathlib import Path\n\ntest_audio_root = Path(\"../input/birdclef-2021/test_soundscapes\")\nsmpl_sub_root = \"../input/birdclef-2021/sample_submission.csv\"\ntarget_path = None\n    \nif not len(list(test_audio_root.glob(\"*.ogg\"))):\n    test_audio_root = Path(\"../input/birdclef-2021/train_soundscapes\")\n    smpl_sub_root = None\n    target_path = Path(\"../input/birdclef-2021/train_soundscape_labels.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:45:23.948846Z","iopub.execute_input":"2023-12-14T05:45:23.949195Z","iopub.status.idle":"2023-12-14T05:45:23.961511Z","shell.execute_reply.started":"2023-12-14T05:45:23.949167Z","shell.execute_reply":"2023-12-14T05:45:23.960041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def colorize(X, eps=1e-6, mean=None, std=None):\n    mean = mean or X.mean()\n    std = std or X.std()\n    X = (X - mean) / (std + eps)\n    \n    _min, _max = X.min(), X.max()\n\n    if (_max - _min) > eps:\n        V = np.clip(X, _min, _max)\n        V = 255 * (V - _min) / (_max - _min)\n        V = V.astype(np.uint8)\n    else:\n        V = np.zeros_like(X, dtype=np.uint8)\n\n    return V","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:45:23.964577Z","iopub.execute_input":"2023-12-14T05:45:23.965965Z","iopub.status.idle":"2023-12-14T05:45:23.975840Z","shell.execute_reply.started":"2023-12-14T05:45:23.965918Z","shell.execute_reply":"2023-12-14T05:45:23.974225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CalcMelSpec:\n    def __init__(self, sr, n_mels, fmin, fmax, **kwargs):\n        self.sr = sr\n        self.n_mels = n_mels\n        self.fmin = fmin\n        self.fmax = fmax\n        kwargs[\"n_fft\"] = kwargs.get(\"n_fft\", self.sr//10)\n        kwargs[\"hop_length\"] = kwargs.get(\"hop_length\", self.sr//(10*4))\n        self.kwargs = kwargs\n\n    def __call__(self, y):\n        melspec = lb.feature.melspectrogram(y=y, sr=self.sr, n_mels=self.n_mels, fmin=self.fmin, fmax=self.fmax, **self.kwargs,)\n        melspec = lb.power_to_db(melspec).astype(np.float32)\n        return melspec","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:45:23.977233Z","iopub.execute_input":"2023-12-14T05:45:23.977613Z","iopub.status.idle":"2023-12-14T05:45:23.991566Z","shell.execute_reply.started":"2023-12-14T05:45:23.977580Z","shell.execute_reply":"2023-12-14T05:45:23.990414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import soundfile as sf\n\nclass BirdCLEFDataset(Dataset):\n    def __init__(self, data, sr=FD, n_mels=128, fmin=0, fmax=None, duration=duration, step=None, res_type=\"kaiser_fast\", resample=True):\n        \n        self.data = data\n        \n        self.sr = sr\n        self.n_mels = n_mels\n        self.fmin = fmin\n        self.fmax = fmax or self.sr//2\n\n        self.duration = duration\n        self.audio_length = self.duration*self.sr\n        self.step = step or self.audio_length\n        \n        self.res_type = res_type\n        self.resample = resample\n\n        self.mel_spec_computer = CalcMelSpec(sr=self.sr, n_mels=self.n_mels, fmin=self.fmin,\n                                                 fmax=self.fmax)\n    def __len__(self):\n        return len(self.data)\n    \n    @staticmethod\n    def normalize(image):\n        image = image.astype(\"float32\", copy=False) / 255.0\n        image = np.stack([image, image, image])\n        return image\n    \n    def audio_to_image(self, audio):\n        melspec = self.mel_spec_computer(audio) \n        image = colorize(melspec)\n        image = self.normalize(image)\n        return image\n\n    def read_file(self, filepath):\n        audio, orig_sr = sf.read(filepath, dtype=\"float32\")\n\n        if self.resample and orig_sr != self.sr:\n            audio = lb.resample(audio, orig_sr, self.sr, res_type=self.res_type)\n          \n        audios = []\n        for i in range(self.audio_length, len(audio) + self.step, self.step):\n            start = max(0, i - self.audio_length)\n            end = start + self.audio_length\n            audios.append(audio[start:end])\n            \n        if len(audios[-1]) < self.audio_length:\n            audios = audios[:-1]\n            \n        images = [self.audio_to_image(audio) for audio in audios]\n        images = np.stack(images)\n        \n        return images\n    \n        \n    def __getitem__(self, idx):\n        return self.read_file(self.data.loc[idx, \"filepath\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:45:23.993351Z","iopub.execute_input":"2023-12-14T05:45:23.993865Z","iopub.status.idle":"2023-12-14T05:45:24.014920Z","shell.execute_reply.started":"2023-12-14T05:45:23.993819Z","shell.execute_reply":"2023-12-14T05:45:24.012632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = pd.DataFrame(\n     [(path.stem, *path.stem.split(\"_\"), path) for path in Path(test_audio_root).glob(\"*.ogg\")],\n    columns = [\"filename\", \"id\", \"site\", \"date\", \"filepath\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:45:24.017782Z","iopub.execute_input":"2023-12-14T05:45:24.018296Z","iopub.status.idle":"2023-12-14T05:45:24.034841Z","shell.execute_reply.started":"2023-12-14T05:45:24.018252Z","shell.execute_reply":"2023-12-14T05:45:24.033711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.read_csv(\"../input/birdclef-2021/train_metadata.csv\")\n\nlabels_id = {label: label_id for label_id,label in enumerate(sorted(df_train[\"primary_label\"].unique()))}\nINV_LABEL_IDS = {val: key for key,val in labels_id.items()}\n\ntest_data = BirdCLEFDataset(data=data)","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:45:24.037488Z","iopub.execute_input":"2023-12-14T05:45:24.038589Z","iopub.status.idle":"2023-12-14T05:45:24.448546Z","shell.execute_reply.started":"2023-12-14T05:45:24.038546Z","shell.execute_reply":"2023-12-14T05:45:24.447360Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision.models import efficientnet_b0\n\ndef load_model(checkpoint_path, num_classes=num_labels):\n    model = efficientnet_b0(pretrained=False)\n    model.classifier[1] = torch.nn.Linear(1280, num_labels)\n    dummy_device = torch.device(\"cpu\")\n    d = torch.load(checkpoint_path, map_location=dummy_device)\n    for key in list(d.keys()):\n        d[key.replace(\"model.\", \"\")] = d.pop(key)\n    model.load_state_dict(d)\n    model = model.to(device)\n    model = model.eval()\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:45:24.451727Z","iopub.execute_input":"2023-12-14T05:45:24.452212Z","iopub.status.idle":"2023-12-14T05:45:24.459789Z","shell.execute_reply.started":"2023-12-14T05:45:24.452177Z","shell.execute_reply":"2023-12-14T05:45:24.458841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint_paths = [\n    Path(\"../input/pretrained/model0.pt\"), Path(\"../input/pretrained/model1.pt\"),\n    Path(\"../input/pretrained/model2.pt\"), Path(\"../input/pretrained/model3.pt\"),\n    Path(\"../input/pretrained/model4.pt\"),\n]\n\nnets = [load_model(checkpoint_path.as_posix()) for checkpoint_path in checkpoint_paths]","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:45:24.460975Z","iopub.execute_input":"2023-12-14T05:45:24.461902Z","iopub.status.idle":"2023-12-14T05:45:25.470079Z","shell.execute_reply.started":"2023-12-14T05:45:24.461867Z","shell.execute_reply":"2023-12-14T05:45:25.469076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef get_thresh_preds(out, thresh=None):\n    thresh = thresh or threshold\n    o = (-out).argsort(1)\n    npreds = (out > thresh).sum(1)\n    preds = []\n    for oo, npred in zip(o, npreds):\n        preds.append(oo[:npred].cpu().numpy().tolist())\n    return preds","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:45:25.472187Z","iopub.execute_input":"2023-12-14T05:45:25.472975Z","iopub.status.idle":"2023-12-14T05:45:25.481109Z","shell.execute_reply.started":"2023-12-14T05:45:25.472928Z","shell.execute_reply":"2023-12-14T05:45:25.480145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_bird_names(preds):\n    bird_names = []\n    for pred in preds:\n        if not pred:\n            bird_names.append(\"nocall\")\n        else:\n            bird_names.append(\" \".join([INV_LABEL_IDS[bird_id] for bird_id in pred]))\n    return bird_names","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:45:25.483317Z","iopub.execute_input":"2023-12-14T05:45:25.483865Z","iopub.status.idle":"2023-12-14T05:45:25.495445Z","shell.execute_reply.started":"2023-12-14T05:45:25.483821Z","shell.execute_reply":"2023-12-14T05:45:25.494064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(nets, test_data, names=True):\n    preds = []\n    with torch.no_grad():\n        for idx in  tqdm(list(range(len(test_data)))):\n            xb = torch.from_numpy(test_data[idx]).to(device)\n            pred = 0.\n            for net in nets:\n                o = net(xb)\n                o = torch.nn.functional.softmax(o, dim=1)\n                pred += o\n\n            pred /= len(nets)\n            \n            if names:\n                pred = get_bird_names(get_thresh_preds(pred))\n\n            preds.append(pred)\n    return preds","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:45:25.497524Z","iopub.execute_input":"2023-12-14T05:45:25.498220Z","iopub.status.idle":"2023-12-14T05:45:25.508103Z","shell.execute_reply.started":"2023-12-14T05:45:25.498175Z","shell.execute_reply":"2023-12-14T05:45:25.506858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_probas = predict(nets, test_data, names=False)\npreds = [get_bird_names(get_thresh_preds(pred, thresh=threshold)) for pred in pred_probas]","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:45:25.509432Z","iopub.execute_input":"2023-12-14T05:45:25.510471Z","iopub.status.idle":"2023-12-14T05:53:11.563060Z","shell.execute_reply.started":"2023-12-14T05:45:25.510433Z","shell.execute_reply":"2023-12-14T05:53:11.561815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preds_as_df(data, preds):\n    sub = {\n        \"row_id\": [],\n        \"birds\": [],\n    }\n    \n    for row, pred in zip(data.itertuples(False), preds):\n        row_id = [f\"{row.id}_{row.site}_{5*i}\" for i in range(1, len(pred)+1)]\n        sub[\"birds\"] += pred\n        sub[\"row_id\"] += row_id\n        \n    sub = pd.DataFrame(sub)\n    \n    if smpl_sub_root:\n        sample_sub = pd.read_csv(smpl_sub_root, usecols=[\"row_id\"])\n        sub = sample_sub.merge(sub, on=\"row_id\", how=\"left\")\n        sub[\"birds\"] = sub[\"birds\"].fillna(\"nocall\")\n    return sub","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:53:11.564809Z","iopub.execute_input":"2023-12-14T05:53:11.566096Z","iopub.status.idle":"2023-12-14T05:53:11.576256Z","shell.execute_reply.started":"2023-12-14T05:53:11.566051Z","shell.execute_reply":"2023-12-14T05:53:11.575000Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = preds_as_df(data, preds)\nsub.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-12-14T05:53:11.578244Z","iopub.execute_input":"2023-12-14T05:53:11.578698Z","iopub.status.idle":"2023-12-14T05:53:11.603947Z","shell.execute_reply.started":"2023-12-14T05:53:11.578659Z","shell.execute_reply":"2023-12-14T05:53:11.602541Z"},"trusted":true},"execution_count":null,"outputs":[]}]}