{"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\n","metadata":{"id":"C0vy_lvxVQf5","outputId":"0ad3c95b-49a6-4e37-c750-8b7f79e76f57","papermill":{"duration":1.406395,"end_time":"2021-04-23T05:33:11.704364","exception":false,"start_time":"2021-04-23T05:33:10.297969","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.03272,"end_time":"2021-04-23T05:33:11.769577","exception":false,"start_time":"2021-04-23T05:33:11.736857","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Notes","metadata":{"papermill":{"duration":0.031848,"end_time":"2021-04-23T05:33:11.832687","exception":false,"start_time":"2021-04-23T05:33:11.800839","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"Days back, I've shared this [infernece kernel](https://www.kaggle.com/kneroma/clean-fast-simple-bird-identifier-inference). But its weights are static as you can't retrain the model. In this work, I'm gonna release the training notebook which is almost my internal training pipeline. I removed some experimentation ideas to make things clearer and straightforward. Don't mind adding new ideas at your side as well.","metadata":{"id":"py6GyjyZVark","papermill":{"duration":0.03109,"end_time":"2021-04-23T05:33:11.895227","exception":false,"start_time":"2021-04-23T05:33:11.864137","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"To make the training faster, we cached the training set into RAM. The whole training records are already [converted into handy  melspecs images](https://www.kaggle.com/kneroma/kkiller-birdclef-2021). These images are from 7 seconds extracts (training on 7 seconds seems to be more effective than 5 seconds). Longer records are truncated into random 7x10 seconds.\n\n**If one is interessted in to the whole records' melspecs** (no truncation):\n* https://www.kaggle.com/kneroma/kkiller-birdclef-mels-computer-d7-part1\n* https://www.kaggle.com/kneroma/kkiller-birdclef-mels-computer-d7-part2\n* https://www.kaggle.com/kneroma/kkiller-birdclef-mels-computer-d7-part3\n* https://www.kaggle.com/kneroma/kkiller-birdclef-mels-computer-d7-part4","metadata":{"papermill":{"duration":0.031735,"end_time":"2021-04-23T05:33:11.958291","exception":false,"start_time":"2021-04-23T05:33:11.926556","status":"completed"},"tags":[]}},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.031274,"end_time":"2021-04-23T05:33:12.020964","exception":false,"start_time":"2021-04-23T05:33:11.98969","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Tips & suggestions\n* You can choose a wide set of models from the **get_model** interface : [\"resnest*\", \"resnet*\", \"resnext*\", \"efficientnet*\" ...]\n* You can change the learning rate scheduler: OneCycle ? ReduceOnPlateau ?\n* Adds secondary labels\n* Use train & test metadata (dates, positions (longitude, latitude), ...)\n* Add melspecs augmentation","metadata":{"papermill":{"duration":0.03122,"end_time":"2021-04-23T05:33:12.083498","exception":false,"start_time":"2021-04-23T05:33:12.052278","status":"completed"},"tags":[]}},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.029867,"end_time":"2021-04-23T05:33:12.144214","exception":false,"start_time":"2021-04-23T05:33:12.114347","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**For Colab training, you just have to uncomment the first cells**","metadata":{"papermill":{"duration":0.030018,"end_time":"2021-04-23T05:33:12.204319","exception":false,"start_time":"2021-04-23T05:33:12.174301","status":"completed"},"tags":[]}},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.030844,"end_time":"2021-04-23T05:33:12.266509","exception":false,"start_time":"2021-04-23T05:33:12.235665","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Versions","metadata":{"papermill":{"duration":0.030804,"end_time":"2021-04-23T05:33:12.3284","exception":false,"start_time":"2021-04-23T05:33:12.297596","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"* **v1** : initial version\n* **v3** : enable training on whole (no truncation) record melspecs","metadata":{"papermill":{"duration":0.044248,"end_time":"2021-04-23T05:33:12.404103","exception":false,"start_time":"2021-04-23T05:33:12.359855","status":"completed"},"tags":[]}},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.05023,"end_time":"2021-04-23T05:33:12.502426","exception":false,"start_time":"2021-04-23T05:33:12.452196","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from google.colab import drive\n# drive.mount('/content/drive')","metadata":{"id":"oYPb42V-Vaza","outputId":"f9846f5e-bb8b-407c-d15e-bab4b8604301","papermill":{"duration":0.056894,"end_time":"2021-04-23T05:33:12.608506","exception":false,"start_time":"2021-04-23T05:33:12.551612","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ! pip install --upgrade --force-reinstall --no-deps  kaggle > /dev/null\n# ! mkdir ~/.kaggle\n# ! cp \"/content/drive/My Drive/Kaggle/kaggle.json\" ~/.kaggle/\n# ! chmod 600 ~/.kaggle/kaggle.json","metadata":{"id":"IRn4fOj5XDp6","papermill":{"duration":0.065168,"end_time":"2021-04-23T05:33:12.733668","exception":false,"start_time":"2021-04-23T05:33:12.6685","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%time\n\n# import os\n# if not os.path.exists(\"/content/datasets/audio_images\"):\n#   !mkdir datasets\n#   !kaggle datasets download -d kneroma/kkiller-birdclef-2021\n#   !unzip /content//kkiller-birdclef-2021.zip -d datasets","metadata":{"id":"KAnAPecOXDtf","outputId":"fb981400-b8c5-45bc-9e7d-5d4a77f2550e","papermill":{"duration":0.066101,"end_time":"2021-04-23T05:33:12.858052","exception":false,"start_time":"2021-04-23T05:33:12.791951","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q pysndfx SoundFile audiomentations pretrainedmodels efficientnet_pytorch resnest","metadata":{"id":"Yn1Ybf15VAqW","papermill":{"duration":13.004387,"end_time":"2021-04-23T05:33:25.919056","exception":false,"start_time":"2021-04-23T05:33:12.914669","status":"completed"},"tags":[],"trusted":true},"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\n\nfrom IPython.display import Audio\nfrom sklearn.metrics import label_ranking_average_precision_score\n\nfrom tqdm.notebook import tqdm\nimport joblib","metadata":{"id":"2dt7oG43VAqc","papermill":{"duration":3.200761,"end_time":"2021-04-23T05:33:29.151263","exception":false,"start_time":"2021-04-23T05:33:25.950502","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from efficientnet_pytorch import EfficientNet\nimport pretrainedmodels\nimport resnest.torch as resnest_torch\n","metadata":{"id":"162Vl9uxe1Mj","papermill":{"duration":1.298377,"end_time":"2021-04-23T05:33:30.48129","exception":false,"start_time":"2021-04-23T05:33:29.182913","status":"completed"},"tags":[],"trusted":true},"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\nseed_everything()","metadata":{"id":"Q39ZsGAhVAqe","papermill":{"duration":0.040696,"end_time":"2021-04-23T05:33:30.553222","exception":false,"start_time":"2021-04-23T05:33:30.512526","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_CLASSES = 397\nSR = 32_000\nDURATION = 10\n\nMAX_READ_SAMPLES = 5 # Each record will have 10 melspecs at most, you can increase this on Colab with High Memory Enabled\n\n# # For colab\n# DATA_ROOT = Path(\"/content/datasets/\")\n# TRAIN_IMAGES_ROOT = Path(\"/content/datasets/audio_images\")\n# TRAIN_LABELS_FILE = Path(\"/content/datasets/rich_train_metadata.csv\")\n# MODEL_ROOT = Path(\"/content/drive/My Drive/Kaggle/BirdClef2021/models\")\n\nDATA_ROOT = Path(\"../input/birdclef-2021\")\n# TRAIN_IMAGES_ROOT = Path(\"../input/kkiller-birdclef-2021/audio_images\")\n# TRAIN_LABELS_FILE = Path(\"../input/kkiller-birdclef-2021/rich_train_metadata.csv\")\n\nMEL_PATHS = sorted(Path(\"../input\").glob(\"kkiller-birdclef-mels-computer-d7-part?/rich_train_metadata.csv\"))\nTRAIN_LABEL_PATHS = sorted(Path(\"../input\").glob(\"kkiller-birdclef-mels-computer-d7-part?/LABEL_IDS.json\"))\n\nMODEL_ROOT = Path(\".\")","metadata":{"id":"2NfkUn9SCWs6","papermill":{"duration":0.094659,"end_time":"2021-04-23T05:33:30.678927","exception":false,"start_time":"2021-04-23T05:33:30.584268","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.030981,"end_time":"2021-04-23T05:33:30.741227","exception":false,"start_time":"2021-04-23T05:33:30.710246","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_BATCH_SIZE = 100\nTRAIN_NUM_WORKERS = 2\n\nVAL_BATCH_SIZE = 128\nVAL_NUM_WORKERS = 2\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nprint(\"Device:\", DEVICE)","metadata":{"id":"Iu56f-7VVAqf","outputId":"0f3fa344-0ed4-47d8-f3c0-218cf3bf5a78","papermill":{"duration":0.552535,"end_time":"2021-04-23T05:33:31.326345","exception":false,"start_time":"2021-04-23T05:33:30.77381","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"GuyjwJnACWs6","outputId":"cdaca87a-567c-4299-9840-7b3cb06be13f","papermill":{"duration":0.031859,"end_time":"2021-04-23T05:33:31.390297","exception":false,"start_time":"2021-04-23T05:33:31.358438","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":{"papermill":{"duration":0.040479,"end_time":"2021-04-23T05:33:31.462342","exception":false,"start_time":"2021-04-23T05:33:31.421863","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# df = pd.read_csv(TRAIN_LABELS_FILE, nrows=None)\n# df[\"secondary_labels\"] = df[\"secondary_labels\"].apply(literal_eval)\n# LABEL_IDS = {label: label_id for label_id,label in enumerate(sorted(df[\"primary_label\"].unique()))}\n\n# print(df.shape)\n# df.head()","metadata":{"id":"Kmh6xx5_NCjJ","outputId":"3e92e880-2e37-46c7-ffc0-5551e2641e7b","papermill":{"duration":0.037875,"end_time":"2021-04-23T05:33:31.531497","exception":false,"start_time":"2021-04-23T05:33:31.493622","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"LABEL_IDS, df = get_df()\n\nprint(df.shape)\ndf.head()","metadata":{"papermill":{"duration":3.719029,"end_time":"2021-04-23T05:33:35.281949","exception":false,"start_time":"2021-04-23T05:33:31.56292","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"W_Xf_natBGGL","papermill":{"duration":0.032534,"end_time":"2021-04-23T05:33:35.348458","exception":false,"start_time":"2021-04-23T05:33:35.315924","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[\"primary_label\"].value_counts()","metadata":{"id":"ZRz-DwbNVAqg","outputId":"42aa9ba1-7fc3-42b6-e818-c4a04277fc51","papermill":{"duration":0.054526,"end_time":"2021-04-23T05:33:35.435815","exception":false,"start_time":"2021-04-23T05:33:35.381289","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[\"label_id\"].min(), df[\"label_id\"].max()","metadata":{"id":"n68WeAa0VAqh","outputId":"31adfe31-5bfb-473c-e79d-272d5a8b8cf3","papermill":{"duration":0.041905,"end_time":"2021-04-23T05:33:35.510621","exception":false,"start_time":"2021-04-23T05:33:35.468716","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"xh4hfWZhuglm","papermill":{"duration":0.032978,"end_time":"2021-04-23T05:33:35.577001","exception":false,"start_time":"2021-04-23T05:33:35.544023","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    \n    \n    if \"resnest\" in name:\n        model = getattr(resnest_torch, name)(pretrained=False)\n    elif \"wsl\" in name:\n        model = torch.hub.load(\"facebookresearch/WSL-Images\", 'resnext101_32x8d_wsl')\n    elif name.startswith(\"resnext\") or  name.startswith(\"resnet\"):\n        model = torch.hub.load(\"pytorch/vision:v0.6.0\", 'resnext50_32x4d', pretrained=True)\n    elif name.startswith(\"tf_efficientnet_b\"):\n        model = getattr(timm.models.efficientnet, 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":{"id":"OGPDuihmVAqi","papermill":{"duration":0.044611,"end_time":"2021-04-23T05:33:35.654862","exception":false,"start_time":"2021-04-23T05:33:35.610251","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_data(df):\n    def load_row(row):\n        # impath = TRAIN_IMAGES_ROOT/f\"{row.primary_label}/{row.filename}.npy\"\n        return row.filename, np.load(str(row.impath))[:MAX_READ_SAMPLES]\n    pool = joblib.Parallel(4)\n    mapper = joblib.delayed(load_row)\n    tasks = [mapper(row) for row in df.itertuples(False)]\n    res = pool(tqdm(tasks))\n    res = dict(res)\n    return res","metadata":{"id":"7HYQwAyBCWs8","papermill":{"duration":0.042033,"end_time":"2021-04-23T05:33:35.730156","exception":false,"start_time":"2021-04-23T05:33:35.688123","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# We cache the train set to reduce training time\n\naudio_image_store = load_data(df)\nlen(audio_image_store)","metadata":{"id":"Vw19bB7mCWs9","outputId":"09a5e374-7e5c-4c92-91e4-b60313cbb1a9","papermill":{"duration":278.253556,"end_time":"2021-04-23T05:38:14.017297","exception":false,"start_time":"2021-04-23T05:33:35.763741","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"shape:\", next(iter(audio_image_store.values())).shape)\nlbd.specshow(next(iter(audio_image_store.values()))[0])","metadata":{"id":"4TNYmT7XCWs9","outputId":"b5052897-ca12-4c25-a49c-3799a88b9f9d","papermill":{"duration":0.192339,"end_time":"2021-04-23T05:38:14.244535","exception":false,"start_time":"2021-04-23T05:38:14.052196","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"_tSt9iC7CWs9","papermill":{"duration":0.037627,"end_time":"2021-04-23T05:38:14.320237","exception":false,"start_time":"2021-04-23T05:38:14.28261","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.Series([len(x) for x in audio_image_store.values()]).value_counts()","metadata":{"id":"bUUqP5KMBZkc","outputId":"ab2ca6c4-8aed-4b02-ed52-485b7af70ae1","papermill":{"duration":0.102172,"end_time":"2021-04-23T05:38:14.459906","exception":false,"start_time":"2021-04-23T05:38:14.357734","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"rOTGe4dbBapl","papermill":{"duration":0.038494,"end_time":"2021-04-23T05:38:14.536372","exception":false,"start_time":"2021-04-23T05:38:14.497878","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BirdClefDataset(Dataset):\n\n    def __init__(self, audio_image_store, meta, sr=SR, is_train=True, num_classes=NUM_CLASSES, duration=DURATION):\n        \n        self.audio_image_store = audio_image_store\n        self.meta = meta.copy().reset_index(drop=True)\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    \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.meta)\n    \n    def __getitem__(self, idx):\n        row = self.meta.iloc[idx]\n        image = self.audio_image_store[row.filename]\n\n        image = image[np.random.choice(len(image))]\n        image = self.normalize(image)\n        \n        \n        t = np.zeros(self.num_classes, dtype=np.float32) + 0.0025 # Label smoothing\n        t[row.label_id] = 0.995\n        \n        return image, t","metadata":{"id":"OWSkCXyhCWs-","papermill":{"duration":0.048677,"end_time":"2021-04-23T05:38:14.62304","exception":false,"start_time":"2021-04-23T05:38:14.574363","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = BirdClefDataset(audio_image_store, meta=df, sr=SR, duration=DURATION, is_train=True)\nlen(df)","metadata":{"id":"Np-56XrXVAqm","outputId":"84e45885-200f-4c98-a9a1-23ef459e5ade","papermill":{"duration":0.072969,"end_time":"2021-04-23T05:38:14.734195","exception":false,"start_time":"2021-04-23T05:38:14.661226","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x, y = ds[np.random.choice(len(ds))]\n# x, y = ds[0]\nx.shape, y.shape, np.where(y >= 0.5)","metadata":{"id":"UNVUxIMpVAqm","outputId":"96e0930f-51f7-4268-9939-2d81999c002f","papermill":{"duration":0.060457,"end_time":"2021-04-23T05:38:14.833469","exception":false,"start_time":"2021-04-23T05:38:14.773012","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lbd.specshow(x[0])","metadata":{"id":"bPwrCoRyCWs-","outputId":"388a8eaf-f373-4e05-c57d-be4f6975c890","papermill":{"duration":0.125247,"end_time":"2021-04-23T05:38:14.997466","exception":false,"start_time":"2021-04-23T05:38:14.872219","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y[:5]","metadata":{"id":"GBr32Q9FCWs_","outputId":"141ba3fd-99b1-4878-8462-2a6eb3cb5e19","papermill":{"duration":0.049544,"end_time":"2021-04-23T05:38:15.089058","exception":false,"start_time":"2021-04-23T05:38:15.039514","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.04181,"end_time":"2021-04-23T05:38:15.173526","exception":false,"start_time":"2021-04-23T05:38:15.131716","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training the model","metadata":{"id":"F56zXq8CVAqn","papermill":{"duration":0.041892,"end_time":"2021-04-23T05:38:15.257645","exception":false,"start_time":"2021-04-23T05:38:15.215753","status":"completed"},"tags":[]}},{"cell_type":"code","source":"","metadata":{"id":"THm438BwMTeR","papermill":{"duration":0.042031,"end_time":"2021-04-23T05:38:15.341741","exception":false,"start_time":"2021-04-23T05:38:15.29971","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def one_step( xb, yb, net, criterion, optimizer, scheduler=None):\n  xb, yb = xb.to(DEVICE), yb.to(DEVICE)\n        \n  optimizer.zero_grad()\n  o = net(xb)\n  loss = criterion(o, yb)\n  loss.backward()\n  optimizer.step()\n  \n  with torch.no_grad():\n      l = loss.item()\n\n      o = o.sigmoid()\n      yb = (yb >= 0.5 )*1.0\n      lrap = label_ranking_average_precision_score(yb.cpu().numpy(), o.cpu().numpy())\n\n      o = (o >= 0.5)*1.0\n\n      prec = (o*yb).sum()/(1e-6 + o.sum())\n      rec = (o*yb).sum()/(1e-6 + yb.sum())\n      f1 = 2*prec*rec/(1e-6+prec+rec)\n\n  if  scheduler is not None:\n    scheduler.step()\n\n  return l, lrap, f1.item(), rec.item(), prec.item()","metadata":{"id":"9Kjy1uquIGZw","papermill":{"duration":0.051879,"end_time":"2021-04-23T05:38:15.435881","exception":false,"start_time":"2021-04-23T05:38:15.384002","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef evaluate(net, criterion, val_laoder):\n    net.eval()\n\n    os, y = [], []\n    val_laoder = tqdm(val_laoder, leave = False, total=len(val_laoder))\n\n    for icount, (xb, yb) in  enumerate(val_laoder):\n\n        y.append(yb.to(DEVICE))\n\n        xb = xb.to(DEVICE)\n        o = net(xb)\n\n        os.append(o)\n\n    y = torch.cat(y)\n    o = torch.cat(os)\n\n    l = criterion(o, y).item()\n    \n    o = o.sigmoid()\n    y = (y > 0.5)*1.0\n\n    lrap = label_ranking_average_precision_score(y.cpu().numpy(), o.cpu().numpy())\n\n    o = (o > 0.5)*1.0\n\n    prec = ((o*y).sum()/(1e-6 + o.sum())).item()\n    rec = ((o*y).sum()/(1e-6 + y.sum())).item()\n    f1 = 2*prec*rec/(1e-6+prec+rec)\n\n    return l, lrap, f1, rec, prec, \n    ","metadata":{"id":"q9v79J0pvXy1","papermill":{"duration":0.052542,"end_time":"2021-04-23T05:38:15.530748","exception":false,"start_time":"2021-04-23T05:38:15.478206","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def one_epoch(net, criterion, optimizer, scheduler, train_laoder, val_laoder):\n  net.train()\n  l, lrap, prec, rec, f1, icount = 0.,0.,0.,0., 0., 0\n  train_laoder = tqdm(train_laoder, leave = False)\n  epoch_bar = train_laoder\n  \n  for (xb, yb) in  epoch_bar:\n      # epoch_bar.set_description(\"----|----|----|----|---->\")\n      _l, _lrap, _f1, _rec, _prec = one_step(xb, yb, net, criterion, optimizer)\n      l += _l\n      lrap += _lrap\n      f1 += _f1\n      rec += _rec\n      prec += _prec\n\n      icount += 1\n        \n      if hasattr(epoch_bar, \"set_postfix\") and not icount%10:\n          epoch_bar.set_postfix(\n            loss=\"{:.6f}\".format(l/icount),\n            lrap=\"{:.3f}\".format(lrap/icount),\n            prec=\"{:.3f}\".format(prec/icount),\n            rec=\"{:.3f}\".format(rec/icount),\n            f1=\"{:.3f}\".format(f1/icount),\n          )\n  \n  scheduler.step()\n\n  l /= icount\n  lrap /= icount\n  f1 /= icount\n  rec /= icount\n  prec /= icount\n  \n  l_val, lrap_val, f1_val, rec_val, prec_val = evaluate(net, criterion, val_laoder)\n  \n  return (l, l_val), (lrap, lrap_val), (f1, f1_val), (rec, rec_val), (prec, prec_val)","metadata":{"id":"qeDgf4LdLWGN","papermill":{"duration":0.053436,"end_time":"2021-04-23T05:38:15.626513","exception":false,"start_time":"2021-04-23T05:38:15.573077","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"IY4ET5V0RMJm","papermill":{"duration":0.042567,"end_time":"2021-04-23T05:38:15.711514","exception":false,"start_time":"2021-04-23T05:38:15.668947","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AutoSave:\n  def __init__(self, top_k=2, metric=\"f1\", mode=\"min\", root=None, name=\"ckpt\"):\n    self.top_k = top_k\n    self.logs = []\n    self.metric = metric\n    self.mode = mode\n    self.root = Path(root or MODEL_ROOT)\n    assert self.root.exists()\n    self.name = name\n\n    self.top_models = []\n    self.top_metrics = []\n\n  def log(self, model, metrics):\n    metric = metrics[self.metric]\n    rank = self.rank(metric)\n\n    self.top_metrics.insert(rank+1, metric)\n    if len(self.top_metrics) > self.top_k:\n      self.top_metrics.pop(0)\n\n    self.logs.append(metrics)\n    self.save(model, metric, rank, metrics[\"epoch\"])\n\n\n  def save(self, model, metric, rank, epoch):\n    t = time.strftime(\"%Y%m%d%H%M%S\")\n    name = \"{}_epoch_{:02d}_{}_{:.04f}_{}\".format(self.name, epoch, self.metric, metric, t)\n    name = re.sub(r\"[^\\w_-]\", \"\", name) + \".pth\"\n    path = self.root.joinpath(name)\n\n    old_model = None\n    self.top_models.insert(rank+1, name)\n    if len(self.top_models) > self.top_k:\n      old_model = self.root.joinpath(self.top_models[0])\n      self.top_models.pop(0)      \n\n    torch.save(model.state_dict(), path.as_posix())\n\n    if old_model is not None:\n      old_model.unlink()\n\n    self.to_json()\n\n\n  def rank(self, val):\n    r = -1\n    for top_val in self.top_metrics:\n      if val <= top_val:\n        return r\n      r += 1\n\n    return r\n  \n  def to_json(self):\n    # t = time.strftime(\"%Y%m%d%H%M%S\")\n    name = \"{}_logs\".format(self.name)\n    name = re.sub(r\"[^\\w_-]\", \"\", name) + \".json\"\n    path = self.root.joinpath(name)\n\n    with path.open(\"w\") as f:\n      json.dump(self.logs, f, indent=2)\n","metadata":{"id":"Cz7XPBvtPLO1","papermill":{"duration":0.057159,"end_time":"2021-04-23T05:38:15.811332","exception":false,"start_time":"2021-04-23T05:38:15.754173","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def one_fold(model_name, fold, train_set, val_set, epochs=10, save=True, save_root=None):\n\n  save_root = Path(save_root) or MODEL_ROOT\n\n  saver = AutoSave(root=save_root, name=f\"birdclef_{model_name}_fold{fold}\", metric=\"f1_val\")\n\n  net = get_model(model_name).to(DEVICE)\n\n  criterion = nn.BCEWithLogitsLoss()\n\n  optimizer = optim.Adam(net.parameters(), lr=8e-4)\n  scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, eta_min=1e-5, T_max=epochs)\n\n  train_data = BirdClefDataset(audio_image_store, meta=df.iloc[train_set].reset_index(drop=True),\n                           sr=SR, duration=DURATION, is_train=True)\n  train_laoder = DataLoader(train_data, batch_size=TRAIN_BATCH_SIZE, num_workers=TRAIN_NUM_WORKERS, shuffle=True, pin_memory=True)\n\n  val_data = BirdClefDataset(audio_image_store, meta=df.iloc[val_set].reset_index(drop=True),  sr=SR, duration=DURATION, is_train=False)\n  val_laoder = DataLoader(val_data, batch_size=VAL_BATCH_SIZE, num_workers=VAL_NUM_WORKERS, shuffle=False)\n\n  epochs_bar = tqdm(list(range(epochs)), leave=False)\n  for epoch  in epochs_bar:\n    epochs_bar.set_description(f\"--> [EPOCH {epoch:02d}]\")\n    net.train()\n\n    (l, l_val), (lrap, lrap_val), (f1, f1_val), (rec, rec_val), (prec, prec_val) = one_epoch(\n        net=net,\n        criterion=criterion,\n        optimizer=optimizer,\n        scheduler=scheduler,\n        train_laoder=train_laoder,\n        val_laoder=val_laoder,\n      )\n\n    epochs_bar.set_postfix(\n    loss=\"({:.6f}, {:.6f})\".format(l, l_val),\n    prec=\"({:.3f}, {:.3f})\".format(prec, prec_val),\n    rec=\"({:.3f}, {:.3f})\".format(rec, rec_val),\n    f1=\"({:.3f}, {:.3f})\".format(f1, f1_val),\n    lrap=\"({:.3f}, {:.3f})\".format(lrap, lrap_val),\n    )\n\n    print(\n        \"[{epoch:02d}] loss: {loss} lrap: {lrap} f1: {f1} rec: {rec} prec: {prec}\".format(\n            epoch=epoch,\n            loss=\"({:.6f}, {:.6f})\".format(l, l_val),\n            prec=\"({:.3f}, {:.3f})\".format(prec, prec_val),\n            rec=\"({:.3f}, {:.3f})\".format(rec, rec_val),\n            f1=\"({:.3f}, {:.3f})\".format(f1, f1_val),\n            lrap=\"({:.3f}, {:.3f})\".format(lrap, lrap_val),\n        )\n    )\n\n    if save:\n      metrics = {\n          \"loss\": l, \"lrap\": lrap, \"f1\": f1, \"rec\": rec, \"prec\": prec,\n          \"loss_val\": l_val, \"lrap_val\": lrap_val, \"f1_val\": f1_val, \"rec_val\": rec_val, \"prec_val\": prec_val,\n          \"epoch\": epoch,\n      }\n\n      saver.log(net, metrics)","metadata":{"id":"8X1dt_aWNi6F","papermill":{"duration":0.059133,"end_time":"2021-04-23T05:38:15.91288","exception":false,"start_time":"2021-04-23T05:38:15.853747","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"NiEiYTjaSCaH","papermill":{"duration":0.0427,"end_time":"2021-04-23T05:38:15.998726","exception":false,"start_time":"2021-04-23T05:38:15.956026","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model_name, epochs=30, save=True, n_splits=10, seed=177, save_root=None, suffix=\"\", folds=None):\n  gc.collect()\n  torch.cuda.empty_cache()\n\n  save_root = save_root or MODEL_ROOT/f\"{model_name}{suffix}\"\n  save_root.mkdir(exist_ok=True, parents=True)\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        \n      one_fold(model_name, fold=fold, train_set=train_set , val_set=val_set , epochs=epochs, save=save, save_root=save_root)\n    \n      gc.collect()\n      torch.cuda.empty_cache()","metadata":{"id":"ljqr4e2zQmzB","papermill":{"duration":0.052887,"end_time":"2021-04-23T05:38:16.094779","exception":false,"start_time":"2021-04-23T05:38:16.041892","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODEL_NAMES = [\n      \"resnest50\",\n] ","metadata":{"id":"aqN6xL7gVAqq","papermill":{"duration":0.048296,"end_time":"2021-04-23T05:38:16.185676","exception":false,"start_time":"2021-04-23T05:38:16.13738","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for model_name in MODEL_NAMES:\n  print(\"\\n\\n###########################################\", model_name.upper())\n  train(model_name, epochs=30, suffix=f\"_sr{SR}_d{DURATION}_v1_v1\", folds=[0])\n","metadata":{"id":"WyFnAQGWELb_","outputId":"d676c647-65e4-43d1-a8f3-69f3d949e727","papermill":{"duration":1266.003605,"end_time":"2021-04-23T05:59:22.231994","exception":false,"start_time":"2021-04-23T05:38:16.228389","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"Q46u71ImEL4E","papermill":{"duration":0.047745,"end_time":"2021-04-23T05:59:22.327344","exception":false,"start_time":"2021-04-23T05:59:22.279599","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}