{"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":"! pip install pytorch-lightning -q\n! pip list | grep torch\n! nvidia-smi","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-06-08T01:56:27.866962Z","iopub.execute_input":"2022-06-08T01:56:27.86746Z","iopub.status.idle":"2022-06-08T01:56:42.448423Z","shell.execute_reply.started":"2022-06-08T01:56:27.867348Z","shell.execute_reply":"2022-06-08T01:56:42.447327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! ls /kaggle/input/plant-pathology-2021-fgvc8-960px","metadata":{"execution":{"iopub.status.busy":"2022-06-08T01:56:55.10267Z","iopub.execute_input":"2022-06-08T01:56:55.103395Z","iopub.status.idle":"2022-06-08T01:56:55.778495Z","shell.execute_reply.started":"2022-06-08T01:56:55.103355Z","shell.execute_reply":"2022-06-08T01:56:55.777507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%matplotlib inline\n\nimport os\nimport json\nimport pandas as pd\nfrom pprint import pprint\n\nbase_path = '/kaggle/input/plant-pathology-2021-fgvc8-960px'\npath_csv = os.path.join(base_path, 'train.csv')\ntrain_data = pd.read_csv(path_csv)\nprint(train_data.head())","metadata":{"execution":{"iopub.status.busy":"2022-06-08T01:57:25.664452Z","iopub.execute_input":"2022-06-08T01:57:25.665287Z","iopub.status.idle":"2022-06-08T01:57:25.713642Z","shell.execute_reply.started":"2022-06-08T01:57:25.665244Z","shell.execute_reply":"2022-06-08T01:57:25.712839Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\n\ntrain_data['nb_classes'] = [len(lbs.split(\" \")) for lbs in train_data['labels']]\nlb_hist = dict(zip(range(10), np.bincount(train_data['nb_classes'])))\npprint(lb_hist)","metadata":{"execution":{"iopub.status.busy":"2022-06-08T01:57:27.429518Z","iopub.execute_input":"2022-06-08T01:57:27.430274Z","iopub.status.idle":"2022-06-08T01:57:27.455784Z","shell.execute_reply.started":"2022-06-08T01:57:27.430237Z","shell.execute_reply":"2022-06-08T01:57:27.454966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import itertools\nimport seaborn as sns\n\nlabels_all = list(itertools.chain(*[lbs.split(\" \") for lbs in train_data['labels']]))\n\nax = sns.countplot(y=sorted(labels_all), orient='v')\nax.grid()","metadata":{"execution":{"iopub.status.busy":"2022-06-08T01:57:29.493425Z","iopub.execute_input":"2022-06-08T01:57:29.493795Z","iopub.status.idle":"2022-06-08T01:57:30.301771Z","shell.execute_reply.started":"2022-06-08T01:57:29.493764Z","shell.execute_reply":"2022-06-08T01:57:30.301008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels_unique = set(labels_all)\nprint(f\"unique labels: {labels_unique}\")\ntrain_data['labels_sorted'] = [\" \".join(sorted(lbs.split(\" \"))) for lbs in train_data['labels']]\n\nlabels_combine = {}\nfor comb in train_data['labels_sorted']:\n    labels_combine[comb] = labels_combine.get(comb, 0) + 1\n\nshow_counts = '\\n'.join(sorted(f'\\t{k}: {v}' for k, v in labels_combine.items()))\nprint(f\"unique combinations: \\n\" + show_counts)\nprint(f\"total: {sum(labels_combine.values())}\")","metadata":{"execution":{"iopub.status.busy":"2022-06-08T01:57:32.330013Z","iopub.execute_input":"2022-06-08T01:57:32.330371Z","iopub.status.idle":"2022-06-08T01:57:32.356985Z","shell.execute_reply.started":"2022-06-08T01:57:32.330341Z","shell.execute_reply":"2022-06-08T01:57:32.356075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nnb_samples = 6\nn, m = len(np.unique(train_data['labels_sorted'])), nb_samples,\nfig, axarr = plt.subplots(nrows=n, ncols=m, figsize=(m * 2, n * 2))\nfor ilb, (lb, df_) in enumerate(train_data.groupby('labels_sorted')):\n    img_names = list(df_['image'])\n    for i in range(m):\n        img_name = img_names[i]\n        img = plt.imread(os.path.join(base_path, f\"train_images/{img_name}\"))\n        axarr[ilb, i].imshow(img)\n        if i == 0:\n            axarr[ilb, i].set_title(f\"{lb} #{len(df_)}\")\n        axarr[ilb, i].set_xticks([])\n        axarr[ilb, i].set_yticks([])\nplt.axis('off')","metadata":{"execution":{"iopub.status.busy":"2022-06-08T01:57:33.984722Z","iopub.execute_input":"2022-06-08T01:57:33.985247Z","iopub.status.idle":"2022-06-08T01:57:42.337385Z","shell.execute_reply.started":"2022-06-08T01:57:33.985205Z","shell.execute_reply":"2022-06-08T01:57:42.336585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport torch\nfrom PIL import Image\nfrom torch.utils.data import Dataset\n\nclass PlantPathologyDataset(Dataset):\n    def __init__(\n        self,\n        path_csv: str = os.path.join(base_path, 'train.csv'),\n        path_img_dir: str = os.path.join(base_path, 'train_images'),\n        transforms = None,\n        mode: str = 'train',\n        split: float = 0.8,\n    ):\n        self.path_img_dir = path_img_dir\n        self.transforms = transforms\n        self.mode = mode\n\n        self.data = pd.read_csv(path_csv)\n        labels_all = list(itertools.chain(*[lbs.split(\" \") for lbs in self.data['labels']]))\n        self.labels_unique = 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        # shuffle data\n        self.data = self.data.sample(frac=1, random_state=42).reset_index(drop=True)\n        # split dataset\n        assert 0.0 <= split <= 1.0\n        frac = int(split * len(self.data))\n        self.data = self.data[:frac] if mode == 'train' else self.data[frac:]\n        self.img_names = list(self.data['image']) #이 부분을 교차 검증에 쓰자\n        self.labels = list(self.data['labels'])\n\n    def to_one_hot(self, labels: str) -> tuple:\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_path = os.path.join(self.path_img_dir, self.img_names[idx])\n        assert os.path.isfile(img_path)\n        label = self.labels[idx]\n        img = plt.imread(img_path)\n\n        # augmentation\n        if self.transforms:\n            img = self.transforms(Image.fromarray(img))\n        label = self.to_one_hot(label)\n        return img, torch.tensor(label)\n\n    def __len__(self) -> int:\n        return len(self.data)\n\n# ==============================\n# ==============================\n\ndataset = PlantPathologyDataset()\n\n# quick view\nfig = plt.figure(figsize=(9, 6))\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(lb)","metadata":{"execution":{"iopub.status.busy":"2022-06-08T01:58:15.909492Z","iopub.execute_input":"2022-06-08T01:58:15.910002Z","iopub.status.idle":"2022-06-08T01:58:18.727323Z","shell.execute_reply.started":"2022-06-08T01:58:15.909961Z","shell.execute_reply":"2022-06-08T01:58:18.726536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nfile_list = os.listdir(os.path.join(base_path, 'test_images'))\nfile_list","metadata":{"execution":{"iopub.status.busy":"2022-06-08T01:58:22.807037Z","iopub.execute_input":"2022-06-08T01:58:22.807804Z","iopub.status.idle":"2022-06-08T01:58:22.827878Z","shell.execute_reply.started":"2022-06-08T01:58:22.807758Z","shell.execute_reply":"2022-06-08T01:58:22.822654Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data","metadata":{"execution":{"iopub.status.busy":"2022-06-08T03:31:38.66624Z","iopub.execute_input":"2022-06-08T03:31:38.666804Z","iopub.status.idle":"2022-06-08T03:31:38.684506Z","shell.execute_reply.started":"2022-06-08T03:31:38.666767Z","shell.execute_reply":"2022-06-08T03:31:38.683654Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedKFold as SKF\nskf = SKF(n_splits=7)\nfor train_index, test_index in skf.split(train_data, 7):\n     X_train, X_test = train_data['image'][train_index], train_data['image'][test_index]\n     y_train, y_test = train_data['labels'][train_index], train_data['labels'][test_index]","metadata":{"execution":{"iopub.status.busy":"2022-06-08T02:50:48.321094Z","iopub.execute_input":"2022-06-08T02:50:48.321459Z","iopub.status.idle":"2022-06-08T02:50:48.345197Z","shell.execute_reply.started":"2022-06-08T02:50:48.321428Z","shell.execute_reply":"2022-06-08T02:50:48.343935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport torch\nfrom PIL import Image\nfrom torch.utils.data import Dataset\nfrom sklearn.model_selection import KFold\nkf = KFold(n_splits=10)\nclass PlantPathologyDataset(Dataset):\n    def __init__(\n        self,\n        path_csv: str = os.path.join(base_path, 'train.csv'),\n        path_img_dir: str = os.path.join(base_path, 'train_images'),\n        transforms = None,\n        mode: str = 'train',\n        split: float = 0.8,\n    ):\n        self.path_img_dir = path_img_dir\n        self.transforms = transforms\n        self.mode = mode\n\n        self.data = pd.read_csv(path_csv)\n        labels_all = list(itertools.chain(*[lbs.split(\" \") for lbs in self.data['labels']]))\n        self.labels_unique = 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        # shuffle data\n        self.data = self.data.sample(frac=1, random_state=42).reset_index(drop=True)\n        # split dataset\n        assert 0.0 <= split <= 1.0\n        frac = int(split * len(self.data))\n        self.data = self.data[frac:] if mode == 'test' else self.data[:frac]\n        self.img_names = list(self.data['image'])\n        self.labels = list(self.data['labels'])\n        # k-fold dataset\n        for train_index, test_index in kf.split(train_data):\n            x_train, x_test=train_data['image'][train_index],train_data['image'][test_index]\n            y_train, y_test=train_data['labels'][train_index],train_data['labels'][test_index]\n            \n        self.data = pd.concat([x_test,y_test],axis=1) if mode == 'val' else pd.concat([x_test,y_test],axis=1)\n        self.img_names = list(self.data['image'])\n        self.labels = list(self.data['labels'])\n        #kfold split\n    def to_one_hot(self, labels: str) -> tuple:\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_path = os.path.join(self.path_img_dir, self.img_names[idx])\n        #assert os.path.isfile(img_path)\n        label = self.labels[idx]\n        img = plt.imread(img_path)\n\n        # augmentation\n        if self.transforms:\n            img = self.transforms(Image.fromarray(img))\n        label = self.to_one_hot(label)\n        return img, torch.tensor(label)\n\n    def __len__(self) -> int:\n        return len(self.data)\n\n# ==============================\n# ==============================\n\ndataset = PlantPathologyDataset()\n# quick view\nfig = plt.figure(figsize=(9, 6))\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(lb)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T16:27:17.922941Z","iopub.execute_input":"2022-06-07T16:27:17.923504Z","iopub.status.idle":"2022-06-07T16:27:19.022312Z","shell.execute_reply.started":"2022-06-07T16:27:17.923466Z","shell.execute_reply":"2022-06-07T16:27:19.020344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport torch\nfrom PIL import Image\nfrom torch.utils.data import Dataset\n\nclass PlantPathologyTestDataset(Dataset):\n    def __init__(\n        self,\n        path_csv: str = '../input/test-csv'\n        path_img_dir: str = os.path.join(base_path, 'test_images'),\n        transforms = None,\n        mode: str = 'test',\n    ):\n        self.path_img_dir = path_img_dir\n        self.transforms = transforms\n        self.mode = mode\n\n        self.data = pd.read_csv(path_csv)\n        self.labels=train_data['labels']\n        labels_all = list(itertools.chain(*[lbs.split(\" \") for lbs in self.labels]))\n        self.labels_unique = 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\n    def to_one_hot(self, labels: str) -> tuple:\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_path = os.path.join(self.path_img_dir, file_list[idx])\n        assert os.path.isfile(img_path)\n        label = self.labels[idx]\n        img = plt.imread(img_path)\n\n        # augmentation\n        if self.transforms:\n            img = self.transforms(Image.fromarray(img))\n        label = self.to_one_hot(label)\n        return img, torch.tensor(label)\n\n    def __len__(self) -> int:\n        return len(self.data)\n\n# ==============================\n# ==============================\n\ndataset = PlantPathologyTestDataset()\n\n# quick view\nfig = plt.figure(figsize=(9, 6))\nfor i in range(3):\n    img, lb = dataset[i]\n    ax = fig.add_subplot(3, 3, i + 1, xticks=[], yticks=[])\n    ax.imshow(img)\n    ax.set_title(lb)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T15:36:04.077022Z","iopub.execute_input":"2022-06-07T15:36:04.077373Z","iopub.status.idle":"2022-06-07T15:36:07.439886Z","shell.execute_reply.started":"2022-06-07T15:36:04.077345Z","shell.execute_reply":"2022-06-07T15:36:07.438829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision import transforms as T\n\nTRAIN_TRANSFORM = T.Compose([\n    T.Resize(512),\n    T.RandomPerspective(),\n    T.RandomResizedCrop(224),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.ToTensor(),\n    T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),\n    # T.Normalize([0.431, 0.498,  0.313], [0.237, 0.239, 0.227]),  # custom\n])\n\nVALID_TRANSFORM = T.Compose([\n    T.Resize(256),\n    T.CenterCrop(224),\n    T.ToTensor(),\n    T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),\n    # T.Normalize([0.431, 0.498,  0.313], [0.237, 0.239, 0.227]),  # custom\n])","metadata":{"execution":{"iopub.status.busy":"2022-06-07T15:36:10.697572Z","iopub.execute_input":"2022-06-07T15:36:10.698207Z","iopub.status.idle":"2022-06-07T15:36:10.933416Z","shell.execute_reply.started":"2022-06-07T15:36:10.698139Z","shell.execute_reply":"2022-06-07T15:36:10.932463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import multiprocessing as mproc\nimport pytorch_lightning as pl\nfrom torch.utils.data import DataLoader\n\nclass PlantPathologyDM(pl.LightningDataModule):\n    dataset_cls = PlantPathologyDataset\n\n    def __init__(\n        self,\n        path_csv: str = os.path.join(base_path, 'train.csv'),\n        path_img_dir: str = os.path.join(base_path, 'train_images'),\n        batch_size: int = 128,\n        num_workers: int = None,\n    ):\n        super().__init__()\n        self.path_csv = path_csv\n        self.path_img_dir = path_img_dir\n        self.batch_size = batch_size\n        self.num_workers = num_workers if num_workers is not None else mproc.cpu_count()\n        self.train_dataset = None\n        self.valid_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    def setup(self, stage=None):\n        self.train_dataset = self.dataset_cls(self.path_csv, self.path_img_dir, mode='train', transforms=TRAIN_TRANSFORM)\n        print(f\"training dataset: {len(self.train_dataset)}\")\n        self.valid_dataset = self.dataset_cls(self.path_csv, self.path_img_dir, mode='valid', transforms=VALID_TRANSFORM)\n        print(f\"validation dataset: {len(self.valid_dataset)}\")\n\n    def train_dataloader(self):\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):\n        return DataLoader(\n            self.valid_dataset,\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            shuffle=False,\n        )","metadata":{"execution":{"iopub.status.busy":"2022-06-07T15:36:14.827779Z","iopub.execute_input":"2022-06-07T15:36:14.828379Z","iopub.status.idle":"2022-06-07T15:36:20.848545Z","shell.execute_reply.started":"2022-06-07T15:36:14.828344Z","shell.execute_reply":"2022-06-07T15:36:20.847627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import multiprocessing as mproc\nimport pytorch_lightning as pl\nfrom torch.utils.data import DataLoader\n\nclass PlantPathologyDM(pl.LightningDataModule):\n    dataset_cls = PlantPathologyDataset\n\n    def __init__(\n        self,\n        path_csv: str = os.path.join(base_path, 'train.csv'),\n        path_img_dir: str = os.path.join(base_path, 'train_images'),\n        batch_size: int = 128,\n        num_workers: int = None,\n    ):\n        super().__init__()\n        self.path_csv = path_csv\n        self.path_img_dir = path_img_dir\n        self.batch_size = batch_size\n        self.num_workers = num_workers if num_workers is not None else mproc.cpu_count()\n        self.train_dataset = None\n        self.valid_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    def setup(self, stage=None):\n        self.train_dataset = self.dataset_cls(self.path_csv, self.path_img_dir, mode='train', transforms=TRAIN_TRANSFORM)\n        print(f\"training dataset: {len(self.train_dataset)}\")\n        self.valid_dataset = self.dataset_cls(self.path_csv, self.path_img_dir, mode='valid', transforms=VALID_TRANSFORM)\n        print(f\"validation dataset: {len(self.valid_dataset)}\")\n\n    def train_dataloader(self):\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):\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):\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# ==============================\n# ==============================\n\ndm = PlantPathologyDM()\ndm.setup()\nprint(dm.num_classes)\n\n# quick view\nfig = plt.figure(figsize=(3, 7))\nfor imgs, lbs in dm.train_dataloader():\n    print(f'batch labels: {torch.sum(lbs, axis=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":"2022-06-07T15:36:32.379774Z","iopub.execute_input":"2022-06-07T15:36:32.380896Z","iopub.status.idle":"2022-06-07T15:36:45.963392Z","shell.execute_reply.started":"2022-06-07T15:36:32.380858Z","shell.execute_reply":"2022-06-07T15:36:45.962458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import multiprocessing as mproc\nimport pytorch_lightning as pl\nfrom torch.utils.data import DataLoader\n\nclass PlantPathologyTestDM(pl.LightningDataModule):\n    dataset_cls = PlantPathologyTestDataset\n\n    def __init__(\n        self,\n        path_csv: str = '../input/test-csv/test_csv.csv',\n        path_img_dir: str = os.path.join(base_path, 'test_images'),\n        batch_size: int = 128,\n        num_workers: int = None,\n    ):\n        super().__init__()\n        self.path_csv = path_csv\n        self.path_img_dir = path_img_dir\n        self.batch_size = batch_size\n        self.num_workers = num_workers if num_workers is not None else mproc.cpu_count()\n        self.test_datset=None\n    def prepare_data(self):\n        pass\n\n    @property\n    def num_classes(self) -> int:\n        assert self.test_dataset\n        return self.test_dataset.num_classes\n\n    def setup(self, stage=None):\n        self.test_dataset = self.dataset_cls(self.path_csv, self.path_img_dir, mode='test', transforms=VALID_TRANSFORM)\n        print(f\"test dataset: {len(self.test_dataset)}\")\n\n    def test_dataloader(self):\n        return DataLoader(\n            self.test_dataset,\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            shuffle=False ,\n        )\n\n# ==============================\n# ==============================\n\ndm_test = PlantPathologyTestDM()\ndm_test.setup()\nprint(dm_test.num_classes)\n\n# quick view\nfig = plt.figure(figsize=(3, 7))\nfor imgs, lbs in dm_test.test_dataloader():\n    print(len(imgs))\n    print(lbs)\n    # print(f'batch labels: {torch.sum(lbs, axis=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":"2022-06-07T15:37:14.542047Z","iopub.execute_input":"2022-06-07T15:37:14.545423Z","iopub.status.idle":"2022-06-07T15:37:16.552189Z","shell.execute_reply.started":"2022-06-07T15:37:14.5451Z","shell.execute_reply":"2022-06-07T15:37:16.550844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint('Device:', device)\nprint('Current cuda device:', torch.cuda.current_device())\nprint('Count of using GPUs:', torch.cuda.device_count())","metadata":{"execution":{"iopub.status.busy":"2022-06-07T15:37:22.525191Z","iopub.execute_input":"2022-06-07T15:37:22.525795Z","iopub.status.idle":"2022-06-07T15:37:22.602216Z","shell.execute_reply.started":"2022-06-07T15:37:22.525736Z","shell.execute_reply":"2022-06-07T15:37:22.601442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torchmetrics\nimport torchvision\nfrom torch import nn\nfrom torch.nn import functional as F\n\n\nclass LitResnet(nn.Module):\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\n\nclass LitPlantPathology(pl.LightningModule):\n\n    def __init__(self, model, lr: float = 1e-4):\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        self.val_accuracy = torchmetrics.Accuracy()\n        self.val_f1_score = torchmetrics.F1Score(self.num_classes)\n        self.learn_rate = lr\n        self.loss = nn.BCEWithLogitsLoss()\n\n    def forward(self, x):\n        return F.sigmoid(self.model(x))\n\n    def compute_loss(self, y_hat, y):\n        return self.loss(y_hat, y.to(float))\n\n    def training_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(\"train_loss\", loss, prog_bar=True)\n        self.log(\"train_acc\", self.train_accuracy(y_hat, y), prog_bar=False)\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_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.learn_rate)\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, self.trainer.max_epochs, 0)\n        return [optimizer], [scheduler]\n\n# ==============================\n# ==============================\n\n# see: https://pytorch.org/vision/stable/models.html\n\nnet = LitResnet(arch='resnet50', num_classes=dm.num_classes)\n# print(net)\nmodel = LitPlantPathology(model=net)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T15:37:31.166504Z","iopub.execute_input":"2022-06-07T15:37:31.166907Z","iopub.status.idle":"2022-06-07T15:37:32.606325Z","shell.execute_reply.started":"2022-06-07T15:37:31.166874Z","shell.execute_reply":"2022-06-07T15:37:32.605365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pytorch_lightning.callbacks import ModelCheckpoint\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    # tpu_cores=8,\n    callbacks=[ckpt, swa],\n    logger=logger,\n    max_epochs=10,\n    #precision=16,\n    accumulate_grad_batches=8,\n    val_check_interval=0.25,\n    progress_bar_refresh_rate=1,\n    weights_summary='top',\n)\n\n# ==============================\n\n# trainer.tune(model, datamodule=dm)\ntrainer.fit(model=model, datamodule=dm)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T11:34:17.185754Z","iopub.execute_input":"2022-06-07T11:34:17.186058Z","iopub.status.idle":"2022-06-07T11:35:30.132908Z","shell.execute_reply.started":"2022-06-07T11:34:17.186031Z","shell.execute_reply":"2022-06-07T11:35:30.131426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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', 'train_acc']].plot(grid=True, legend=True, xlabel=agg_col)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T11:35:30.135277Z","iopub.execute_input":"2022-06-07T11:35:30.135713Z","iopub.status.idle":"2022-06-07T11:35:30.651486Z","shell.execute_reply.started":"2022-06-07T11:35:30.135657Z","shell.execute_reply":"2022-06-07T11:35:30.649937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.save_checkpoint(\"example.ckpt\")","metadata":{"execution":{"iopub.status.busy":"2022-06-07T11:35:30.65255Z","iopub.status.idle":"2022-06-07T11:35:30.65338Z","shell.execute_reply.started":"2022-06-07T11:35:30.653144Z","shell.execute_reply":"2022-06-07T11:35:30.65317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls input/","metadata":{"execution":{"iopub.status.busy":"2022-06-07T11:35:30.654781Z","iopub.status.idle":"2022-06-07T11:35:30.655261Z","shell.execute_reply.started":"2022-06-07T11:35:30.655025Z","shell.execute_reply":"2022-06-07T11:35:30.655047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MyLightningModule(LitPlantPathology):\n    def __init__(self, *args, **kwargs):\n        super().__init__()\n        self.save_hyperparameters()\ncheckpoint = torch.load(\"output/example.ckpt\")\n#print(checkpoint['state_dict'])","metadata":{"execution":{"iopub.status.busy":"2022-06-07T11:35:30.656729Z","iopub.status.idle":"2022-06-07T11:35:30.657382Z","shell.execute_reply.started":"2022-06-07T11:35:30.657155Z","shell.execute_reply":"2022-06-07T11:35:30.657179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_model = LitPlantPathology()\nnew_weights = new_model.state_dict()\nold_weights = list(torch.load(\"example.ckpt\")['state_dict'].items())\n\ni=0\nfor k, _ in new_weights.items():\n    new_weights[k] = old_weights[i][1]\n    i += 1\n\nnew_model.load_state_dict(new_weights)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T11:35:30.658314Z","iopub.status.idle":"2022-06-07T11:35:30.659126Z","shell.execute_reply.started":"2022-06-07T11:35:30.658864Z","shell.execute_reply":"2022-06-07T11:35:30.658888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_model.freeze()\n  \n# use it for finetuning\ndef forward(self, x):\n    features = pretrained_model(x)\n    classes = classifier(features)\n  \n# or for prediction\nout = pretrai_model(x)\napi_write({'response': out}","metadata":{"execution":{"iopub.status.busy":"2022-06-07T11:35:30.660279Z","iopub.status.idle":"2022-06-07T11:35:30.661072Z","shell.execute_reply.started":"2022-06-07T11:35:30.660831Z","shell.execute_reply":"2022-06-07T11:35:30.660854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"prediction_list = []\nfor imgs, lbs in dm.train_dataloader():\n    output = new_model(imgs)\n    output = new_model.proba(output) # if not part of forward already\n    prediction_list.append(output)","metadata":{"execution":{"iopub.status.busy":"2022-06-07T11:35:30.662216Z","iopub.status.idle":"2022-06-07T11:35:30.663007Z","shell.execute_reply.started":"2022-06-07T11:35:30.662766Z","shell.execute_reply":"2022-06-07T11:35:30.662789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}