{"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":"# iMet Collection 2021 with Lightning ⚡","metadata":{}},{"cell_type":"code","source":"! pip install -q pytorch-lightning\n! pip list | grep torch\n! nvidia-smi","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-06-09T13:49:14.909389Z","iopub.execute_input":"2021-06-09T13:49:14.909868Z","iopub.status.idle":"2021-06-09T13:49:26.993839Z","shell.execute_reply.started":"2021-06-09T13:49:14.909784Z","shell.execute_reply":"2021-06-09T13:49:26.992236Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data view\n\nChecking what data do we have available...","metadata":{}},{"cell_type":"code","source":"# jsu to see what is the data location\n! ls /kaggle/input -l\n! ls /kaggle/input/imet-2021-fgvc8 -l\n\nimport pandas as pd\n\nPATH_DATASET = \"/kaggle/input/imet-2021-fgvc8/\"\npd.read_csv(PATH_DATASET + \"train-from-kaggle.csv\").head()","metadata":{"execution":{"iopub.status.busy":"2021-06-09T13:49:26.999793Z","iopub.execute_input":"2021-06-09T13:49:27.000991Z","iopub.status.idle":"2021-06-09T13:49:28.960586Z","shell.execute_reply.started":"2021-06-09T13:49:27.000937Z","shell.execute_reply":"2021-06-09T13:49:28.959617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset & DataModule\n\nCreating standard PyTorch dataset to define how the data shall be loaded and set representations. We define the sample pair as:\n- RGB image\n- one-hot lable encding\n\nA DataModule standardizes the training, val, test splits, data preparation and transforms. The main advantage is consistent data splits, data preparation and transforms across models.","metadata":{}},{"cell_type":"code","source":"import glob\nimport itertools\nimport logging\nimport multiprocessing as mproc\nimport os\nfrom typing import Dict, List, Optional, Sequence, Tuple, Type, Union\n\nimport tqdm\nimport numpy as np\nimport torch\nimport matplotlib.pyplot as plt\nfrom PIL import Image, ImageFile\nfrom torch import Tensor\nfrom torch.utils.data import DataLoader, Dataset\nfrom joblib import Parallel, delayed\n\n\nImageFile.LOAD_TRUNCATED_IMAGES = True\n\n\ndef get_nb_pixels(img_path: str):\n    try:\n        img = Image.open(img_path)\n        return np.prod(img.size)\n    except Exception:\n        return 0\n\n\nclass IMetDataset(Dataset):\n    \"\"\"The ful dataset with one-hot encoding for multi-label case.\"\"\"\n    IMAGE_SIZE_LIMIT = 1000\n    COL_LABELS = 'attribute_ids'\n    COL_IMAGES = 'id'\n\n    def __init__(\n        self,\n        df_data: Union[str, pd.DataFrame] = 'train-from-kaggle.csv',\n        path_img_dir: str = 'train-1/train-1',\n        transforms=None,\n        mode: str = 'train',\n        split: float = 0.8,\n        uq_labels: Tuple[str] = None,\n        random_state=42,\n        check_imgs: bool = True,\n    ):\n        self.path_img_dir = path_img_dir\n        self.transforms = transforms\n        self.mode = mode\n        self._img_names = None\n        self._raw_labels = None\n\n        # set or load the config table\n        if isinstance(df_data, pd.DataFrame):\n            self.data = df_data\n        elif isinstance(df_data, str):\n            assert os.path.isfile(df_data), f\"missing file: {df_data}\"\n            self.data = pd.read_csv(df_data)\n        else:\n            raise ValueError(f'unrecognised input for DataFrame/CSV: {df_data}')\n\n        # take over existing table or load from file\n        if uq_labels:\n            self.labels_unique = tuple(uq_labels)\n        else:\n            labels_all = list(itertools.chain(*[lbs.split(\" \") for lbs in self.raw_labels]))\n            self.labels_unique = tuple(sorted(set(labels_all)))\n        self.labels_lut = {lb: i for i, lb in enumerate(self.labels_unique)}\n        self.num_classes = len(self.labels_unique)\n\n        # filter/drop too small images\n        if check_imgs:\n            with Parallel(n_jobs=mproc.cpu_count()) as parallel:\n                self.data['pixels'] = parallel(delayed(get_nb_pixels)(os.path.join(self.path_img_dir, im)) for im in self.img_names)\n            nb_small_imgs = sum(self.data['pixels'] < self.IMAGE_SIZE_LIMIT)\n            if nb_small_imgs:\n                logging.warning(f\"found and dropped {nb_small_imgs} too small or invalid images :/\")\n            self.data = self.data[self.data['pixels'] >= self.IMAGE_SIZE_LIMIT]\n        # shuffle data\n        self.data = self.data.sample(frac=1, random_state=random_state).reset_index(drop=True)\n\n        # split dataset\n        assert 0.0 <= split <= 1.0, f\"split {split} is out of range\"\n        frac = int(split * len(self.data))\n        self.data = self.data[:frac] if mode == 'train' else self.data[frac:]\n        # need to reset after another split since it cached\n        self._img_names = None\n        self._raw_labels = None\n        self.labels = self._prepare_labels()\n\n    @property\n    def img_names(self):\n        if not self._img_names:\n            self._img_names = [f\"{n}.png\" if '.' not in n else n for n in self.data[self.COL_IMAGES]]\n        return self._img_names\n\n    @property\n    def raw_labels(self):\n        if not self._raw_labels:\n            self._raw_labels = list(self.data[self.COL_LABELS])\n        return self._raw_labels\n\n    def _prepare_labels(self) -> list:\n        return [torch.tensor(self.to_onehot_encoding(lb)) if lb else None for lb in self.raw_labels]\n\n    def to_onehot_encoding(self, labels: str) -> tuple:\n        # processed with encoding\n        one_hot = [0] * len(self.labels_unique)\n        for lb in labels.split(\" \"):\n            one_hot[self.labels_lut[lb]] = 1\n        return tuple(one_hot)\n\n    def __getitem__(self, idx: int) -> tuple:\n        img_name = self.img_names[idx]\n        img_path = os.path.join(self.path_img_dir, img_name)\n        assert os.path.isfile(img_path)\n        label = self.labels[idx]\n        # todo: find some faster way, do conversion only if needed; im.mode not in (\"L\", \"RGB\")\n        img = Image.open(img_path).convert('RGB')\n\n        # augmentation\n        if self.transforms:\n            img = self.transforms(img)\n                \n        # in case of predictions, return image name as label\n        label = label if label is not None else img_name\n        return img, label\n\n    def __len__(self) -> int:\n        return len(self.data)\n\n\ndataset = IMetDataset(\n    df_data=PATH_DATASET + \"train-from-kaggle.csv\",\n    path_img_dir=PATH_DATASET + \"train-1/train-1\",\n)\n\n# quick view\nfig = plt.figure(figsize=(12, 8))\nfor i in range(9):\n    img, lb = dataset[i]\n    ax = fig.add_subplot(3, 3, i + 1, xticks=[], yticks=[])\n    ax.imshow(img)\n    ax.set_title(f\"img: {img.size}\\n lb: {lb}\")","metadata":{"execution":{"iopub.status.busy":"2021-06-09T13:49:28.963588Z","iopub.execute_input":"2021-06-09T13:49:28.964091Z","iopub.status.idle":"2021-06-09T13:56:47.809486Z","shell.execute_reply.started":"2021-06-09T13:49:28.964043Z","shell.execute_reply":"2021-06-09T13:56:47.808354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pytorch_lightning import LightningDataModule\n\n\nclass IMetDM(LightningDataModule):\n\n    IMAGE_EXTENSIONS = ('.png', '.jpg', '.jpeg')\n\n    def __init__(\n        self,\n        base_path: str,\n        path_csv: str = 'train-from-kaggle.csv',\n        batch_size: int = 128,\n        num_workers: int = None,\n        train_transforms=None,\n        valid_transforms=None,\n        split: float = 0.8,\n    ):\n        super().__init__()\n        # path configurations\n        assert os.path.isdir(base_path), f\"missing folder: {base_path}\"\n        self.train_dir = os.path.join(base_path, 'train-1/train-1')\n        self.test_dir = os.path.join(base_path, 'test/test')\n\n        if not os.path.isfile(path_csv):\n            path_csv = os.path.join(base_path, path_csv)\n        assert os.path.isfile(path_csv), f\"missing table: {path_csv}\"\n        self.path_csv = path_csv\n\n        self.train_transforms = train_transforms\n        self.valid_transforms = valid_transforms\n\n        # other configs\n        self.batch_size = batch_size\n        self.split = split\n        self.num_workers = num_workers if num_workers is not None else mproc.cpu_count()\n        self.labels_unique: Sequence = ...\n        self.lut_label: Dict = ...\n        self.label_histogram: Tensor = ...\n\n        # need to be filled in setup()\n        self.train_dataset = None\n        self.valid_dataset = None\n        self.test_table = []\n        self.test_dataset = None\n\n    def prepare_data(self):\n        pass\n\n    @property\n    def num_classes(self) -> int:\n        assert self.train_dataset and self.valid_dataset\n        return max(self.train_dataset.num_classes, self.valid_dataset.num_classes)\n\n    @staticmethod\n    def onehot_mapping(\n        onehot: Tensor,\n        lut_label: Dict[int, str],\n        thr: float = 0.5,\n        label_required: bool = True,\n    ) -> Union[str, List[str]]:\n        \"\"\"Convert Model outputs to string labels\"\"\"\n        assert lut_label\n        # on case it is not one hot encoding but single label\n        if onehot.nelement() == 1:\n            return lut_label[onehot[0]]\n        labels = [lut_label[i] for i, s in enumerate(onehot) if s >= thr]\n        # in case no reached threshold then take max\n        if not labels and label_required:\n            idx = torch.argmax(onehot).item()\n            labels = [lut_label[idx]]\n        return sorted(labels)\n\n    def onehot_to_labels(self, onehot: Tensor, thr: float = 0.5, label_required: bool = True) -> Union[str, List[str]]:\n        \"\"\"Convert Model outputs to string labels\"\"\"\n        return self.onehot_mapping(onehot, self.lut_label, thr=thr, label_required=label_required)\n\n    def setup(self, *_, **__) -> None:\n        \"\"\"Prepare datasets\"\"\"\n        pbar = tqdm.tqdm(total=4)\n        assert os.path.isdir(self.train_dir), f\"missing folder: {self.train_dir}\"\n        ds = IMetDataset(self.path_csv, self.train_dir, mode='train', split=1.0)\n        self.labels_unique = ds.labels_unique\n        self.lut_label = dict(enumerate(self.labels_unique))\n        pbar.update()\n        \n        ds_defaults = dict(\n            df_data=ds.data,\n            path_img_dir=self.train_dir,\n            split=self.split,\n            uq_labels=self.labels_unique,\n            check_imgs=False,\n        )\n        self.train_dataset = IMetDataset(**ds_defaults, mode='train', transforms=self.train_transforms)\n        logging.info(f\"training dataset: {len(self.train_dataset)}\")\n        pbar.update()\n        self.valid_dataset = IMetDataset(**ds_defaults, mode='valid', transforms=self.valid_transforms)\n        logging.info(f\"validation dataset: {len(self.valid_dataset)}\")\n        pbar.update()\n\n        if not os.path.isdir(self.test_dir):\n            return\n        ls_images = glob.glob(os.path.join(self.test_dir, '*.*'))\n        ls_images = [os.path.basename(p) for p in ls_images if os.path.splitext(p)[-1] in self.IMAGE_EXTENSIONS]\n        self.test_table = [{'id': n, 'attribute_ids': ''} for n in ls_images]\n        self.test_dataset = IMetDataset(\n            df_data=pd.DataFrame(self.test_table),\n            path_img_dir=self.test_dir,\n            split=0,\n            uq_labels=self.labels_unique,\n            mode='test',\n            transforms=self.valid_transforms\n        )\n        logging.info(f\"test dataset: {len(self.test_dataset)}\")\n        pbar.update()\n\n    def train_dataloader(self) -> DataLoader:\n        return DataLoader(\n            self.train_dataset,\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            shuffle=True,\n        )\n\n    def val_dataloader(self) -> DataLoader:\n        return DataLoader(\n            self.valid_dataset,\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            shuffle=False,\n        )\n\n    def test_dataloader(self) -> Optional[DataLoader]:\n        logging.warning('no testing images found')","metadata":{"execution":{"iopub.status.busy":"2021-06-09T13:56:47.811478Z","iopub.execute_input":"2021-06-09T13:56:47.811873Z","iopub.status.idle":"2021-06-09T13:56:49.882464Z","shell.execute_reply.started":"2021-06-09T13:56:47.811833Z","shell.execute_reply":"2021-06-09T13:56:49.881221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\nfrom torchvision import transforms as T\n\n#: default training augmentation\nTORCHVISION_TRAIN_TRANSFORM = T.Compose([\n    T.Resize(size=512, interpolation=Image.BILINEAR),\n    T.RandomRotation(degrees=30),\n    T.RandomPerspective(distortion_scale=0.2),\n    T.RandomResizedCrop(size=224),\n    T.RandomHorizontalFlip(p=0.5),\n    # T.RandomVerticalFlip(p=0.5),\n    # T.ColorJitter(brightness=0.05, contrast=0.05, saturation=0.05, hue=0.05),\n    T.ToTensor(),\n    T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),\n    #T.Normalize(DATASET_IMAGE_MEAN, DATASET_IMAGE_STD),  # custom\n])\n#: default validation augmentation\nTORCHVISION_VALID_TRANSFORM = T.Compose([\n    T.Resize(size=256, interpolation=Image.BILINEAR),\n    T.CenterCrop(size=224),\n    T.ToTensor(),\n    T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),\n    #T.Normalize(DATASET_IMAGE_MEAN, DATASET_IMAGE_STD),  # custom\n])\n\ndm = IMetDM(\n    base_path=PATH_DATASET,\n    batch_size=128,\n    train_transforms=TORCHVISION_TRAIN_TRANSFORM,\n    valid_transforms=TORCHVISION_VALID_TRANSFORM,\n    num_workers=0,\n)\ndm.setup()\n\n# Quick view\nfig = plt.figure(figsize=(3, 7))\nfor imgs, lbs in dm.val_dataloader():\n    batch_lb_sum = torch.sum(lbs, axis=0).numpy()\n    print(f'batch labels: {list(batch_lb_sum[batch_lb_sum > 0])}')\n    print(f'image size: {imgs[0].shape}')\n    for i in range(3):\n        ax = fig.add_subplot(3, 1, i + 1, xticks=[], yticks=[])\n        # print(np.rollaxis(imgs[i].numpy(), 0, 3).shape)\n        ax.imshow(np.rollaxis(imgs[i].numpy(), 0, 3))\n        ax.set_title(lbs[i])\n    break","metadata":{"execution":{"iopub.status.busy":"2021-06-09T13:56:49.884457Z","iopub.execute_input":"2021-06-09T13:56:49.884895Z","iopub.status.idle":"2021-06-09T14:08:58.549766Z","shell.execute_reply.started":"2021-06-09T13:56:49.884843Z","shell.execute_reply":"2021-06-09T14:08:58.548431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CNN Model\n\nWe start with some stanrd CNN models taken from torch vision. Then we define Ligthning module including training and validation step and configure optimizer/schedular.","metadata":{}},{"cell_type":"code","source":"from typing import Optional\n\nimport torch\nimport torchmetrics\nimport torchvision\nfrom pytorch_lightning import LightningModule\nfrom torch import nn, Tensor\nfrom torch.nn import functional as F\n\n\nclass LitResnet(nn.Module):\n\n    def __init__(self, arch: str, pretrained: bool = True, num_classes: int = 6):\n        super().__init__()\n        self.arch = arch\n        self.num_classes = num_classes\n        self.model = torchvision.models.__dict__[arch](pretrained=pretrained)\n        num_features = self.model.fc.in_features\n        self.model.fc = nn.Linear(num_features, num_classes)\n\n    def forward(self, x):\n        return self.model(x)\n\nclass LitMet(LightningModule):\n    \"\"\"\n    This model is meant and tested to be used together with ...\n    \"\"\"\n\n    def __init__(self, model, lr: float = 1e-4, augmentations: Optional[nn.Module] = None):\n        super().__init__()\n        self.model = model\n        self.arch = self.model.arch\n        self.num_classes = self.model.num_classes\n        self.train_accuracy = torchmetrics.Accuracy()\n        _metrics_extra_args = dict(num_classes=self.num_classes, multilabel=True, average='weighted')\n        self.train_precision = torchmetrics.Precision(**_metrics_extra_args)\n        self.train_f1_score = torchmetrics.F1(**_metrics_extra_args)\n        self.val_accuracy = torchmetrics.Accuracy()\n        self.val_precision = torchmetrics.Precision(**_metrics_extra_args)\n        self.val_f1_score = torchmetrics.F1(**_metrics_extra_args)\n        self.learning_rate = lr\n        self.aug = augmentations\n\n    def forward(self, x: Tensor) -> Tensor:\n        return torch.sigmoid(self.model(x))\n\n    def compute_loss(self, y_hat: Tensor, y: Tensor):\n        return F.binary_cross_entropy_with_logits(y_hat, y.to(y_hat.dtype))\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        if self.aug:\n            x = self.aug(x)  # => batched augmentations\n        y_hat = self(x)\n        loss = self.compute_loss(y_hat, y)\n        self.log(\"train_loss\", loss, prog_bar=False)\n        self.log(\"train_acc\", self.train_accuracy(y_hat, y), prog_bar=False)\n        self.log(\"train_prec\", self.train_precision(y_hat, y), prog_bar=False)\n        self.log(\"train_f1\", self.train_f1_score(y_hat, y), prog_bar=True)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        x, y = batch\n        y_hat = self(x)\n        loss = self.compute_loss(y_hat, y)\n        self.log(\"valid_loss\", loss, prog_bar=False)\n        self.log(\"valid_acc\", self.val_accuracy(y_hat, y), prog_bar=True)\n        self.log(\"valid_prec\", self.val_precision(y_hat, y), prog_bar=True)\n        self.log(\"valid_f1\", self.val_f1_score(y_hat, y), prog_bar=True)\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(self.model.parameters(), lr=self.learning_rate)\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, self.trainer.max_epochs, 0)\n        return [optimizer], [scheduler]\n\n\n# see: https://pytorch.org/vision/stable/models.html\nnet = LitResnet(arch='resnet50', num_classes=dm.num_classes)\n# print(net)\n\nmodel = LitMet(model=net, lr=5e-4)","metadata":{"execution":{"iopub.status.busy":"2021-06-09T14:08:58.551938Z","iopub.execute_input":"2021-06-09T14:08:58.552488Z","iopub.status.idle":"2021-06-09T14:09:00.82478Z","shell.execute_reply.started":"2021-06-09T14:08:58.552434Z","shell.execute_reply":"2021-06-09T14:09:00.823511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training\n\nWe use Pytorch Lightning which allow us to drop all the boilet plate code and simplify all training just to use/call Trainer...","metadata":{}},{"cell_type":"code","source":"import pytorch_lightning as pl\nprint(pl.__version__)\n\nlogger = pl.loggers.CSVLogger(save_dir='logs/', name=model.arch)\nswa = pl.callbacks.StochasticWeightAveraging(swa_epoch_start=0.6)\nckpt = pl.callbacks.ModelCheckpoint(\n    monitor='valid_f1',\n    save_top_k=1,\n    save_last=True,\n    # save_weights_only=True,\n    filename='checkpoint/{epoch:02d}-{valid_acc:.4f}-{valid_f1:.4f}',\n    # verbose=False,\n    mode='max',\n)\n\n# ==============================\n\ntrainer = pl.Trainer(\n    # fast_dev_run=True,\n    gpus=1,\n    callbacks=[ckpt],\n    logger=logger,\n    max_epochs=1,\n    precision=16,\n    #overfit_batches=5,\n    auto_lr_find=True,\n    accumulate_grad_batches=24,\n    val_check_interval=0.5,\n    progress_bar_refresh_rate=1,\n    weights_summary='top',\n)\n\n# ==============================\n\n# lr_find_kwargs = dict(min_lr=1e-5, max_lr=1e-2, num_training=25)\n# trainer.tune(model, datamodule=dm, lr_find_kwargs=lr_find_kwargs)\n# print(f\"LR: {model.learning_rate}\")\n\n# ==============================\n\ndm.batch_size = 128\ntrainer.fit(model=model, datamodule=dm)","metadata":{"execution":{"iopub.status.busy":"2021-06-09T14:28:52.781672Z","iopub.execute_input":"2021-06-09T14:28:52.782131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Quick visualization of the training process...","metadata":{}},{"cell_type":"code","source":"metrics = pd.read_csv(f'{trainer.logger.log_dir}/metrics.csv')\nprint(metrics.head())\n\naggreg_metrics = []\nagg_col = \"epoch\"\nfor i, dfg in metrics.groupby(agg_col):\n    agg = dict(dfg.mean())\n    agg[agg_col] = i\n    aggreg_metrics.append(agg)\n\ndf_metrics = pd.DataFrame(aggreg_metrics)\ndf_metrics[['train_loss', 'valid_loss']].plot(grid=True, legend=True, xlabel=agg_col)\ndf_metrics[['valid_f1', 'valid_acc', 'valid_prec', 'train_acc']].plot(grid=True, legend=True, xlabel=agg_col)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}