{"cells":[{"metadata":{},"cell_type":"markdown","source":"## Import the libraries"},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"_kg_hide-input":true,"_kg_hide-output":true},"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport os\nimport sys\nfrom typing import Tuple\nimport PIL\nfrom torch.utils.data import Dataset\nfrom pathlib import Path\nfrom PIL import Image\nfrom PIL.Image import Image as PILImage\nfrom torch.utils.data.dataloader import DataLoader\nimport numpy as np\nimport pandas as pd\nfrom pytorch_lightning import LightningDataModule\nfrom sklearn.model_selection import train_test_split, StratifiedKFold\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensor\nfrom albumentations.pytorch import ToTensorV2\n\nfrom torchvision import models\nimport torch.nn as nn\nimport torch\nimport torch.nn.functional as F\nimport pytorch_lightning as pl\nfrom torch import optim\nfrom pytorch_lightning.callbacks import LearningRateMonitor, ModelCheckpoint\n\n\nfrom argparse import ArgumentParser\n\nimport sys\nsys.path.append('../input/timm-pytorch-image-models/pytorch-image-models-master')\nimport timm\n\npath = Path(\"/kaggle/input/cassava-leaf-disease-classification/\")\nfrom tqdm.auto import tqdm\nfrom scipy.special import softmax\n!pip install ../input/geffnet-100/geffnet-1.0.0-py3-none-any.whl\nimport geffnet\nimport cv2","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Data Block\n\n- Create a Pytorch Dataset.\n- Create a Pytorch Lightning Data Module block which contains all the code for creating data loaders."},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaDataset(Dataset):\n    def __init__(self, path, df, transform=None, transform2=None) -> None:\n        super().__init__()\n        self.df = df\n        self.path = path\n        self.transform = transform\n        self.transform2 = transform2 # трансформации для чужих моделей\n        self.num_workers = 2\n\n    def __getitem__(self, index) -> Tuple[PILImage, int]:\n        img_id= self.df.iloc[index,0]\n        image = Image.open(self.path / img_id)\n        image = np.array(image)\n        if self.transform is not None:\n            transformed = self.transform(image=image)\n            image = transformed[\"image\"]\n            \n        file_path = f'{self.path}/{img_id}'\n        image2 = cv2.imread(file_path)\n        image2 = cv2.cvtColor(image2, cv2.COLOR_BGR2RGB)\n        if self.transform:\n            augmented = self.transform2(image=image2)\n            image2 = augmented['image']\n        return image, image2\n\n    def __len__(self):\n        return self.df.shape[0]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaDataModule(LightningDataModule):\n    def __init__(\n        self,\n        path: str = None,\n        aug_p: float = 0.5,\n        val_pct: float = 0.2,\n        img_sz: int = 224,\n        batch_size: int = 64,\n        num_workers: int = 4,\n        fold_id: int = 0,\n    ):\n        super().__init__()\n        self.path = Path(path)\n        self.aug_p = aug_p\n        self.val_pct = val_pct\n        self.img_sz = img_sz\n        self.batch_size = batch_size\n        self.num_workers = num_workers\n        self.fold_id = fold_id\n\n    def prepare_data(self):\n        # only called on 1 GPU/TPU in distributed\n        df = pd.read_csv(self.path / \"train.csv\")\n        skf = StratifiedKFold(n_splits=5)\n        t = df.label\n        train_index, valid_index = list(skf.split(np.zeros(len(t)), t))[self.fold_id]\n        train_df = df.loc[train_index]\n        valid_df = df.loc[valid_index]\n\n        train_df.to_pickle(\"train_df.pkl\")\n        valid_df.to_pickle(\"valid_df.pkl\")\n\n    def setup(self):\n        # called on every process in DDP\n        self.train_transform, self.test_transform = get_augmentations(\n            p=self.aug_p, image_size=self.img_sz\n        )\n        self.train_df = pd.read_pickle(\"train_df.pkl\")\n        self.valid_df = pd.read_pickle(\"valid_df.pkl\")\n\n    def train_dataloader(self):\n        train_dataset = CassavaDataset(\n            self.path / \"train_images\", df=self.train_df, transform=self.train_transform\n        )\n        return DataLoader(\n            train_dataset,\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            shuffle=True,\n            pin_memory=True,\n        )\n\n    def val_dataloader(self):\n        valid_dataset = CassavaDataset(\n            self.path / \"train_images\", df=self.valid_df, transform=self.test_transform\n        )\n        return DataLoader(\n            valid_dataset,\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            shuffle=False,\n            pin_memory=True,\n        )","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Apply Augmentations\n\n"},{"metadata":{},"cell_type":"markdown","source":"## Create a PyTorch Model\n"},{"metadata":{"trusted":true},"cell_type":"code","source":"ssl_models = [\n    \"resnet18\",\n    \"resnet34\",\n    \"resnet50\",\n    \"mobilenetv2\",\n    \"resnext50\",\n    \"resnext101_32x16d_ssl\",\n]\n\nclass Resnext(nn.Module):\n    def __init__(\n        self,\n        model_name=\"resnet18_ssl\",\n        pool_type=F.adaptive_avg_pool2d,\n        num_classes=1000,\n        kaggle=False,\n    ):\n        super().__init__()\n        self.pool_type = pool_type\n\n        if model_name == \"resnet18\":    \n            backbone = models.resnet18(pretrained=False)\n            self.backbone = nn.Sequential(*list(backbone.children())[:-2])\n            in_features = getattr(backbone, \"fc\").in_features\n            self.classifier = nn.Sequential(\n            nn.Dropout(0.2),\n            nn.Linear(in_features, num_classes))\n            \n        if model_name == \"resnet34\":     \n            backbone = models.resnet34(pretrained=False)\n            self.backbone = nn.Sequential(*list(backbone.children())[:-2])\n            in_features = getattr(backbone, \"fc\").in_features\n            self.classifier = nn.Sequential(\n            nn.Dropout(0.2),\n            nn.Linear(in_features, num_classes))\n            \n        if model_name == \"resnet50\":     \n            backbone = models.resnet50(pretrained=False)\n            self.backbone = nn.Sequential(*list(backbone.children())[:-2])\n            in_features = getattr(backbone, \"fc\").in_features\n            self.classifier = nn.Sequential(\n            nn.Dropout(0.2),\n            nn.Linear(in_features, num_classes))\n        \n        if model_name == \"mobilenetv2\":     \n            backbone = models.mobilenet_v2(pretrained=False)\n            self.backbone = nn.Sequential(*list(backbone.children())[:-1])\n            in_features = 1280 \n            self.classifier = nn.Sequential(\n            nn.Dropout(0.2),\n            nn.Linear(in_features, num_classes))\n            \n        if model_name == \"resnext50\": \n            backbone = timm.create_model(\"resnext50_32x4d\", pretrained=False)\n           # semi_supervised_resnext50_32x4.pth - initial-pipeline\n            self.backbone = nn.Sequential(*list(backbone.children())[:-2])\n            in_features = getattr(backbone, \"fc\").in_features\n            self.classifier = nn.Linear(in_features, num_classes)\n            \n        #self.classifier = nn.Linear(in_features, num_classes) - для старых моделей до M7 TT2 (все T -resnet18)\n\n\n    def forward(self, x):\n        features = self.pool_type(self.backbone(x), 1)\n        features = features.view(x.size(0), -1)\n        return self.classifier(features)\n\ndef get_efficientnet(model_name, pretrained=True, num_classes=5):\n    model = geffnet.create_model(model_name, pretrained=False)\n    model.classifier = nn.Linear(model.classifier.in_features, num_classes)\n    return model","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Pytorch Lightning Module\nCreate a Pytorch Lightning Module where we write the essential parts of our training pipeline like\n"},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaModel(pl.LightningModule):\n    def __init__(\n        self,\n        model_name: str = None,\n        num_classes: int = None,\n        data_path: Path = None,\n        loss_fn=F.cross_entropy,\n        lr=1e-4,\n        wd=1e-6,\n    ):\n        super().__init__()\n        \n        if model_name.find(\"effi\") > -1:\n            self.model = get_efficientnet(model_name)\n        else:\n            self.model = Resnext(model_name=model_name, num_classes=num_classes)\n            \n        self.data_path = data_path\n        self.loss_fn = loss_fn\n        self.lr = lr\n        self.accuracy = pl.metrics.Accuracy()\n        self.wd = wd\n\n    def forward(self, x):\n        return self.model(x)\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        y_hat = self(x)\n        loss = self.loss_fn(y_hat, y)\n        self.log(\"train_loss\", loss, 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.loss_fn(y_hat, y)\n        self.log(\"valid_loss\", loss, prog_bar=True)\n        self.log(\"val_acc\", self.accuracy(y_hat, y), prog_bar=True)\n\n    def configure_optimizers(self):\n        optimizer = optim.AdamW(\n            self.model.parameters(), lr=self.lr, weight_decay=self.wd\n        )\n        scheduler = optim.lr_scheduler.CosineAnnealingLR(\n            optimizer, self.trainer.max_epochs, 0\n        )\n\n        return [optimizer], [scheduler]\n","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Чужая часть"},{"metadata":{"trusted":true},"cell_type":"code","source":"# ====================================================\n# Transforms\n# ====================================================\ndef get_transforms(*, data):\n    if data == 'valid':\n        return A.Compose([\n            A.Resize(CFG.size, CFG.size),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])\n\n# CFG\n# ====================================================\nclass CFG:\n    debug=False\n    num_workers=8\n    model_name='resnext50_32x4d'\n    size=512\n    batch_size=32\n    seed=2020\n    target_size=5\n    target_col='label'\n    n_fold=5\n    trn_fold=[0, 1, 2, 3, 4]\n    inference=True\n\n# ====================================================\n# MODEL\n# ====================================================\nclass CustomResNext(nn.Module):\n    def __init__(self, model_name='resnext50_32x4d', pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        n_features = self.model.fc.in_features\n        self.model.fc = nn.Linear(n_features, CFG.target_size)\n\n    def forward(self, x):\n        x = self.model(x)\n        return x\n\n# ====================================================\n# Helper functions\n# ====================================================\ndef load_state(model_path):\n    model = CustomResNext(CFG.model_name, pretrained=False)\n    try:  # single GPU model_file\n        model.load_state_dict(torch.load(model_path)['model'], strict=True)\n        state_dict = torch.load(model_path)['model']\n    except:  # multi GPU model_file\n        state_dict = torch.load(model_path)['model']\n        state_dict = {k[7:] if k.startswith('module.') else k: state_dict[k] for k in state_dict.keys()}\n\n    return state_dict\n\nMODEL_DIR = '../input/cassava-resnext50-32x4d-weights/'\n\nmodel = CustomResNext(CFG.model_name, pretrained=False)\nstates = [load_state(MODEL_DIR+f'{CFG.model_name}_fold{fold}.pth') for fold in CFG.trn_fold]","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**inference part**"},{"metadata":{"trusted":true},"cell_type":"code","source":"class FocalLoss(nn.Module):\n    \"Focal Loss - https://arxiv.org/abs/1708.02002\"\n\n    def __init__(self, alpha=0.25, gamma=2):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n\n    def forward(self, preds, target):\n\n        ce = F.cross_entropy(preds, target, reduction=\"none\")\n        pt = torch.exp(-ce)\n        return ((1.0 - pt) ** self.gamma * ce).mean()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df = pd.DataFrame()\n#df['image_id'] = list(os.listdir('../input/cassava-leaf-disease-classification/train_images/'))\ndf['image_id'] = list(os.listdir('../input/cassava-leaf-disease-classification/test_images/'))\nprint(df)\n#df = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')[0:1000]\n#print(df)\n#ds = CassavaDataset(path/'train_images',df=df)\nds = CassavaDataset(path/'test_images',df=df)\n\nbatch_size, num_workers = 1, 1\n\nloss_fn = {\"cross_entropy\": F.cross_entropy, \"focal_loss\": FocalLoss()}\n\n\nimagenet_stats = {\"mean\": [0.3, 0.3, 0.3], \"std\": [0.1, 0.1, 0.1]}\n\nvalid_tfms = A.Compose(\n    [A.Resize(640, 480),A.CenterCrop(450, 450), ToTensor(normalize=imagenet_stats)]\n)\ntest_ds = CassavaDataset(path=path / \"test_images\", df=df, transform=valid_tfms, transform2=get_transforms(data='valid'))\n#test_ds = CassavaDataset(path=path / \"train_images\", df=df, transform=valid_tfms)\n\ntest_dl = DataLoader(\n    dataset=test_ds,\n    batch_size=batch_size,\n    num_workers=num_workers,\n    shuffle=False,\n    pin_memory=True,\n)\n\ndevice = torch.device(\"cuda\")\n\nmodel4 = CassavaModel(\n     model_name=\"mobilenetv2\",\n     num_classes=5,\n     data_path=path,\n     lr=0.001,\n     loss_fn=loss_fn[\"focal_loss\"]\n )\nM4 = model4.load_from_checkpoint(\"../input/cassamodels/M8.ckpt\",\n      model_name=\"mobilenetv2\",\n      num_classes=5,\n      data_path=path,\n      lr=0.001,\n      loss_fn=loss_fn[\"focal_loss\"])\n\n\nmodel5 = CassavaModel(\n    model_name=\"resnet34\",\n    num_classes=5,\n    data_path=path,\n    lr=0.001,\n    loss_fn=loss_fn[\"focal_loss\"]\n)\nM5 = model5.load_from_checkpoint(\"../input/cassamodels/TT2.ckpt\",\n    model_name=\"resnet34\",\n    num_classes=5,\n    data_path=path,\n    lr=0.001,\n    loss_fn=loss_fn[\"focal_loss\"])\n\nmodel6 = CassavaModel(\n    model_name=\"resnext50\",\n    num_classes=5,\n    data_path=path,\n    lr=0.001,\n    loss_fn=loss_fn[\"focal_loss\"]\n)\nM6 = model6.load_from_checkpoint(\"../input/cassamodels/T1.ckpt\",\n    model_name=\"resnext50\",\n    num_classes=5,\n    data_path=path,\n    lr=0.001,\n    loss_fn=loss_fn[\"focal_loss\"])\n\nmodel7 = CassavaModel(\n     model_name=\"mobilenetv2\",\n     num_classes=5,\n     data_path=path,\n     lr=0.001,\n     loss_fn=loss_fn[\"focal_loss\"]\n )\nM7 = model7.load_from_checkpoint(\"../input/cassamodels/M10.ckpt\",\n      model_name=\"mobilenetv2\",\n      num_classes=5,\n      data_path=path,\n      lr=0.001,\n      loss_fn=loss_fn[\"focal_loss\"])\n\nmodel8 = CassavaModel(\n     model_name=\"tf_efficientnet_b3_ns\",\n     num_classes=5,\n     data_path=path,\n     lr=0.001,\n     loss_fn=loss_fn[\"focal_loss\"]\n )\nM8 = model8.load_from_checkpoint(\"../input/cassamodels/Eff1.ckpt\",\n      model_name=\"tf_efficientnet_b3_ns\",\n      num_classes=5,\n      data_path=path,\n      lr=0.001,\n      loss_fn=loss_fn[\"focal_loss\"])\n\nmodel9 = CassavaModel(\n     model_name=\"tf_efficientnet_b4_ns\",\n     num_classes=5,\n     data_path=path,\n     lr=0.001,\n     loss_fn=loss_fn[\"focal_loss\"]\n )\nM9 = model9.load_from_checkpoint(\"../input/cassamodels/Eff2.ckpt\",\n      model_name=\"tf_efficientnet_b4_ns\",\n      num_classes=5,\n      data_path=path,\n      lr=0.001,\n      loss_fn=loss_fn[\"focal_loss\"])\n\nmodel10 = CassavaModel(\n     model_name=\"tf_efficientnet_b1_ns\",\n     num_classes=5,\n     data_path=path,\n     lr=0.001,\n     loss_fn=loss_fn[\"focal_loss\"]\n )\nM10 = model10.load_from_checkpoint(\"../input/cassamodels/Eff3.ckpt\",\n      model_name=\"tf_efficientnet_b1_ns\",\n      num_classes=5,\n      data_path=path,\n      lr=0.001,\n      loss_fn=loss_fn[\"focal_loss\"])\n\n\nM4.freeze()\nM4 = M4.to(device)\nM4.eval()\nM5.freeze()\nM5 = M5.to(device)\nM5.eval()\nM6.freeze()\nM6 = M6.to(device)\nM6.eval()\nM7.freeze()\nM7 = M7.to(device)\nM7.eval()\nM8.freeze()\nM8 = M8.to(device)\nM8.eval()\nM9.freeze()\nM9 = M9.to(device)\nM9.eval()\nM10.freeze()\nM10 = M10.to(device)\nM10.eval()\n\nmodel.to(device) # чужая модель\n\n\nprint(\"no TTA\")\npreds = []\ni=0\nwith torch.no_grad():\n    for xb,xb1 in test_dl:\n        xb = xb.to(device)\n        xb1 = xb1.to(device)\n\n        pred4 = M4(xb)\n        pred5 = M5(xb)\n        pred6 = M6(xb)\n        pred7 = M7(xb)\n        pred8 = M8(xb)\n        pred9 = M9(xb)\n        pred10 = M10(xb)\n        \n        avg_preds = []\n\n        pred4 = torch.softmax(pred4, 1).detach().cpu()\n        pred5 = torch.softmax(pred5, 1).detach().cpu()\n        pred6 = torch.softmax(pred6, 1).detach().cpu()\n        pred7 = torch.softmax(pred7, 1).detach().cpu()\n        pred8 = torch.softmax(pred8, 1).detach().cpu()\n        pred9 = torch.softmax(pred9, 1).detach().cpu()\n        pred10 = torch.softmax(pred10, 1).detach().cpu()\n\n        \n        \n        #### предикты чужой модели\n        \n        for state in states:\n            model.load_state_dict(state)\n            model.eval()\n            with torch.no_grad():\n                y_preds = model(xb1)\n            avg_preds.append(torch.softmax(y_preds, 1).detach().cpu())\n        #### предикты чужой модели\n        \n        ansmb = 0.33 * (pred4 + pred5 + pred6 + pred7 + pred8 + pred9 + pred10 + avg_preds[0] + avg_preds[1] + avg_preds[2] + avg_preds[3] + avg_preds[4])\n        \n        preds.extend(ansmb.argmax(1).to(\"cpu\").tolist())\n        \n#df[\"preds\"] = preds\ndf[\"label\"] = preds\ndf.to_csv(\"submission.csv\", index=False)\n#print(df)\n\n","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}