{"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":"markdown","source":"# Summary of this notebook\n\nIn this notebook, we are going to calculate 397dims birdcall probabilities for train_short_audio\n\n# input & output\n\n[input]\n\nbirdclef-2021 (original data)\n\nmelspectrogram multilabel classifier models (Ⅰ)\n\n7sec clip melspectrogram images of train_short_audio\n(https://www.kaggle.com/datasets/huyhoang333/birdclef2021augmentedaudio-melspec-p1 to -p4)\n\nnocall detector output for train_short_audio\n\nhttps://www.kaggle.com/datasets/huyhoang333/augmented-train-short-audio-nocall-fold0to4\n\nsklearn library (To use StratifiedGroupKfold, we have to install scikit-learn 1.0.dev0)\n\nhttps://www.kaggle.com/namakemono/scikit-learn-10dev0\n\n[output]\n\n397dims birdcall probabilities for train_short_audio (with some more features)","metadata":{"papermill":{"duration":0.019255,"end_time":"2021-06-03T07:14:09.159488","exception":false,"start_time":"2021-06-03T07:14:09.140233","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install -q pysndfx SoundFile audiomentations pretrainedmodels efficientnet_pytorch resnest\n!pip install ../input/scikit-learn-10dev0/scikit_learn-1.0.dev0-cp37-cp37m-manylinux2010_x86_64.whl\n!pip install timm","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:09.208163Z","iopub.status.busy":"2021-06-03T07:14:09.207601Z","iopub.status.idle":"2021-06-03T07:14:37.987274Z","shell.execute_reply":"2021-06-03T07:14:37.986311Z"},"id":"Yn1Ybf15VAqW","outputId":"c737c334-3974-472b-f1b7-010bedd167c4","papermill":{"duration":28.809512,"end_time":"2021-06-03T07:14:37.987445","exception":false,"start_time":"2021-06-03T07:14:09.177933","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport librosa as lb\nimport librosa.display as lbd\nimport soundfile as sf\nfrom  soundfile import SoundFile\nimport pandas as pd\nfrom  IPython.display import Audio\nfrom pathlib import Path\n\nimport torch\nfrom torch import nn, optim\nfrom  torch.utils.data import Dataset, DataLoader\n\nfrom resnest.torch import resnest50\n\nfrom matplotlib import pyplot as plt\n\nimport os, random, gc\nimport re, time, json\nfrom  ast import literal_eval\n\nfrom IPython.display import Audio\nfrom sklearn.metrics import label_ranking_average_precision_score\n\nfrom tqdm.notebook import tqdm\nimport joblib\n\nimport timm\nfrom sklearn.model_selection import StratifiedGroupKFold\n\nfrom efficientnet_pytorch import EfficientNet\nimport pretrainedmodels\nimport resnest.torch as resnest_torch","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:38.040395Z","iopub.status.busy":"2021-06-03T07:14:38.039566Z","iopub.status.idle":"2021-06-03T07:14:42.999349Z","shell.execute_reply":"2021-06-03T07:14:42.998744Z"},"id":"2dt7oG43VAqc","papermill":{"duration":4.98933,"end_time":"2021-06-03T07:14:42.999503","exception":false,"start_time":"2021-06-03T07:14:38.010173","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not os.path.exists(\"resnest50-528c19ca.pth\"):\n    !wget  \"https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-resnest/resnest50-528c19ca.pth\"","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:43.059913Z","iopub.status.busy":"2021-06-03T07:14:43.057985Z","iopub.status.idle":"2021-06-03T07:14:45.632745Z","shell.execute_reply":"2021-06-03T07:14:45.633202Z"},"papermill":{"duration":2.611937,"end_time":"2021-06-03T07:14:45.633372","exception":false,"start_time":"2021-06-03T07:14:43.021435","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    \nseed_everything()","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:45.690216Z","iopub.status.busy":"2021-06-03T07:14:45.689379Z","iopub.status.idle":"2021-06-03T07:14:45.694813Z","shell.execute_reply":"2021-06-03T07:14:45.694404Z"},"id":"Q39ZsGAhVAqe","papermill":{"duration":0.036159,"end_time":"2021-06-03T07:14:45.694934","exception":false,"start_time":"2021-06-03T07:14:45.658775","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_CLASSES = 397\nSR = 32_000\nDURATION = 7\n\nMAX_READ_SAMPLES = 15 # Each record will have 10 melspecs at most, you can increase this on Colab with High Memory Enabled","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:45.749965Z","iopub.status.busy":"2021-06-03T07:14:45.748974Z","iopub.status.idle":"2021-06-03T07:14:45.751879Z","shell.execute_reply":"2021-06-03T07:14:45.751463Z"},"id":"KjE3paAQHRSW","papermill":{"duration":0.031497,"end_time":"2021-06-03T07:14:45.751988","exception":false,"start_time":"2021-06-03T07:14:45.720491","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    def __init__(self, debug:bool):\n        self.debug = debug\n        \n        self.epochs = 1 if self.debug else 50\n\n        self.max_distance = None # choose from [10, 20, None]\n        if self.max_distance is not None:\n            self.sites = [\"SSW\"] # choose multiples from [\"COL\", \"COR\", \"SNE\", \"SSW\"]\n        else:\n            self.sites = None\n        self.max_duration = None # choose from [15, 30, 60, None]\n        self.min_rating = None # choose from [3, 4, None], best: 3?\n        self.max_spieces = None # choose from [100, 200, 300, None], best: 300?\n        self.confidence_ub = 0.995 # Probability of birdsong occurrence, default: 0.995, choose from [0.5, 0.7, 0.9, 0.995]\n        self.use_high_confidence_only = False # Whether to use only frames that are likely to be ringing (False performed better).\n        self.use_mixup = True\n        self.mixup_alpha = 5.0 # 0.5\n        self.grouped_by_author = True\n        # self.folds = [4]\n\n        self.suffix = f\"sr{SR}_d{DURATION}\"\n        if self.max_spieces:\n            self.suffix += f\"_spices-{self.max_spieces}\"\n        if self.min_rating:\n            self.suffix += f\"_rating-{self.min_rating}\"\n        if self.use_high_confidence_only:\n            self.suffix += f\"_high-confidence-only\"\n        if self.use_mixup:\n            self.suffix += f\"_miixup-{self.mixup_alpha}\"\n        if self.grouped_by_author:\n            self.suffix += f\"_grouped-by-auther\"\n\n    def to_dict(self):\n        return {\n            \"debug\": self.debug,\n            \"epochs\": self.epochs,\n            \"max_distance\": self.max_distance,\n            \"sites\": self.sites,\n            \"max_duration\": self.max_duration,\n            \"min_rating\": self.min_rating,\n            \"max_spieces\": self.max_spieces,\n            \"confidence_ub\": self.confidence_ub,\n            \"use_high_confidence_only\": self.use_high_confidence_only,\n            \"use_mixup\": self.use_mixup,\n            \"mixup_alpha\": self.mixup_alpha,\n            \"suffix\": self.suffix,\n            \"grouped_by_author\": self.grouped_by_author\n        }\n\nconfig = Config(debug=False)\nfrom pprint import pprint\npprint(config.to_dict())","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:45.81611Z","iopub.status.busy":"2021-06-03T07:14:45.815224Z","iopub.status.idle":"2021-06-03T07:14:45.820591Z","shell.execute_reply":"2021-06-03T07:14:45.820158Z"},"id":"a3ew2nhVq5Py","outputId":"97576deb-c494-488b-fbb1-5b569a47fb11","papermill":{"duration":0.0437,"end_time":"2021-06-03T07:14:45.82071","exception":false,"start_time":"2021-06-03T07:14:45.77701","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODEL_NAMES = [\n    # \"resnext101_32x8d_wsl\",\n    # 'efficientnet_b0',\n    \"resnest50\",\n    # \"densenet121\",\n] ","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:45.876667Z","iopub.status.busy":"2021-06-03T07:14:45.875832Z","iopub.status.idle":"2021-06-03T07:14:45.878773Z","shell.execute_reply":"2021-06-03T07:14:45.87835Z"},"id":"tWgVCLLJwMFk","papermill":{"duration":0.032737,"end_time":"2021-06-03T07:14:45.878892","exception":false,"start_time":"2021-06-03T07:14:45.846155","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MEL_PATHS = sorted(Path(\"../input\").glob(\"birdclef2021augmentedaudio-melspec-p?/rich_train_metadata.csv\"))\nTRAIN_LABEL_PATHS = sorted(Path(\"../input\").glob(\"birdclef2021augmentedaudio-melspec-p?/LABEL_IDS.json\"))\n\nMODEL_ROOT = Path(\".\")","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:45.936098Z","iopub.status.busy":"2021-06-03T07:14:45.935549Z","iopub.status.idle":"2021-06-03T07:14:45.957219Z","shell.execute_reply":"2021-06-03T07:14:45.956717Z"},"id":"2NfkUn9SCWs6","papermill":{"duration":0.052521,"end_time":"2021-06-03T07:14:45.957356","exception":false,"start_time":"2021-06-03T07:14:45.904835","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_BATCH_SIZE = 50 # 16\nTRAIN_NUM_WORKERS = 2\n\nVAL_BATCH_SIZE = 50 # 16 # 128\nVAL_NUM_WORKERS = 2\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nprint(\"Device:\", DEVICE)","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:46.057073Z","iopub.status.busy":"2021-06-03T07:14:46.055557Z","iopub.status.idle":"2021-06-03T07:14:46.060418Z","shell.execute_reply":"2021-06-03T07:14:46.060784Z"},"id":"Iu56f-7VVAqf","outputId":"0adfad52-a0a0-43ee-b0c9-c807a8da023b","papermill":{"duration":0.076428,"end_time":"2021-06-03T07:14:46.060931","exception":false,"start_time":"2021-06-03T07:14:45.984503","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint_paths = [\n    # model1\n    Path(\"../input/augmentedclassifiermodels/augmented-clefmodel/birdclef_resnest26d_fold0_epoch_26_f1_val_04734_20230707180942.pth\"),\n]","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:46.117355Z","iopub.status.busy":"2021-06-03T07:14:46.116631Z","iopub.status.idle":"2021-06-03T07:14:46.119517Z","shell.execute_reply":"2021-06-03T07:14:46.118985Z"},"id":"GuyjwJnACWs6","papermill":{"duration":0.03241,"end_time":"2021-06-03T07:14:46.11964","exception":false,"start_time":"2021-06-03T07:14:46.08723","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_df(mel_paths=MEL_PATHS, train_label_paths=TRAIN_LABEL_PATHS):\n  df = None\n  LABEL_IDS = {}\n    \n  for file_path in mel_paths:\n    temp = pd.read_csv(str(file_path), index_col=0)\n    temp[\"impath\"] = temp.apply(lambda row: file_path.parent/\"audio_images/{}/{}.npy\".format(row.primary_label, row.filename), axis=1) \n    df = temp if df is None else df.append(temp)\n    \n  df[\"secondary_labels\"] = df[\"secondary_labels\"].apply(literal_eval)\n\n  for file_path in train_label_paths:\n    with open(str(file_path)) as f:\n      LABEL_IDS.update(json.load(f))\n\n  return LABEL_IDS, df","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:46.179391Z","iopub.status.busy":"2021-06-03T07:14:46.178602Z","iopub.status.idle":"2021-06-03T07:14:46.181193Z","shell.execute_reply":"2021-06-03T07:14:46.181555Z"},"id":"9BcwFTpLp4CZ","papermill":{"duration":0.035477,"end_time":"2021-06-03T07:14:46.181688","exception":false,"start_time":"2021-06-03T07:14:46.146211","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import List\ndef get_locations() -> List[dict]:\n    return [{\n        \"site\": \"COL\",\n        \"latitude\": 5.57,\n        \"longitude\": -75.85\n    }, {\n        \"site\": \"COR\",\n        \"latitude\": 10.12,\n        \"longitude\": -84.51\n    }, {\n        \"site\": \"SNE\",\n        \"latitude\": 38.49,\n        \"longitude\": -119.95\n    }, {\n        \"site\": \"SSW\",\n        \"latitude\": 42.47,\n        \"longitude\": -76.45\n    }]\n\ndef is_in_site(row, sites, max_distance):\n    for location in get_locations():\n        if location[\"site\"] in sites:\n            x = (row[\"latitude\"] - location[\"latitude\"])\n            y = (row[\"longitude\"] - location[\"longitude\"])\n            r = (x**2 + y**2) ** 0.5\n            if r < max_distance:\n                return True\n    return False","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:46.242607Z","iopub.status.busy":"2021-06-03T07:14:46.241806Z","iopub.status.idle":"2021-06-03T07:14:46.244501Z","shell.execute_reply":"2021-06-03T07:14:46.244071Z"},"id":"JayFwqVAI8A_","papermill":{"duration":0.036916,"end_time":"2021-06-03T07:14:46.24461","exception":false,"start_time":"2021-06-03T07:14:46.207694","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"LABEL_IDS, df = get_df()\n\nif config.grouped_by_author:\n    kf = StratifiedGroupKFold(n_splits=5)\n    x = df[[\"latitude\", \"longitude\"]].values\n    y = df[\"label_id\"].values\n    groups = df[\"author\"].values\n    df[\"fold\"] = -1\n    for kfold_index, (train_index, valid_index) in enumerate(kf.split(x, y, groups)):\n        df.loc[valid_index, \"fold\"] = kfold_index\n\nif config.debug:\n    df = df.head(100)\n\nprint(\"before:%d\" % len(df))\n# Within a certain distance of the target area\nif config.max_distance is not None:\n    df = df[df.apply(lambda row: is_in_site(row, config.sites, config.max_distance), axis=1)]\n# Number of Species\nif config.max_spieces is not None:\n    s = df[\"primary_label\"].value_counts().head(config.max_spieces)\n    df = df[df[\"primary_label\"].isin(s.index)]\n# Rating is above a certain value\nif config.min_rating is not None:\n    df = df[df[\"rating\"] >= config.min_rating]\n# Within a certain amount of recording time\nif config.max_duration is not None:\n    df = df[df[\"duration\"] < config.max_duration]\ndf = df.reset_index(drop=True)\nprint(\"after:%d\" % len(df))\n\nprint(df.shape)\ndf.head()","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:46.307963Z","iopub.status.busy":"2021-06-03T07:14:46.307424Z","iopub.status.idle":"2021-06-03T07:14:51.066186Z","shell.execute_reply":"2021-06-03T07:14:51.065571Z"},"id":"CM4lmCl2p4Ca","outputId":"7787fcdb-3b54-40cd-87db-8b21118c6bbb","papermill":{"duration":4.79471,"end_time":"2021-06-03T07:14:51.066346","exception":false,"start_time":"2021-06-03T07:14:46.271636","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model(name, num_classes=NUM_CLASSES):\n    \"\"\"\n    Loads a pretrained model. \n    Supports ResNest, ResNext-wsl, EfficientNet, ResNext and ResNet.\n\n    Arguments:\n        name {str} -- Name of the model to load\n\n    Keyword Arguments:\n        num_classes {int} -- Number of classes to use (default: {1})\n\n    Returns:\n        torch model -- Pretrained model\n    \"\"\"\n    if \"resnest\" in name:\n        if not os.path.exists(\"resnest50-528c19ca.pth\"):\n            !wget https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-resnest/resnest50-528c19ca.pth\n    \n        pretrained_weights = torch.load('resnest50-528c19ca.pth')\n        model = getattr(resnest_torch, name)(pretrained=False)\n        model.load_state_dict(pretrained_weights)\n    elif \"wsl\" in name:\n        model = torch.hub.load(\"facebookresearch/WSL-Images\", name)\n    elif name.startswith(\"resnext\") or  name.startswith(\"resnet\"):\n        model = torch.hub.load(\"pytorch/vision:v0.6.0\", name, pretrained=True)\n    elif name.startswith(\"efficientnet_b\"):\n        model = getattr(timm.models.efficientnet, name)(pretrained=True)\n    elif name.startswith(\"densenet\"):\n        model = getattr(timm.models.densenet, name)(pretrained=True)\n    elif \"efficientnet-b\" in name:\n        model = EfficientNet.from_pretrained(name)\n    else:\n        model = pretrainedmodels.__dict__[name](pretrained='imagenet')\n\n    if hasattr(model, \"fc\"):\n        nb_ft = model.fc.in_features\n        model.fc = nn.Linear(nb_ft, num_classes)\n    elif hasattr(model, \"_fc\"):\n        nb_ft = model._fc.in_features\n        model._fc = nn.Linear(nb_ft, num_classes)\n    elif hasattr(model, \"classifier\"):\n        nb_ft = model.classifier.in_features\n        model.classifier = nn.Linear(nb_ft, num_classes)\n    elif hasattr(model, \"last_linear\"):\n        nb_ft = model.last_linear.in_features\n        model.last_linear = nn.Linear(nb_ft, num_classes)\n\n    return model","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:51.142868Z","iopub.status.busy":"2021-06-03T07:14:51.141487Z","iopub.status.idle":"2021-06-03T07:14:51.143803Z","shell.execute_reply":"2021-06-03T07:14:51.144252Z"},"id":"OGPDuihmVAqi","papermill":{"duration":0.047392,"end_time":"2021-06-03T07:14:51.144397","exception":false,"start_time":"2021-06-03T07:14:51.097005","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BirdClefDataset(Dataset):\n    def __init__(\n        self,\n        meta,\n        sr=SR,\n        is_train=True,\n        num_classes=NUM_CLASSES,\n        duration=DURATION\n    ):\n        self.meta = meta.copy().reset_index(drop=True)\n        records = []\n        for idx, row in tqdm(self.meta.iterrows(), total=len(self.meta)):\n            images = np.load(str(row[\"impath\"]))\n            for i, image in enumerate(images):\n                seconds = i * duration\n                records.append({\n                    \"filename\": row[\"filename\"],\n                    \"impath\": row[\"impath\"],\n                    \"seconds\": seconds,\n                    \"index\": i\n                })\n        self.records = records\n        self.sr = sr\n        self.is_train = is_train\n        self.num_classes = num_classes\n        self.duration = duration\n        self.audio_length = self.duration*self.sr\n        self.eps = 0.0025\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 __len__(self):\n        return len(self.records)\n\n    def __getitem__(self, idx):\n        row = self.records[idx]\n        image = np.load(str(row[\"impath\"]))[row[\"index\"]]\n        image = self.normalize(image)\n        return image, row[\"filename\"], row[\"seconds\"]","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:51.209039Z","iopub.status.busy":"2021-06-03T07:14:51.208534Z","iopub.status.idle":"2021-06-03T07:14:51.212006Z","shell.execute_reply":"2021-06-03T07:14:51.212433Z"},"id":"sqpRO4bFVG26","papermill":{"duration":0.040663,"end_time":"2021-06-03T07:14:51.212567","exception":false,"start_time":"2021-06-03T07:14:51.171904","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"nocall_df = pd.read_csv(\"../input/augmented-train-short-audio-nocall-fold0to4/augmented_nocalldetection_for_shortaudio_fold0.csv\")","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:51.276647Z","iopub.status.busy":"2021-06-03T07:14:51.276128Z","iopub.status.idle":"2021-06-03T07:14:51.881287Z","shell.execute_reply":"2021-06-03T07:14:51.880324Z"},"id":"1sdNWK4RVTeY","papermill":{"duration":0.641394,"end_time":"2021-06-03T07:14:51.881431","exception":false,"start_time":"2021-06-03T07:14:51.240037","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nds = BirdClefDataset(meta=df, sr=SR, duration=DURATION, is_train=True)\nlen(df)","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:14:51.943216Z","iopub.status.busy":"2021-06-03T07:14:51.942293Z","iopub.status.idle":"2021-06-03T07:24:04.766199Z","shell.execute_reply":"2021-06-03T07:24:04.765563Z"},"id":"Np-56XrXVAqm","outputId":"f02597b7-41ab-41c1-b6c7-74600fe53fcb","papermill":{"duration":552.857082,"end_time":"2021-06-03T07:24:04.766371","exception":false,"start_time":"2021-06-03T07:14:51.909289","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def add_tail(model, num_classes):\n    if hasattr(model, \"fc\"):\n        nb_ft = model.fc.in_features\n        model.fc = nn.Linear(nb_ft, num_classes)\n    elif hasattr(model, \"_fc\"):\n        nb_ft = model._fc.in_features\n        model._fc = nn.Linear(nb_ft, num_classes)\n    elif hasattr(model, \"classifier\"):\n        nb_ft = model.classifier.in_features\n        model.classifier = nn.Linear(nb_ft, num_classes)\n    elif hasattr(model, \"last_linear\"):\n        nb_ft = model.last_linear.in_features\n        model.last_linear = nn.Linear(nb_ft, num_classes)\n    return model\n\ndef load_net(checkpoint_path, num_classes=NUM_CLASSES):\n    if \"resnest50\" in checkpoint_path:\n        net = resnest50(pretrained=False)\n    elif \"resnest26d\" in checkpoint_path:\n        net = timm.models.resnest26d(pretrained=False)\n    elif \"resnext101_32x8d_wsl\" in checkpoint_path:\n        net = torch.hub.load(\"facebookresearch/WSL-Images\", \"resnext101_32x8d_wsl\")\n    elif \"efficientnet_b0\" in checkpoint_path:\n        net = getattr(timm.models.efficientnet, \"efficientnet_b0\")(pretrained=False)\n    elif \"densenet121\" in checkpoint_path:\n        net = timm.models.densenet121(pretrained=False)\n    else:\n        raise ValueError(\"Unexpected checkpont name: %s\" % checkpoint_path)\n    net = add_tail(net, num_classes)\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    net.load_state_dict(d)\n    net = net.to(DEVICE)\n    net = net.eval()\n    return net","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:24:04.835248Z","iopub.status.busy":"2021-06-03T07:24:04.834567Z","iopub.status.idle":"2021-06-03T07:24:04.837611Z","shell.execute_reply":"2021-06-03T07:24:04.837189Z"},"id":"Z50-hLSoegMq","papermill":{"duration":0.042033,"end_time":"2021-06-03T07:24:04.837721","exception":false,"start_time":"2021-06-03T07:24:04.795688","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef predict(net, criterion, val_laoder):\n    net.eval()\n    records = []\n    val_laoder = tqdm(val_laoder, leave = False, total=len(val_laoder))\n    for icount, (xb, filename, seconds) in enumerate(val_laoder):\n        xb = xb.to(DEVICE)\n        prob = net(xb)\n        prob = torch.sigmoid(prob)\n        records.append({\n            \"prob\": prob,\n            \"filename\": filename,\n            \"seconds\": seconds\n        })\n    return records","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:24:04.901639Z","iopub.status.busy":"2021-06-03T07:24:04.900842Z","iopub.status.idle":"2021-06-03T07:24:04.904549Z","shell.execute_reply":"2021-06-03T07:24:04.904157Z"},"id":"vxRZD0AGr2i_","papermill":{"duration":0.038069,"end_time":"2021-06-03T07:24:04.904662","exception":false,"start_time":"2021-06-03T07:24:04.866593","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def one_fold(checkpoint_path, fold, train_set, val_set, epochs=20, save=True, save_root=None):\n    net = load_net(checkpoint_path)\n    criterion = nn.BCEWithLogitsLoss()\n    val_data = BirdClefDataset(\n        meta=df.iloc[val_set].reset_index(drop=True),\n        sr=SR,\n        duration=DURATION,\n        is_train=False\n    )\n    val_laoder = DataLoader(val_data, batch_size=VAL_BATCH_SIZE, num_workers=VAL_NUM_WORKERS, shuffle=False)\n    y_preda = predict(net, criterion, val_laoder)\n    return y_preda\n","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:24:04.968592Z","iopub.status.busy":"2021-06-03T07:24:04.967898Z","iopub.status.idle":"2021-06-03T07:24:04.971581Z","shell.execute_reply":"2021-06-03T07:24:04.971081Z"},"id":"gz72cyTbkb3p","papermill":{"duration":0.038132,"end_time":"2021-06-03T07:24:04.971703","exception":false,"start_time":"2021-06-03T07:24:04.933571","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict_for_oof(checkpoint_path, epochs=20, save=True, n_splits=5, seed=177, save_root=None, suffix=\"\", folds=None):\n  gc.collect()\n  torch.cuda.empty_cache()\n\n  fold_bar = tqdm(df.reset_index().groupby(\"fold\").index.apply(list).items(), total=df.fold.max()+1)\n  \n  for fold, val_set in fold_bar:\n      if folds and not fold in folds:\n        continue\n      \n      print(f\"\\n############################### [FOLD {fold}]\")\n      fold_bar.set_description(f\"[FOLD {fold}]\")\n      train_set = np.setdiff1d(df.index, val_set)\n      records = one_fold(checkpoint_path, fold=fold, train_set=train_set , val_set=val_set , epochs=epochs, save=save, save_root=save_root)\n      gc.collect()\n      torch.cuda.empty_cache()\n      return records","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:24:05.037731Z","iopub.status.busy":"2021-06-03T07:24:05.037038Z","iopub.status.idle":"2021-06-03T07:24:05.040517Z","shell.execute_reply":"2021-06-03T07:24:05.040098Z"},"id":"oi6o_EYskb04","papermill":{"duration":0.040254,"end_time":"2021-06-03T07:24:05.040627","exception":false,"start_time":"2021-06-03T07:24:05.000373","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def to_call_prob(row):\n    i = row[\"seconds\"] // DURATION\n    call_prob = float(row[\"nocalldetection\"].split()[i])\n    return call_prob","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:24:05.104471Z","iopub.status.busy":"2021-06-03T07:24:05.103909Z","iopub.status.idle":"2021-06-03T07:24:05.10779Z","shell.execute_reply":"2021-06-03T07:24:05.107395Z"},"id":"hECAVdswTRd-","papermill":{"duration":0.03782,"end_time":"2021-06-03T07:24:05.107906","exception":false,"start_time":"2021-06-03T07:24:05.070086","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def to_birds(row):\n    if row[\"call_prob\"] < 0.5:\n        return \"nocall\"\n    res = [row[\"primary_label\"]] + eval(row[\"secondary_labels\"])\n    return \" \".join(res)","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:24:05.171009Z","iopub.status.busy":"2021-06-03T07:24:05.170498Z","iopub.status.idle":"2021-06-03T07:24:05.17458Z","shell.execute_reply":"2021-06-03T07:24:05.174164Z"},"id":"7qR90aHdI56w","papermill":{"duration":0.037468,"end_time":"2021-06-03T07:24:05.174688","exception":false,"start_time":"2021-06-03T07:24:05.13722","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"INV_LABEL_IDS = {v:k for k, v in LABEL_IDS.items()}\ncolumns = [INV_LABEL_IDS[i] for i in range(len(LABEL_IDS))]\nmetadata_df = pd.read_csv(\"../input/birdclef-2021/train_metadata.csv\")\nnocall_df = pd.read_csv(\"../input/augmented-train-short-audio-nocall-fold0to4/augmented_nocalldetection_for_shortaudio_fold0.csv\")","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:24:05.243462Z","iopub.status.busy":"2021-06-03T07:24:05.242854Z","iopub.status.idle":"2021-06-03T07:24:06.320888Z","shell.execute_reply":"2021-06-03T07:24:06.32044Z"},"id":"AamWh4cKMD2o","papermill":{"duration":1.11721,"end_time":"2021-06-03T07:24:06.321019","exception":false,"start_time":"2021-06-03T07:24:05.203809","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\nfilename_to_nocalldetection = {}\nfilepath_list = list(glob.glob(\"../input/augmented-train-short-audio-nocall-fold0to4/*.csv\"))\nfor filepath in filepath_list:\n    nocall_df = pd.read_csv(filepath)\n    probs = nocall_df[\"nocalldetection\"].apply(\n        lambda _: list(\n            map(float, _.split())\n        )\n    ).tolist()\n    for k, v in zip(nocall_df[\"filename\"].tolist(), probs):\n        if not k in filename_to_nocalldetection:\n            filename_to_nocalldetection[k] = v\n        else:\n            w = filename_to_nocalldetection[k]\n            for i in range( len(w)):\n                w[i] += v[i]\n            filename_to_nocalldetection[k] = w\n\nfor k, v in filename_to_nocalldetection.items():\n    for i in range(len(v)):\n        filename_to_nocalldetection[k][i] /= len(filepath_list)\n\nfor k, v in filename_to_nocalldetection.items():\n    filename_to_nocalldetection[k] = \" \".join(map(str, v))","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:24:06.388233Z","iopub.status.busy":"2021-06-03T07:24:06.387519Z","iopub.status.idle":"2021-06-03T07:24:13.109646Z","shell.execute_reply":"2021-06-03T07:24:13.109193Z"},"id":"G8sJHDogMY3w","papermill":{"duration":6.759227,"end_time":"2021-06-03T07:24:13.109778","exception":false,"start_time":"2021-06-03T07:24:06.350551","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"nocall_df = pd.DataFrame(filename_to_nocalldetection.items(), columns=[\"filename\", \"nocalldetection\"])","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:24:13.191254Z","iopub.status.busy":"2021-06-03T07:24:13.190667Z","iopub.status.idle":"2021-06-03T07:24:13.218119Z","shell.execute_reply":"2021-06-03T07:24:13.217593Z"},"id":"mQQ7s-0jMYqi","papermill":{"duration":0.078512,"end_time":"2021-06-03T07:24:13.218244","exception":false,"start_time":"2021-06-03T07:24:13.139732","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"nocall_df.head()","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:24:13.289076Z","iopub.status.busy":"2021-06-03T07:24:13.288395Z","iopub.status.idle":"2021-06-03T07:24:13.291663Z","shell.execute_reply":"2021-06-03T07:24:13.292147Z"},"id":"XGwpHLOPMYfp","outputId":"bbb80044-1723-4487-e195-6a30de24eea4","papermill":{"duration":0.044421,"end_time":"2021-06-03T07:24:13.292306","exception":false,"start_time":"2021-06-03T07:24:13.247885","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"filepath_list = []\nfor checkpoint_path in checkpoint_paths:\n    print(\"\\n\\n###########################################\", checkpoint_path)\n    # Find out which fold it is from the name of the model file.\n    fold = -1\n    for i in range(5):\n        if f\"fold{i}\" in checkpoint_path.stem:\n            fold = i\n            break\n    print(\"target validation fold is %d\" % fold)\n    if fold == -1:\n        raise ValueError(\"Unexpected fold value\")\n    # Run on the fold that is the target of oof.\n    records_list = predict_for_oof(checkpoint_path.as_posix(), epochs=config.epochs, suffix=config.suffix, folds=[fold])\n    dfs = []\n    for records in records_list:\n        prob = records[\"prob\"].to(\"cpu\").numpy()\n        _df = pd.DataFrame(prob)\n        _df.columns = columns\n        _df[\"seconds\"] = records[\"seconds\"].to(\"cpu\").numpy().tolist()\n        _df[\"filename\"] = list(records[\"filename\"])\n        dfs.append(_df)\n    oof_df = pd.concat(dfs)\n    oof_df = pd.merge(oof_df, metadata_df, how=\"left\", on=[\"filename\"])\n    oof_df = pd.merge(oof_df, nocall_df[[\"filename\", \"nocalldetection\"]], how=\"left\", on=[\"filename\"])\n    oof_df[\"call_prob\"] = oof_df.apply(to_call_prob, axis=1)\n    oof_df[\"birds\"] = oof_df.apply(to_birds, axis=1)\n    filepath = \"%s.csv\" % checkpoint_path.stem\n    print(f\"Save to {filepath}\")\n    oof_df.drop(\n        columns=[\n            'scientific_name',\n            'common_name',\n            'license',\n            'time',\n            'url',\n            'nocalldetection',\n        ]\n    ).to_csv(filepath, index=False)\n    filepath_list.append(filepath)","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:24:13.690955Z","iopub.status.busy":"2021-06-03T07:24:13.68991Z","iopub.status.idle":"2021-06-03T07:36:37.834509Z","shell.execute_reply":"2021-06-03T07:36:37.834019Z"},"id":"B-qNEj90kbxS","outputId":"d9907df7-bc94-4d7e-e77d-5f3d97ef7ce0","papermill":{"duration":744.512346,"end_time":"2021-06-03T07:36:37.834652","exception":false,"start_time":"2021-06-03T07:24:13.322306","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"oof_df.drop(\n    columns=[\n        'scientific_name',\n        'common_name',\n        'license',\n        'time',\n        'url',\n        'nocalldetection',\n    ]\n)","metadata":{"execution":{"iopub.execute_input":"2021-06-03T07:36:38.000836Z","iopub.status.busy":"2021-06-03T07:36:38.000087Z","iopub.status.idle":"2021-06-03T07:36:38.114824Z","shell.execute_reply":"2021-06-03T07:36:38.115322Z"},"id":"KHfXUdHyYwkc","papermill":{"duration":0.247708,"end_time":"2021-06-03T07:36:38.115476","exception":false,"start_time":"2021-06-03T07:36:37.867768","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"PYic6Ti5bc4P","papermill":{"duration":0.032343,"end_time":"2021-06-03T07:36:38.181668","exception":false,"start_time":"2021-06-03T07:36:38.149325","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}