{"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":"!nvidia-smi","metadata":{"id":"ikMF-nQSdTbF","outputId":"041c6767-86ab-46e4-9b22-a4840121ecca","papermill":{"duration":1.301859,"end_time":"2021-02-16T15:59:41.363841","exception":false,"start_time":"2021-02-16T15:59:40.061982","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-24T23:16:27.328247Z","iopub.execute_input":"2022-11-24T23:16:27.328625Z","iopub.status.idle":"2022-11-24T23:16:28.461712Z","shell.execute_reply.started":"2022-11-24T23:16:27.328589Z","shell.execute_reply":"2022-11-24T23:16:28.460290Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    import resnest\nexcept ModuleNotFoundError:\n    !pip install -q \"../input/resnest50-fast-package/resnest-0.0.6b20200701/resnest\"","metadata":{"execution":{"iopub.status.busy":"2022-11-24T23:16:28.463837Z","iopub.execute_input":"2022-11-24T23:16:28.464541Z","iopub.status.idle":"2022-11-24T23:16:28.473121Z","shell.execute_reply.started":"2022-11-24T23:16:28.464478Z","shell.execute_reply":"2022-11-24T23:16:28.471955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport librosa as lb\nimport soundfile as sf\nimport pandas as pd\nimport cv2\nfrom pathlib import Path\nimport re\nimport librosa as lb\nimport librosa.display as lbd\nimport torch\nfrom torch import nn\nfrom  torch.utils.data import Dataset, DataLoader\nimport torchvision.models as models\nfrom tqdm.notebook import tqdm\nimport sys\nsys.path.append('../input/timm-pytorch-image-models/pytorch-image-models-master')\nsys.path.append('../input/vision060/vision-0.6.0')\nimport timm\nimport time\nfrom resnest.torch import resnest50, resnest101, resnest50_fast_1s1x64d","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","id":"kSCcqxf0c_O7","papermill":{"duration":4.153186,"end_time":"2021-02-16T15:59:57.894303","exception":false,"start_time":"2021-02-16T15:59:53.741117","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-24T23:16:29.457581Z","iopub.execute_input":"2022-11-24T23:16:29.457932Z","iopub.status.idle":"2022-11-24T23:16:29.464513Z","shell.execute_reply.started":"2022-11-24T23:16:29.457898Z","shell.execute_reply":"2022-11-24T23:16:29.463348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configs","metadata":{"papermill":{"duration":0.046996,"end_time":"2021-02-16T15:59:59.377834","exception":false,"start_time":"2021-02-16T15:59:59.330838","status":"completed"},"tags":[]}},{"cell_type":"code","source":"NUM_CLASSES = 397\nSR = 32_000\nDURATION = 5\nTHRESH = 0.25\n\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"DEVICE:\", DEVICE)\n\nTEST_AUDIO_ROOT = Path(\"../input/birdclef-2021/test_soundscapes\")\nSAMPLE_SUB_PATH = \"../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    SAMPLE_SUB_PATH = None\n    # SAMPLE_SUB_PATH = \"../input/birdclef-2021/sample_submission.csv\"\n    TARGET_PATH = Path(\"../input/birdclef-2021/train_soundscape_labels.csv\")","metadata":{"id":"rJhYZVIDc_O9","papermill":{"duration":0.440614,"end_time":"2021-02-16T15:59:59.86408","exception":false,"start_time":"2021-02-16T15:59:59.423466","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-24T23:16:30.561152Z","iopub.execute_input":"2022-11-24T23:16:30.561533Z","iopub.status.idle":"2022-11-24T23:16:30.570710Z","shell.execute_reply.started":"2022-11-24T23:16:30.561495Z","shell.execute_reply":"2022-11-24T23:16:30.569634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data","metadata":{"papermill":{"duration":0.046026,"end_time":"2021-02-16T16:00:00.054107","exception":false,"start_time":"2021-02-16T16:00:00.008081","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class MelSpecComputer:\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\n        melspec = lb.feature.melspectrogram(\n            y, sr=self.sr, n_mels=self.n_mels, fmin=self.fmin, fmax=self.fmax, **self.kwargs,\n        )\n\n        melspec = lb.power_to_db(melspec).astype(np.float32)\n        return melspec","metadata":{"id":"4fwNhdbJc_O-","papermill":{"duration":0.056975,"end_time":"2021-02-16T16:00:00.157168","exception":false,"start_time":"2021-02-16T16:00:00.100193","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-24T23:16:32.506182Z","iopub.execute_input":"2022-11-24T23:16:32.506537Z","iopub.status.idle":"2022-11-24T23:16:32.513029Z","shell.execute_reply.started":"2022-11-24T23:16:32.506503Z","shell.execute_reply":"2022-11-24T23:16:32.512096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def mono_to_color(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\n\ndef crop_or_pad(y, length):\n    if len(y) < length:\n        y = np.concatenate([y, length - np.zeros(len(y))])\n    elif len(y) > length:\n        y = y[:length]\n    return y","metadata":{"id":"Nk0haSTIc_O_","papermill":{"duration":0.065406,"end_time":"2021-02-16T16:00:00.271104","exception":false,"start_time":"2021-02-16T16:00:00.205698","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-24T23:16:36.133041Z","iopub.execute_input":"2022-11-24T23:16:36.133369Z","iopub.status.idle":"2022-11-24T23:16:36.141147Z","shell.execute_reply.started":"2022-11-24T23:16:36.133337Z","shell.execute_reply":"2022-11-24T23:16:36.139828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BirdCLEFDataset(Dataset):\n    def __init__(self, data, sr=SR, 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 = MelSpecComputer(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 = mono_to_color(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":"2022-11-24T23:16:50.678000Z","iopub.execute_input":"2022-11-24T23:16:50.678440Z","iopub.status.idle":"2022-11-24T23:16:50.692322Z","shell.execute_reply.started":"2022-11-24T23:16:50.678377Z","shell.execute_reply":"2022-11-24T23:16:50.691113Z"},"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\"]\n)\nprint(data.shape)\ndata.head()","metadata":{"id":"rVwlOrbxc_PD","outputId":"9b53ec99-634a-4b30-f9b5-75036184482b","papermill":{"duration":0.162256,"end_time":"2021-02-16T16:00:01.102152","exception":false,"start_time":"2021-02-16T16:00:00.939896","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-24T23:16:51.473560Z","iopub.execute_input":"2022-11-24T23:16:51.473891Z","iopub.status.idle":"2022-11-24T23:16:51.495917Z","shell.execute_reply.started":"2022-11-24T23:16:51.473860Z","shell.execute_reply":"2022-11-24T23:16:51.495032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.read_csv(\"../input/birdclef-2021/train_metadata.csv\")\n\nLABEL_IDS = {label: label_id for label_id,label in enumerate(sorted(df_train[\"primary_label\"].unique()))}\nINV_LABEL_IDS = {val: key for key,val in LABEL_IDS.items()}\n\nLABEL_IDS_N = {label: label_id for label_id,label in enumerate(sorted(df_train[\"primary_label\"].unique()))}\nLABEL_IDS_N.update({'nocall':397})\nINV_LABEL_IDS_N = {val: key for key,val in LABEL_IDS_N.items()}","metadata":{"execution":{"iopub.status.busy":"2022-11-24T23:16:51.946041Z","iopub.execute_input":"2022-11-24T23:16:51.946376Z","iopub.status.idle":"2022-11-24T23:16:52.174165Z","shell.execute_reply.started":"2022-11-24T23:16:51.946344Z","shell.execute_reply":"2022-11-24T23:16:52.173243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{"id":"BNsivZZZc_PG","papermill":{"duration":0.048675,"end_time":"2021-02-16T16:00:08.424781","exception":false,"start_time":"2021-02-16T16:00:08.376106","status":"completed"},"tags":[]}},{"cell_type":"code","source":"test_data = BirdCLEFDataset(data=data)\nlen(test_data), test_data[0].shape","metadata":{"id":"AzlOErOmc_PH","papermill":{"duration":0.058312,"end_time":"2021-02-16T16:00:08.531199","exception":false,"start_time":"2021-02-16T16:00:08.472887","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-24T23:16:53.044394Z","iopub.execute_input":"2022-11-24T23:16:53.044766Z","iopub.status.idle":"2022-11-24T23:16:56.473879Z","shell.execute_reply.started":"2022-11-24T23:16:53.044735Z","shell.execute_reply":"2022-11-24T23:16:56.472670Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_net(checkpoint_path, num_classes=NUM_CLASSES):\n    if \"resnest50\" in checkpoint_path:\n        print('res50')\n        net = resnest50(pretrained=False)\n        net.fc = nn.Linear(net.fc.in_features, num_classes)\n    if \"resnest50_fast_1s1x64d\" in checkpoint_path:\n        print('res50_fast')\n        net = resnest50_fast_1s1x64d(pretrained=False)\n        net.fc = nn.Linear(net.fc.in_features, num_classes)\n    elif \"resnest101\" in checkpoint_path:\n        print('res101')\n        net = resnest101(pretrained=False)\n        net.fc = nn.Linear(net.fc.in_features, num_classes)\n    elif 'tf_efficientnet_b' in checkpoint_path:\n        print('tf_b0')\n        net = timm.create_model('tf_efficientnet_b0_ns', pretrained = False)\n        net.classifier = nn.Linear(net.classifier.in_features,num_classes)\n    elif \"resnext\" in checkpoint_path:\n        print('resx50')\n        net = timm.create_model('resnext50_32x4d', pretrained=False)\n        net.fc = nn.Linear(net.fc.in_features, num_classes)\n    elif \"densenet121\" in checkpoint_path and 'seresnet50' not in checkpoint_path:\n        print('d121')\n        net = getattr(timm.models.densenet, \"densenet121\")(pretrained=False)\n        net.classifier = nn.Linear(net.classifier.in_features,num_classes)\n    elif 'seresnet50' in checkpoint_path:\n        print('ser')\n        net = getattr(timm.models.resnet, 'seresnet50')(pretrained=False)\n        net.fc = nn.Linear(net.fc.in_features, num_classes)\n    dummy_device = torch.device(\"cpu\")\n    d = torch.load(checkpoint_path, map_location=dummy_device)\n    net.load_state_dict(d['net'])\n    net = net.to(DEVICE)\n    net = net.eval()\n    return net","metadata":{"papermill":{"duration":0.05945,"end_time":"2021-02-16T16:00:14.565804","exception":false,"start_time":"2021-02-16T16:00:14.506354","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-25T00:00:07.209224Z","iopub.execute_input":"2022-11-25T00:00:07.209575Z","iopub.status.idle":"2022-11-25T00:00:07.219962Z","shell.execute_reply.started":"2022-11-25T00:00:07.209544Z","shell.execute_reply":"2022-11-25T00:00:07.218277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_net_nocall(checkpoint_path, num_classes=398):\n    if \"resnest50\" in checkpoint_path:\n        print('res50')\n        net = resnest50(pretrained=False)\n        net.fc = nn.Linear(net.fc.in_features, num_classes)\n        \n    dummy_device = torch.device(\"cpu\")\n    d = torch.load(checkpoint_path, map_location=dummy_device)\n    net.load_state_dict(d['net'])\n    net = net.to(DEVICE)\n    net = net.eval()\n    return net","metadata":{"execution":{"iopub.status.busy":"2022-11-24T23:47:53.135063Z","iopub.execute_input":"2022-11-24T23:47:53.135380Z","iopub.status.idle":"2022-11-24T23:47:53.140844Z","shell.execute_reply.started":"2022-11-24T23:47:53.135348Z","shell.execute_reply":"2022-11-24T23:47:53.139839Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint_paths_resnest = [\n    \n#     Path('../input/resnest50add/resnest50_fold1_e49_augs_mix3_3.pth'),\n#     Path('../input/resnest50add/resnest50_fold1_e50_augs_mix3_3.pth'),\n    Path('../input/resnest50add/resnest50_fold1_e51_augs_mix3_3.pth'),\n    \n    Path('../input/resnestbirdclefnew/resnest50_fold2_e46_augs_mix5.pth'),\n#     Path('../input/resnest50add/resnest50_fold2_e50_cont_augs_mix3_3.pth'),\n    \n     Path('../input/resnestbirdclefnew/resnest50_fold3_e50_augs_mix5_3.pth'),\n    \n#     Path('../input/resnest50add/resnest50_fold4_e48_augs_mix3_3.pth'), \n    Path('../input/resnest50add/resnest50_fold4_e57_augs_mix3_3.pth'),\n    Path('../input/resnest50add/resnest50_fold0_e3_augs_mix3_3_8.pth')\n]\nnets_resnest = [\n        load_net(checkpoint_path.as_posix()) for checkpoint_path in checkpoint_paths_resnest\n]","metadata":{"id":"ayVLRTzLc_PI","papermill":{"duration":14.506299,"end_time":"2021-02-16T16:00:30.268614","exception":false,"start_time":"2021-02-16T16:00:15.762315","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-25T00:00:16.581651Z","iopub.execute_input":"2022-11-25T00:00:16.581996Z","iopub.status.idle":"2022-11-25T00:00:20.534701Z","shell.execute_reply.started":"2022-11-25T00:00:16.581962Z","shell.execute_reply":"2022-11-25T00:00:20.533645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint_paths_other = [\n\n    Path('../input/birdclef-efficientnet/tf_efficientnet_b0_ns_fold0_e53_cont_augs_mix5.pth'),\n    Path('../input/birdclef-efficientnet/tf_efficientnet_b0_ns_fold1_e49_cont_augs_mix5.pth'),\n    Path('../input/birdclef-efficientnet/tf_efficientnet_b0_ns_fold3_e50_augs_mix5.pth'),\n    Path('/kaggle/input/birdclef-efficientnet/tf_efficientnet_b0_ns_fold0_e54_cont_augs_mix5.pth')\n]\n\n\nnets_other = [\n        load_net(checkpoint_path.as_posix()) for checkpoint_path in checkpoint_paths_other\n]","metadata":{"execution":{"iopub.status.busy":"2022-11-25T00:00:20.536707Z","iopub.execute_input":"2022-11-25T00:00:20.537105Z","iopub.status.idle":"2022-11-25T00:00:21.496170Z","shell.execute_reply.started":"2022-11-25T00:00:20.537055Z","shell.execute_reply":"2022-11-25T00:00:21.495273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint_paths_nocall = [\n    Path('/kaggle/input/densenet121/densenet121_fold1_e54_cont_augs_mix5.pth'),\n    Path('/kaggle/input/densenet121/seresnet50_fold3_e46_cont_augs_mix5.pth'),\n    Path('../input/densenet121/resnest101_fold0_e52_augs_mix3_3.pth'),\n]\n\n\nnets_otherall = [\n        load_net(checkpoint_path.as_posix()) for checkpoint_path in checkpoint_paths_nocall\n]","metadata":{"execution":{"iopub.status.busy":"2022-11-25T00:00:51.000625Z","iopub.execute_input":"2022-11-25T00:00:51.000995Z","iopub.status.idle":"2022-11-25T00:00:53.803679Z","shell.execute_reply.started":"2022-11-25T00:00:51.000957Z","shell.execute_reply":"2022-11-25T00:00:53.802840Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(nets_resnest), len(nets_other), len(nets_otherall)","metadata":{"execution":{"iopub.status.busy":"2022-11-25T00:06:59.390753Z","iopub.execute_input":"2022-11-25T00:06:59.391108Z","iopub.status.idle":"2022-11-25T00:06:59.396871Z","shell.execute_reply.started":"2022-11-25T00:06:59.391078Z","shell.execute_reply":"2022-11-25T00:06:59.396019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef get_thresh_preds(out, thresh=None):\n    thresh = thresh or THRESH\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":{"papermill":{"duration":0.059925,"end_time":"2021-02-16T16:00:30.912547","exception":false,"start_time":"2021-02-16T16:00:30.852622","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-11-25T00:07:29.570939Z","iopub.execute_input":"2022-11-25T00:07:29.571288Z","iopub.status.idle":"2022-11-25T00:07:29.577773Z","shell.execute_reply.started":"2022-11-25T00:07:29.571255Z","shell.execute_reply":"2022-11-25T00:07:29.576536Z"},"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":"2022-11-25T00:07:29.757072Z","iopub.execute_input":"2022-11-25T00:07:29.757393Z","iopub.status.idle":"2022-11-25T00:07:29.762982Z","shell.execute_reply.started":"2022-11-25T00:07:29.757364Z","shell.execute_reply":"2022-11-25T00:07:29.761692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_bird_names_nocall(preds):\n    bird_names = []\n    for pred in preds:\n        if not pred:\n            bird_names.append(\"!\")\n        else:\n            bird_names.append(\" \".join([INV_LABEL_IDS_N[bird_id] for bird_id in pred]))\n    return bird_names","metadata":{"execution":{"iopub.status.busy":"2022-11-25T00:07:29.940944Z","iopub.execute_input":"2022-11-25T00:07:29.941304Z","iopub.status.idle":"2022-11-25T00:07:29.946899Z","shell.execute_reply.started":"2022-11-25T00:07:29.941271Z","shell.execute_reply":"2022-11-25T00:07:29.945965Z"},"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 SAMPLE_SUB_PATH:\n        sample_sub = pd.read_csv(SAMPLE_SUB_PATH, 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":"2022-11-25T00:07:30.048974Z","iopub.execute_input":"2022-11-25T00:07:30.049304Z","iopub.status.idle":"2022-11-25T00:07:30.057469Z","shell.execute_reply.started":"2022-11-25T00:07:30.049274Z","shell.execute_reply":"2022-11-25T00:07:30.055634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(test_data, names=False):\n    big_preds = []\n    for value,nets in enumerate([nets_resnest, nets_other, nets_otherall]):\n        print(value)\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.sigmoid(o)\n                    \n                    pred += o\n\n                pred /= len(nets)\n                preds.append(pred)\n        big_preds.append(preds)\n    return big_preds","metadata":{"execution":{"iopub.status.busy":"2022-11-25T00:07:30.188211Z","iopub.execute_input":"2022-11-25T00:07:30.188600Z","iopub.status.idle":"2022-11-25T00:07:30.199899Z","shell.execute_reply.started":"2022-11-25T00:07:30.188555Z","shell.execute_reply":"2022-11-25T00:07:30.198520Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_probas = predict(test_data, names=False)\nprint(len(pred_probas))\n# print(pred_probas)","metadata":{"execution":{"iopub.status.busy":"2022-11-25T00:07:30.313693Z","iopub.execute_input":"2022-11-25T00:07:30.314089Z","iopub.status.idle":"2022-11-25T00:11:23.220742Z","shell.execute_reply.started":"2022-11-25T00:07:30.314030Z","shell.execute_reply":"2022-11-25T00:11:23.219578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check weights carefully\nweights = np.array([0.4,0.4,.2])\n# weights = np.array([1.])\n# weights = np.array([1.,0.])\nsub_score = np.sum(pred_probas*weights[:, None], 0)\n\nsub_score = sub_score.tolist()\nlen(sub_score)","metadata":{"execution":{"iopub.status.busy":"2022-11-25T00:12:59.170779Z","iopub.execute_input":"2022-11-25T00:12:59.171137Z","iopub.status.idle":"2022-11-25T00:12:59.182571Z","shell.execute_reply.started":"2022-11-25T00:12:59.171106Z","shell.execute_reply":"2022-11-25T00:12:59.181500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = [get_bird_names(get_thresh_preds(pred, thresh=0.25)) for pred in sub_score]\n# preds[:2]","metadata":{"execution":{"iopub.status.busy":"2022-11-25T00:13:00.864144Z","iopub.execute_input":"2022-11-25T00:13:00.864486Z","iopub.status.idle":"2022-11-25T00:13:00.941936Z","shell.execute_reply.started":"2022-11-25T00:13:00.864453Z","shell.execute_reply":"2022-11-25T00:13:00.941087Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = preds_as_df(data, preds)\nprint(sub.shape)","metadata":{"execution":{"iopub.status.busy":"2022-11-25T00:13:01.603945Z","iopub.execute_input":"2022-11-25T00:13:01.604277Z","iopub.status.idle":"2022-11-25T00:13:01.611624Z","shell.execute_reply.started":"2022-11-25T00:13:01.604243Z","shell.execute_reply":"2022-11-25T00:13:01.610683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub.head()","metadata":{"execution":{"iopub.status.busy":"2022-11-25T00:13:01.922296Z","iopub.execute_input":"2022-11-25T00:13:01.922685Z","iopub.status.idle":"2022-11-25T00:13:01.932487Z","shell.execute_reply.started":"2022-11-25T00:13:01.922650Z","shell.execute_reply":"2022-11-25T00:13:01.931605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-11-25T00:13:02.191720Z","iopub.execute_input":"2022-11-25T00:13:02.192020Z","iopub.status.idle":"2022-11-25T00:13:02.203944Z","shell.execute_reply.started":"2022-11-25T00:13:02.191992Z","shell.execute_reply":"2022-11-25T00:13:02.202925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":" sub.shape","metadata":{"execution":{"iopub.status.busy":"2022-11-25T00:13:02.435189Z","iopub.execute_input":"2022-11-25T00:13:02.435519Z","iopub.status.idle":"2022-11-25T00:13:02.441385Z","shell.execute_reply.started":"2022-11-25T00:13:02.435486Z","shell.execute_reply":"2022-11-25T00:13:02.440353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Small validation","metadata":{}},{"cell_type":"code","source":"def get_metrics(s_true, s_pred):\n    s_true = set(s_true.split())\n    s_pred = set(s_pred.split())\n    n, n_true, n_pred = len(s_true.intersection(s_pred)), len(s_true), len(s_pred)\n    \n    prec = n/n_pred\n    rec = n/n_true\n    f1 = 2*prec*rec/(prec + rec) if prec + rec else 0\n    \n    return {\"f1\": f1, \"prec\": prec, \"rec\": rec, \"n_true\": n_true, \"n_pred\": n_pred, \"n\": n}","metadata":{"execution":{"iopub.status.busy":"2022-11-25T00:13:03.572605Z","iopub.execute_input":"2022-11-25T00:13:03.572968Z","iopub.status.idle":"2022-11-25T00:13:03.579547Z","shell.execute_reply.started":"2022-11-25T00:13:03.572924Z","shell.execute_reply":"2022-11-25T00:13:03.578431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TARGET_PATH:\n    sub_target = pd.read_csv(TARGET_PATH)\n    sub_target = sub_target.merge(sub, how=\"left\", on=\"row_id\")\n    \n    print(sub_target[\"birds_x\"].notnull().sum(), sub_target[\"birds_x\"].notnull().sum())\n    assert sub_target[\"birds_x\"].notnull().all()\n    assert sub_target[\"birds_y\"].notnull().all()\n    \n    df_metrics = pd.DataFrame([get_metrics(s_true, s_pred) for s_true, s_pred in zip(sub_target.birds_x, sub_target.birds_y)])\n    \n    print(df_metrics.mean())\n    ","metadata":{"execution":{"iopub.status.busy":"2022-11-25T00:13:04.039609Z","iopub.execute_input":"2022-11-25T00:13:04.039935Z","iopub.status.idle":"2022-11-25T00:13:04.074950Z","shell.execute_reply.started":"2022-11-25T00:13:04.039904Z","shell.execute_reply":"2022-11-25T00:13:04.073970Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_target[sub_target.birds_x != \"nocall\"].shape","metadata":{"execution":{"iopub.status.busy":"2022-11-24T23:30:16.062607Z","iopub.execute_input":"2022-11-24T23:30:16.063060Z","iopub.status.idle":"2022-11-24T23:30:16.073602Z","shell.execute_reply.started":"2022-11-24T23:30:16.063014Z","shell.execute_reply":"2022-11-24T23:30:16.072293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"correct = sub_target[sub_target['birds_x'] == sub_target['birds_y']]\ncorrect.shape[0],correct[correct['birds_y']!='nocall'].shape[0]","metadata":{"execution":{"iopub.status.busy":"2022-11-24T23:30:16.075334Z","iopub.execute_input":"2022-11-24T23:30:16.076013Z","iopub.status.idle":"2022-11-24T23:30:16.087506Z","shell.execute_reply.started":"2022-11-24T23:30:16.075954Z","shell.execute_reply":"2022-11-24T23:30:16.086134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_target['birds_y'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2021-05-20T03:21:01.874768Z","iopub.execute_input":"2021-05-20T03:21:01.875109Z","iopub.status.idle":"2021-05-20T03:21:01.883705Z","shell.execute_reply.started":"2021-05-20T03:21:01.87508Z","shell.execute_reply":"2021-05-20T03:21:01.882871Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_target['birds_x'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2021-05-20T03:19:46.272521Z","iopub.execute_input":"2021-05-20T03:19:46.27286Z","iopub.status.idle":"2021-05-20T03:19:46.281525Z","shell.execute_reply.started":"2021-05-20T03:19:46.272828Z","shell.execute_reply":"2021-05-20T03:19:46.280384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}