{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"!pip -q install ../input/timm-0-1-30/timm-0.1.30-py3-none-any.whl","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport cv2\nimport timm\nfrom matplotlib import pyplot as plt\nfrom sklearn.model_selection import StratifiedKFold\nfrom torch.utils.data import DataLoader, Dataset\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport torchvision.models as models\nimport pytorch_lightning as pl\nfrom pytorch_lightning.core.lightning import LightningModule\nfrom pytorch_lightning.callbacks import ModelCheckpoint\n%matplotlib inline\nfrom pylab import rcParams","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"df_train = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')\ndata_dir = '../input/cassava-leaf-disease-classification/train_images/'","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"DEBUG = True  # set it False for full training","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class config:\n    FOLD_ID = 0\n    IMAGE_SIZE = 256\n    BATCH_SIZE = 32\n    EPOCHS = 10\n    LR = 1e-3\n    NWORKERS = 24\n    SEED = 42\n    NSPLITS = 5\n    NCLASSES = 5\n    T_max = 10\ncfg = config","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"DEVICE = ('cuda' if torch.cuda.is_available() else 'cpu')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"if DEBUG:\n    df_train = df_train[:200]\n    \nskf = StratifiedKFold(cfg.NSPLITS, random_state=cfg.SEED, shuffle=True)\ndf_train['fold'] = -1\nfor fold, (train_idx, valid_idx) in enumerate(skf.split(df_train, df_train.label)):\n    df_train.loc[valid_idx, 'fold'] = fold","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df_train.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"transforms_train = A.Compose([\n    A.RandomRotate90(),\n    A.Flip(),\n    A.Transpose(),\n    \n    A.OneOf([\n            A.CLAHE(clip_limit=2),\n            A.IAASharpen(),\n            A.IAAEmboss(),\n            A.RandomBrightnessContrast(),            \n        ], p=0.3),\n        A.HueSaturationValue(p=0.3),\n    \n    A.Resize(cfg.IMAGE_SIZE, cfg.IMAGE_SIZE),\n    A.Normalize(),\n    ToTensorV2()\n])\n\n\ntransforms_valid = A.Compose([\n    A.Resize(cfg.IMAGE_SIZE, cfg.IMAGE_SIZE),\n    A.Normalize(),\n    ToTensorV2()\n])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class CLDDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        label = row.label\n        image = cv2.imread(data_dir + row.image_id)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        if self.transform is not None:\n            auged = self.transform(image=image)\n            image = auged['image']\n        \n        return image, label","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dataset = CLDDataset(df_train, transforms_train)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"rcParams['figure.figsize'] = 20,10\nfor i in range(2):\n    fig, ax = plt.subplots(1,5)\n    for p in range(5):\n        img, label = dataset[i*5+p]\n        ax[p].imshow(img.permute(1,2,0))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = timm.create_model('tf_efficientnet_b0_ns', pretrained=True, num_classes=cfg.NCLASSES)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class CLDC(LightningModule):\n    def __init__(self, df, model, criterion):\n        super().__init__()\n        self.df = df\n        self.net = model\n        self.criterion = criterion\n\n\n    def forward(self, x):\n        output = self.net(x)\n        return output\n\n    def training_step(self, batch, batch_idx):\n        imgs, labels = batch\n\n        x = self(imgs)\n        loss = self.criterion(x, labels)\n        \n        self.log('train_loss', loss, on_step=True, on_epoch=True, prog_bar=True, logger=True)\n        \n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        imgs, labels = batch\n        x = self(imgs)\n        val_loss = self.criterion(x, labels)\n        \n        self.log('val_loss', val_loss, on_step=True, on_epoch=True, prog_bar=True, logger=True)\n        \n        return val_loss\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=cfg.LR)\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=cfg.T_max)\n        return [optimizer], [scheduler]\n\n    def train_dataloader(self):\n        train_set = CLDDataset(df=self.df[self.df['fold'] != cfg.FOLD_ID], transform=transforms_train)\n        train_loader = DataLoader(train_set, batch_size=cfg.BATCH_SIZE, drop_last=True)\n\n        return train_loader\n\n    def val_dataloader(self):\n        val_set = CLDDataset(df=self.df[self.df['fold'] == cfg.FOLD_ID], transform=transforms_valid)\n        val_loader = DataLoader(val_set, batch_size=cfg.BATCH_SIZE)\n\n        return val_loader","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"cldc = CLDC(df=df_train, model=model, criterion=criterion)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"checkpoint_callback = ModelCheckpoint(\n    save_top_k=1,\n    verbose=True,\n    monitor='val_loss',\n    mode='min',\n    save_weights_only=False\n)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"trainer = pl.Trainer(\n    gpus=-1,\n    max_epochs=cfg.EPOCHS,\n    benchmark=True,\n    amp_level='O1',\n    # auto_lr_find=True,\n    checkpoint_callback=checkpoint_callback\n)\ntrainer.fit(cldc)","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}