{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"!pip install prettyprinter -q","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import os\nimport sys\nimport numpy as np\nimport pandas as pd\nimport random\nfrom pathlib import Path\nfrom typing import Union, List, Callable\nimport logging\nfrom queue import Queue\nfrom tqdm import tqdm\nfrom sklearn.metrics import f1_score\nfrom logging import getLogger, Formatter, FileHandler, StreamHandler, INFO\nfrom pprint import pformat\nfrom prettyprinter import pprint, install_extras\n\nimport torch\nimport torch.nn as nn\nfrom torch import Tensor\nfrom dataclasses import dataclass, field\nfrom torch.utils.data import Dataset, DataLoader\n \nfrom PIL import Image\nimport torchvision\nfrom torchvision import transforms, models\n\nimport matplotlib.pyplot as plt\n%matplotlib inline\n\ninstall_extras(\n    include=[\n        \"dataclasses\",\n    ],\n    warn_on_error=True\n)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"@dataclass\nclass Config:\n    # paths\n    base_dir: Path = Path(\".\").absolute().parent\n    train_imgs_dir: Path = base_dir / \"input/birdsong-log-mel-spectrograms\"\n    train_df: Path = train_imgs_dir / \"train_all.csv\"\n    output_dir = base_dir / \"working\"\n    checkpoint_dir = output_dir / \"checkpoint_dir\"\n    \n    def __post_init__(self):\n        # create the directories if they don't exist\n        self.checkpoint_dir.mkdir(exist_ok=True, parents=True)\n    \n    # training\n    seed: int = 100\n    bs: int = 64\n    num_epochs: int = 5\n    lr: float = 0.003\n    mixed_precision: bool = False\n    opt_level: str = 'O1'\n    step_scheduler: Callable = None\n    device: str = 'cuda'\n    # number of best models for saving\n    num_best_models: int = 3\n        \n    # transforming\n    mean: List[float] = field(default_factory=list)\n    std: List[float] = field(default_factory=list)\n    max_freqmask_width: int = 15\n    max_timemask_width: int = 15\n    timemask_p: float = 0.4\n    freqmask_p: float = 0.4\n        \n        \nconfig = Config()\nconfig.mean += [0.485, 0.456, 0.406]\nconfig.std += [0.229, 0.224, 0.225]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def create_logger():\n    log_file = (config.checkpoint_dir / \"logfile.txt\").as_posix()\n\n    # logger\n    logger_ = getLogger(log_file)\n    logger_.setLevel(INFO)\n\n    # formatter\n    fmr = Formatter(\"[%(levelname)s] %(asctime)s >>\\t%(message)s\")\n\n    # file handler\n    fh = FileHandler(log_file)\n    fh.setLevel(INFO)\n    fh.setFormatter(fmr)\n\n    # stream handler\n    ch = StreamHandler()\n    ch.setLevel(INFO)\n    ch.setFormatter(fmr)\n\n    logger_.addHandler(fh)\n    logger_.addHandler(ch)\n    \n    return logger_\n\nLOGGER = create_logger()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n        \ndef to_numpy(tensor: Union[Tensor, Image.Image, np.array]) -> np.ndarray:\n    if type(tensor) == np.array or type(tensor) == np.ndarray:\n        return np.array(tensor)\n    elif type(tensor) == Image.Image:\n        return np.array(tensor)\n    elif type(tensor) == Tensor:\n        return tensor.cpu().detach().numpy()\n    else:\n        raise ValueError(msg)\n        \ndef seed_everything(seed):\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        torch.backends.cudnn.benchmark = True\n        \nseed_everything(config.seed)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class FrequencyMask(object):\n    def __init__(self, max_width=config.max_freqmask_width, \n                       use_mean=bool(random.randint(0, 1))):\n        \n        self.max_width = max_width\n        self.use_mean = use_mean\n        \n    def __call__(self, tensor_obj):\n        tensor = tensor_obj.detach().clone()\n        start = random.randrange(0, tensor.shape[2])\n        end = start + random.randrange(1, self.max_width)\n       \n        if self.use_mean:\n            tensor[:, start:end, :] = tensor.mean()\n        else:\n            tensor[:, start:end, :] = 0\n        return tensor #[C, H, W]\n    \n    def __repr__(self):\n        format_string = self.__class__.__name__ + \"(max_width=\"\n        format_string += str(self.max_width) + \")\"\n        format_string += \"use_mean=\" + (str(self.use_mean) + \")\")\n        return format_string\n    \n\nclass TimeMask(object):\n    def __init__(self, max_width=config.max_timemask_width, \n                       use_mean=bool(random.randint(0, 1))):\n        \n        self.max_width = max_width\n        self.use_mean = use_mean\n        \n    def __call__(self, tensor_obj):\n        tensor = tensor_obj.detach().clone()\n        start = random.randrange(0, tensor.shape[1])\n        end = start + random.randrange(1, self.max_width)\n        \n        if self.use_mean:\n            tensor[:, :, start:end] = tensor.mean()\n        else:\n            tensor[:, :, start:end] = 0\n        return tensor\n    \n    def __repr__(self):\n        format_string = self.__class__.__name__ + \"(max_width=\"\n        format_string += str(self.max_width) + \")\"\n        format_string += \" use_mean=\" + (str(self.use_mean) + \")\")\n        return format_string","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class BSImageData(Dataset):\n    def __init__(self, data, eval_fold=0, train=True):\n        self.data = data\n        \n        if train:\n            files = self.data[self.data.fold != eval_fold].reset_index(drop=True)\n            # train transforms\n            self.transforms = transforms.Compose([\n                transforms.ToTensor(),\n                transforms.RandomApply([FrequencyMask()], p=config.freqmask_p),\n                transforms.RandomApply([TimeMask()], p=config.timemask_p),\n                transforms.Normalize(\n                    mean=config.mean,\n                    std=config.std\n                )\n            ])\n        else: \n            files = self.data[self.data.fold == eval_fold].reset_index(drop=True)\n            # eval transforms\n            self.transforms = transforms.Compose([\n                transforms.ToTensor(),\n                transforms.Normalize(\n                    mean=config.mean,\n                    std=config.std\n                )\n            ])\n            \n        self.items = files[\"im_path\"].values\n        self.labels = files[\"ebird_code\"].values\n        self.length = len(self.items)\n        \n        \n    def __getitem__(self, index):\n        fname = self.items[index]\n        label = self.labels[index]\n        img = Image.open(fname)\n        img = (np.array(img.convert('RGB')) / 255.).astype(np.float32)\n        return (self.transforms(img), label)\n            \n    def __len__(self):\n        return self.length","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample_transforms = [\n    ( \"Original\", [transforms.ToTensor()]), \n    ( FrequencyMask(), [transforms.ToTensor(), FrequencyMask(use_mean=True)] ),\n    ( TimeMask(), [transforms.ToTensor(), TimeMask(use_mean=True)] ) \n]\n\ndef vis_imgs():\n    plt.figure(figsize=[20,14])\n    for i in range(3):\n        img = random.choice(list((config.train_imgs_dir / \"fold_0\" / \"fold_0\" / \"aldfly\").glob('*')))\n        img = Image.open(img)\n        img = transforms.Compose( sample_transforms[i][1] )(img)\n        img = np.array(transforms.ToPILImage()(img))\n        ax = plt.subplot(1, 3, i+1)\n        ax.set_title(sample_transforms[i][0], fontsize=12)\n        ax.imshow(img)\n        ax.set_xlabel(\"Time\")\n        ax.set_ylabel(\"Hz\")\n        \n    plt.suptitle(\"Mel Spectrograms with different masks\", y=0.78,  fontsize=16)   \n    plt.tight_layout()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"vis_imgs()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def build_model():\n    resnet = models.resnet50(pretrained=True)\n    for param in resnet.parameters():\n        param.requires_grad = False\n\n    resnet.fc = nn.Sequential(nn.Linear(resnet.fc.in_features, 500),\n                         nn.ReLU(),\n                         nn.Dropout(), nn.Linear(500, 264))\n\n    return resnet","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_learning_rate(optimizer):\n    for param_group in optimizer.param_groups:\n        return param_group[\"lr\"]\n\ndef train_fn(data_loader, model, criterion, optimizer, epoch, config, device=config.device):\n    \n    model.train()\n    loss_handler = AverageMeter()\n    score_handler = AverageMeter()\n\n    pbar = tqdm(total=len(data_loader) * config.bs)\n    pbar.set_description(\n        \" Epoch {}, lr: {:.2e}\".format(epoch + 1, get_learning_rate(optimizer))\n    )\n\n    for i, (inputs, target) in enumerate(data_loader):\n\n        optimizer.zero_grad()\n        inputs = inputs.to(device)\n        target = target.to(device)\n        out = model(inputs)\n\n        loss = criterion(out, target)\n\n        if config.mixed_precision:\n            with amp.scale_loss(loss, optimizer) as scaled_loss:\n                scaled_loss.backward()\n        else:\n            loss.backward()\n        \n        optimizer.step()\n\n        if config.step_scheduler:\n            scheduler.step()\n\n        loss_handler.update(loss.item())                         \n        score_handler.update( f1_score(to_numpy(target), to_numpy(torch.argmax(out, dim=1)), average='micro') )\n\n        current_lr = get_learning_rate(optimizer)\n        batch_size = len(inputs)\n        pbar.update(batch_size)\n        pbar.set_postfix(loss=f\"{loss_handler.avg:.5f}\")\n\n    pbar.close()\n    return loss_handler, score_handler.avg","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def eval_fn(data_loader, model, criterion, device=config.device):\n    model.eval()\n    loss_handler = AverageMeter()\n    score_handler = AverageMeter()\n    \n    with torch.no_grad():\n        tk0 = tqdm(data_loader, total=len(data_loader))\n        for inputs, target in tk0:\n\n            inputs = inputs.to(device)\n            target = target.to(device)\n            out = model(inputs)\n            loss = criterion(out, target)\n\n            loss_handler.update(loss.item())\n            score_handler.update( f1_score(to_numpy(target), to_numpy(torch.argmax(out, dim=1)), average=\"micro\") )\n\n            tk0.set_postfix(loss=f\"{loss_handler.avg:.5f}\")\n\n    return loss_handler, score_handler.avg","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def main(config, fold):\n    \n    LOGGER.info(f\"\\n{pformat(config.__dict__)}\\n\")\n    LOGGER.info(\"\\nHERE GOES! 🚀\")\n    LOGGER.info(f\"eval fold: {fold}\")\n    df = pd.read_csv(config.train_df)\n    ebird_dct = {}\n    for i, label in enumerate(df.ebird_code.unique()):\n        ebird_dct[label] = i\n\n    train_data = df.loc[:, [\"im_path\", \"ebird_code\", \"fold\"]]\n    train_data.ebird_code = train_data.ebird_code.map(ebird_dct)\n\n    train_ds = BSImageData(train_data, eval_fold=fold)\n    eval_ds = BSImageData(train_data, eval_fold=fold, train=False)\n\n    train_dl = DataLoader(train_ds, batch_size=config.bs, num_workers=4, shuffle=True)\n    eval_dl = DataLoader(eval_ds, batch_size=config.bs, num_workers=4, shuffle=False)\n\n    device = config.device\n    model = build_model()\n    model = model.to(device)\n    optimizer = torch.optim.Adam(model.parameters(), lr=config.lr)\n    if config.mixed_precision:\n        model, optimizer = amp.initialize(\n                model, optimizer, opt_level=config.opt_level\n        )\n\n    criterion = nn.CrossEntropyLoss()\n    model_queue = Queue()\n    for epoch in range(config.num_epochs):\n        train_loss, train_score = train_fn(\n            train_dl, model, criterion, optimizer, config=config, epoch=epoch\n        )\n\n        valid_loss, valid_score = eval_fn(eval_dl, model, criterion)\n\n        LOGGER.info(\n            f\"|EPOCH {epoch+1}| F1_train {train_score:.5f}| F1_valid {valid_score:.5f}|\"\n        )\n        \n        best_loss = float(\"inf\")\n        if valid_loss.avg < best_loss:\n            best_loss = valid_loss.avg\n            print(f\"New best model in epoch {epoch+1}\")\n            mname = f\"{epoch+1}_best_model_{best_loss:.5f}.pth\"\n            torch.save(model.state_dict(), config.checkpoint_dir / mname)\n            \n            model_queue.put(mname)\n            if model_queue.qsize() > config.num_best_models:\n                mname_to_del = model_queue.get()\n                (config.checkpoint_dir / mname_to_del).unlink()\n                LOGGER.info(f\"{mname_to_del} deleted\")\n                \n    LOGGER.info(f\"saved models: {model_queue.queue}\")\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"main(config, fold=0)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!cat ../working/checkpoint_dir/logfile.txt","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}