{"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":"# About this notebook\n\n## Version 1\n- PyTorch Efficientnet_b0 starter code\n- 5 folds\n- 30 epochs\n- batch size 64 no accumulation\n- Custom learning scheduler\n- With data augmentation\n\n## Version 2\n- PyTorch Efficientnet_b0 starter code\n- 5 folds\n- 30 epochs\n- batch size 64 no accumulation\n- Custom learning scheduler\n- With data augmentation\n\n## Version 3\n- PyTorch Efficientnet_b0 starter code\n- 5 folds\n- 30 epochs\n- batch size 64 no accumulation\n- Custom learning scheduler\n- With data augmentation\n- CrossEntropyLoss\n- meta features 'Age' 'variety'\n\n# Improvements maybe\n- Use ArcFace or add triplet loss with cross entropy for accuracy improvement\n- Use meta featues 'Age' 'variety' to improove the performance of the models\n- Use focal Loss (already implemented)\n\n# acknowledgement\n- Y.NAKAMA great [notebook](https://www.kaggle.com/yasufuminakama/herbarium-2020-pytorch-resnet18-train/notebook)\n\nIf this notebook is helpful, feel free to upvote :)","metadata":{}},{"cell_type":"code","source":"!pip install -q --upgrade wandb\n!pip install timm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-25T23:58:57.798535Z","iopub.execute_input":"2022-07-25T23:58:57.799313Z","iopub.status.idle":"2022-07-25T23:59:20.491569Z","shell.execute_reply.started":"2022-07-25T23:58:57.799216Z","shell.execute_reply":"2022-07-25T23:59:20.490286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport cv2 as cv\nfrom matplotlib import pyplot as plt\nimport seaborn as sns\n\npd.options.display.max_columns = 300","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:20.493930Z","iopub.execute_input":"2022-07-25T23:59:20.494576Z","iopub.status.idle":"2022-07-25T23:59:21.217634Z","shell.execute_reply.started":"2022-07-25T23:59:20.494516Z","shell.execute_reply":"2022-07-25T23:59:21.216654Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Quick EDA","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv('../input/paddy-disease-classification/train.csv')\n#test = pd.read_csv('../input/melanomacsv/test_clean.csv')\nsubmission = pd.read_csv('../input/paddy-disease-classification/sample_submission.csv')\ntrain_dir = '../input/paddy-disease-classification/train_images/'\n\n#train['path_jpeg'] = train['label'].apply(lambda x: os.path.join('../input/paddy-disease-classification/train_images',f'{x}'))\n#test['path_jpeg'] = test['dcm_name'].apply(lambda x: os.path.join('../input/jpeg-melanoma-256x256/test',f'{x}.jpg'))","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:21.218891Z","iopub.execute_input":"2022-07-25T23:59:21.220113Z","iopub.status.idle":"2022-07-25T23:59:21.257597Z","shell.execute_reply.started":"2022-07-25T23:59:21.220070Z","shell.execute_reply":"2022-07-25T23:59:21.256661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['path_jpeg'] = train_df.apply(lambda row: train_dir + row['label'] + '/' + row['image_id'], axis=1)","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:21.260118Z","iopub.execute_input":"2022-07-25T23:59:21.261632Z","iopub.status.idle":"2022-07-25T23:59:21.438039Z","shell.execute_reply.started":"2022-07-25T23:59:21.261593Z","shell.execute_reply":"2022-07-25T23:59:21.437092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i in range(5):\n    image = cv.imread(train_df.loc[i, 'path_jpeg'])\n    image = cv.cvtColor(image, cv.COLOR_BGR2RGB)\n    label = train_df.loc[i, 'label']\n    plt.imshow(image)\n    plt.title(f\"{label}\")\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:21.439619Z","iopub.execute_input":"2022-07-25T23:59:21.439982Z","iopub.status.idle":"2022-07-25T23:59:22.789317Z","shell.execute_reply.started":"2022-07-25T23:59:21.439946Z","shell.execute_reply":"2022-07-25T23:59:22.788308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn import preprocessing\n\nle = preprocessing.LabelEncoder()\nle.fit(train_df['label'])\ntrain_df['label'] = le.transform(train_df['label'])","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:22.792432Z","iopub.execute_input":"2022-07-25T23:59:22.792714Z","iopub.status.idle":"2022-07-25T23:59:22.861964Z","shell.execute_reply.started":"2022-07-25T23:59:22.792684Z","shell.execute_reply":"2022-07-25T23:59:22.861034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"le.fit(train_df['variety'])\ntrain_df['variety'] = le.transform(train_df['variety'])","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:22.863586Z","iopub.execute_input":"2022-07-25T23:59:22.863926Z","iopub.status.idle":"2022-07-25T23:59:22.873114Z","shell.execute_reply.started":"2022-07-25T23:59:22.863892Z","shell.execute_reply":"2022-07-25T23:59:22.872019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Directory settings","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:22.874606Z","iopub.execute_input":"2022-07-25T23:59:22.875343Z","iopub.status.idle":"2022-07-25T23:59:22.882084Z","shell.execute_reply.started":"2022-07-25T23:59:22.875306Z","shell.execute_reply":"2022-07-25T23:59:22.881196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"class CFG:\n    apex=False\n    debug=False\n    print_freq=100\n    size=256\n    num_workers=2\n    scheduler='CosineAnnealingLR' # ['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts','OneCycleLR']\n    epochs=30\n    # CosineAnnealingLR params\n    cosanneal_params={\n        'T_max':10,\n        'eta_min':1e-4*0.5,\n        'last_epoch':-1\n    }\n    #ReduceLROnPlateau params\n    reduce_params={\n        'mode':'min',\n        'factor':0.1,\n        'patience':6,\n        'eps':1e-6,\n        'verbose':True\n    }\n    #ReduceLROnPlateau params for auc\n    reduce_params_for_auc={\n        'mode':'max',\n        'factor':0.1,\n        'patience':3,\n        'eps':1e-6,\n        'verbose':True\n    }\n    # CosineAnnealingWarmRestarts params\n    cosanneal_res_params={\n        'T_0':3,\n        'eta_min':1e-6,\n        'T_mult':1,\n        'last_epoch':-1\n    }\n    # OneCycleLR params\n    onecycle_params={\n        'pct_start':0.1,\n        'div_factor':1e1,\n        'max_lr':1e-3,\n        'steps_per_epoch':3, \n        'epochs':3\n    }\n    #batch_size=64\n    momentum=0.9\n    lr=1e-3\n    weight_decay=1e-4\n    gradient_accumulation_steps=1\n    max_grad_norm=1000\n    nfolds=5\n    trn_folds=[0, 1, 2, 3, 4]\n    model_name='efficientnet_b0'     #'vit_base_patch32_224_in21k' 'tf_efficientnetv2_b0' 'resnext50_32x4d' 'resnet50d' 'efficientnet_b0'\n    preds_col = ['bacterial_leaf_blight', 'bacterial_leaf_streak',\n               'bacterial_panicle_blight', 'blast', 'brown_spot', 'dead_heart',\n               'downy_mildew', 'hispa', 'normal', 'tungro']\n    train=True\n    early_stop=True\n    target_size=len(preds_col)\n    scale=30.0\n    margin=0.50\n    easy_margin=False\n    ls_eps=0.0\n    fc_dim=512\n    early_stopping_steps=5\n    grad_cam=False\n    seed=42\n    image_size = 256\n    batch_size = 64\n","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:22.883511Z","iopub.execute_input":"2022-07-25T23:59:22.883886Z","iopub.status.idle":"2022-07-25T23:59:22.897632Z","shell.execute_reply.started":"2022-07-25T23:59:22.883850Z","shell.execute_reply":"2022-07-25T23:59:22.896567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Library\n# ====================================================\nimport sys\nimport os\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\nfrom torch.cuda import amp\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn import preprocessing\nfrom sklearn import metrics\nfrom sklearn.metrics import roc_auc_score, roc_curve, f1_score, accuracy_score, log_loss\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold\n\nfrom tqdm.auto import tqdm\nfrom functools import partial\n\nimport cv2\nfrom PIL import Image\nfrom PIL import ImageFile\n# sometimes, you will have images without an ending bit\n# this takes care of those kind of (corrupt) images\nImageFile.LOAD_TRUNCATED_IMAGES = True\nimport albumentations as A \nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD, AdamW\nfrom torch.optim.optimizer import Optimizer\nimport torchvision.models as models\nfrom torch.nn.parameter import Parameter\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\n\nimport albumentations as A\nfrom torchvision import transforms ,datasets\nfrom torchvision.utils import make_grid\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\nfrom albumentations import (ToFloat, Normalize, VerticalFlip, HorizontalFlip, Compose, Resize,\n                            RandomBrightnessContrast, HueSaturationValue, Blur, GaussNoise,\n                            Rotate, RandomResizedCrop, Cutout, ShiftScaleRotate)\n\n\nimport timm\n\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nVERSION = 3","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:22.902680Z","iopub.execute_input":"2022-07-25T23:59:22.903913Z","iopub.status.idle":"2022-07-25T23:59:26.558612Z","shell.execute_reply.started":"2022-07-25T23:59:22.903872Z","shell.execute_reply":"2022-07-25T23:59:26.557538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# W&B","metadata":{}},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nwandb_api = user_secrets.get_secret(\"wandb_key\")\n\nimport wandb\nwandb.login(key=wandb_api)\n\ndef class2dict(f):\n    return dict((name, getattr(f, name)) for name in dir(f) if not name.startswith('__'))\n\nrun = wandb.init(project=\"Paddy Doctor Competition\", \n                 name=f\"{CFG.model_name} batch size\",\n                 config=class2dict(CFG),\n                 group=CFG.model_name,\n                 job_type=\"train\")","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:26.560383Z","iopub.execute_input":"2022-07-25T23:59:26.561264Z","iopub.status.idle":"2022-07-25T23:59:31.816248Z","shell.execute_reply.started":"2022-07-25T23:59:26.561221Z","shell.execute_reply":"2022-07-25T23:59:31.815232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Utils\n# ====================================================\ndef get_score(y_true, y_pred):\n    score = log_loss(y_true, y_pred)\n    return score\n\ndef init_logger(log_file=OUTPUT_DIR+'train.log'):\n    from logging import getLogger, INFO, FileHandler,  Formatter,  StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = init_logger()\n\n\ndef set_seed(seed = 1234):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed()","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:31.818344Z","iopub.execute_input":"2022-07-25T23:59:31.819022Z","iopub.status.idle":"2022-07-25T23:59:31.833548Z","shell.execute_reply.started":"2022-07-25T23:59:31.818983Z","shell.execute_reply":"2022-07-25T23:59:31.832596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV schem","metadata":{}},{"cell_type":"code","source":"%%time\ntrain_df[\"fold\"] = -1\ntrain_df = train_df.sample(frac=1).reset_index(drop=True)\ny = train_df.label.values\nkf = StratifiedKFold(n_splits=5)\nfor f, (t_, v_) in enumerate(kf.split(X=train_df, y=y)):\n    train_df.loc[v_, 'fold'] = f","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:31.835404Z","iopub.execute_input":"2022-07-25T23:59:31.836269Z","iopub.status.idle":"2022-07-25T23:59:31.861658Z","shell.execute_reply.started":"2022-07-25T23:59:31.836231Z","shell.execute_reply":"2022-07-25T23:59:31.859988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class classificationDataset:\n    def __init__(self,dataframe ,image_paths, targets ,resize=None, augmentations=None):\n\n        self.dataframe = dataframe\n        self.image_paths = image_paths\n        self.targets = targets\n        self.resize = resize\n        self.augmentations = augmentations\n    \n    def __len__(self):\n        return len(self.image_paths)\n    \n    def __getitem__(self, item):\n        \n        image = Image.open(self.image_paths[item])\n        targets = self.targets[item]\n        csv_data = np.array(self.dataframe.iloc[item][['variety','age']].values, dtype=np.float32)\n        \n        if self.resize is not None:\n            image = image.resize(\n                (self.resize[1], self.resize[0]), resample=Image.BILINEAR\n            )\n            \n        image = np.array(image)\n        \n        if self.augmentations is not None:\n            augmented = self.augmentations(image=image)\n            image = augmented[\"image\"]\n\n            \n        return image, np.array(csv_data) , torch.tensor(targets)","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:31.862848Z","iopub.execute_input":"2022-07-25T23:59:31.863192Z","iopub.status.idle":"2022-07-25T23:59:31.877990Z","shell.execute_reply.started":"2022-07-25T23:59:31.863155Z","shell.execute_reply":"2022-07-25T23:59:31.876640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transforms","metadata":{}},{"cell_type":"code","source":"mean = (0.485, 0.456, 0.406)\nstd = (0.229, 0.224, 0.225)\n\ntrain_aug = A.Compose(\n    [\n            A.Resize(height=256, width=256),\n            A.ShiftScaleRotate(rotate_limit=90, scale_limit = [0.8, 1.2]),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.HueSaturationValue(sat_shift_limit=[0.7, 1.3], hue_shift_limit=[-0.1, 0.1]),\n\n            #A.RandomBrightnessContrast(brightness_limit=[0.7, 1.3],contrast_limit= [0.7, 1.3]),\n            A.Normalize(),\n            A.pytorch.transforms.ToTensorV2()\n    ]\n)\ntta_aug = A.Compose([\n            A.Resize(height=256, width=256),\n            A.ShiftScaleRotate(rotate_limit=90, scale_limit = [0.8, 1.2]),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.HueSaturationValue(sat_shift_limit=[0.7, 1.3], hue_shift_limit=[-0.1, 0.1]),\n            A.CenterCrop(224, 224),\n            #A.RandomBrightnessContrast(brightness_limit=[0.7, 1.3],contrast_limit= [0.7, 1.3]),\n            A.Normalize(),\n            A.pytorch.transforms.ToTensorV2()\n])\n\nvalid_aug = A.Compose(\n    [       A.Resize(height=256, width=256),   \n            A.Normalize(),\n            A.pytorch.transforms.ToTensorV2()\n    ]\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:31.880044Z","iopub.execute_input":"2022-07-25T23:59:31.880708Z","iopub.status.idle":"2022-07-25T23:59:31.908758Z","shell.execute_reply.started":"2022-07-25T23:59:31.880673Z","shell.execute_reply":"2022-07-25T23:59:31.907340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Transforms\n# ====================================================\ndef get_transforms(*, data):\n    \n    if data == 'train':\n        return A.Compose(\n        [\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            A.Flip(p=0.5),\n            \n            #A.Cutout(p=0.5),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.Rotate(limit=180, p=0.5),\n            A.ShiftScaleRotate(\n                shift_limit = 0.1, scale_limit=0.1, rotate_limit=45, p=0.5\n            ),\n           \n            ToTensorV2(p=1.0),\n        ]\n    )\n\n    elif 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        ])","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:31.919013Z","iopub.execute_input":"2022-07-25T23:59:31.919287Z","iopub.status.idle":"2022-07-25T23:59:31.941494Z","shell.execute_reply.started":"2022-07-25T23:59:31.919262Z","shell.execute_reply":"2022-07-25T23:59:31.940311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_folds = train_df[train_df['fold'] != 0]\ntrain_images_path = train_folds.path_jpeg.values\ntrain_targets = train_folds.label.values\n\ntrain_dataset = classificationDataset(dataframe=train_folds,\n                                image_paths=train_images_path,\n                                targets=train_targets,\n                                augmentations=get_transforms(data='train'))\nfor i in range(5):\n    plt.figure(figsize=(4, 4))\n    image1,csv_data ,label = train_dataset[i]\n    plt.imshow(image1.permute(2,1,0))\n    plt.title(f'label: {label}')\n    plt.show() ","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:31.943146Z","iopub.execute_input":"2022-07-25T23:59:31.943840Z","iopub.status.idle":"2022-07-25T23:59:33.262592Z","shell.execute_reply.started":"2022-07-25T23:59:31.943803Z","shell.execute_reply":"2022-07-25T23:59:33.261608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss","metadata":{}},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, alpha=1, gamma=2, logits=True, reduce=True):\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.logits = logits\n        self.reduce = reduce\n\n    def forward(self, inputs, targets):\n        if self.logits:\n            BCE_loss = F.cross_entropy(inputs, targets, reduce=False)\n        else:\n            BCE_loss = F.cross_entropy(inputs, targets, reduce=False)\n        pt = torch.exp(-BCE_loss)\n        F_loss = self.alpha * (1-pt)**self.gamma * BCE_loss\n\n        if self.reduce:\n            return torch.mean(F_loss)\n        else:\n            return F_loss","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:33.264150Z","iopub.execute_input":"2022-07-25T23:59:33.264716Z","iopub.status.idle":"2022-07-25T23:59:33.278811Z","shell.execute_reply.started":"2022-07-25T23:59:33.264680Z","shell.execute_reply":"2022-07-25T23:59:33.277776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"sigmoid = torch.nn.Sigmoid()\n\nclass Swish(torch.autograd.Function):\n    @staticmethod\n    def forward(ctx, i):\n        result = i * sigmoid(i)\n        ctx.save_for_backward(i)\n        return result\n    @staticmethod\n    def backward(ctx, grad_output):\n        i = ctx.saved_variables[0]\n        sigmoid_i = sigmoid(i)\n        return grad_output * (sigmoid_i * (1 + i * (1 - sigmoid_i)))\nswish = Swish.apply\n\nclass Swish_module(nn.Module):\n    def forward(self, x):\n        return swish(x)\nswish_layer = Swish_module()\n\n\nclass CustomModel(nn.Module):\n    def __init__(self, model_name=CFG.model_name , out_features=CFG.target_size, pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained, in_chans=3)\n        \n        in_features = self.model.classifier.in_features\n        \n        self.model.classifier = nn.Linear(in_features, CFG.fc_dim)\n        \n        \n        self.csv = nn.Sequential(nn.Linear(2, 250),\n                                 nn.BatchNorm1d(250),\n                                 Swish_module(),\n                                 nn.Dropout(p=0.2),\n                                 \n                                 nn.Linear(250, 250),\n                                 nn.BatchNorm1d(250),\n                                 Swish_module(),\n                                 nn.Dropout(p=0.2))\n        \n        \n        self.probs = nn.Linear(CFG.fc_dim + 250, CFG.target_size)\n    \n    def forward(self, x, csv_data):\n        batch_size = x.shape[0]\n        #print(f'input : {x.shape}')\n        \n        features = self.model(x)\n        #print(f'features : {features.shape}')\n        \n        out_csv = self.csv(csv_data)\n        #print(f'out_csv : {out_csv.shape}')\n        \n        image_csv_data = torch.cat((features, out_csv), dim=1)\n        #print(f\"image_csv_data shape  :{image_csv_data.shape} \")\n        \n        probs   = self.probs(image_csv_data)\n        \n        return probs","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:33.280443Z","iopub.execute_input":"2022-07-25T23:59:33.281142Z","iopub.status.idle":"2022-07-25T23:59:33.301464Z","shell.execute_reply.started":"2022-07-25T23:59:33.281105Z","shell.execute_reply":"2022-07-25T23:59:33.300440Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Helper functions\n# ====================================================\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\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\n\ndef asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\n\ndef timeSince(since, percent):\n    now = time.time()\n    s = now - since\n    es = s / (percent)\n    rs = es - s\n    return '%s (remain %s)' % (asMinutes(s), asMinutes(rs))\n\n\ndef train_fn(fold, train_loader, model, criterion, optimizer, epoch, scheduler, device):\n    if CFG.apex:\n        scaler = GradScaler()\n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    # switch to train mode\n    model.train()\n    start = end = time.time()\n    global_step = 0\n    for step, (images,csv_data ,labels) in enumerate(train_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        \n        images = images.to(device).float()\n        labels = labels.to(device).long()\n        csv_data = csv_data.to(device).float()\n        \n        batch_size = labels.size(0)\n        if CFG.apex:\n            with autocast():\n                y_preds = model(images, csv_data)\n                loss = criterion(y_preds, labels)\n        else:\n            y_preds = model(images, csv_data)\n            loss = criterion(y_preds, labels)\n        # record loss\n        losses.update(loss.item(), batch_size)\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        if CFG.apex:\n            scaler.scale(loss).backward()\n        else:\n            loss.backward()\n        grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.max_grad_norm)\n        if (step + 1) % CFG.gradient_accumulation_steps == 0:\n            if CFG.apex:\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                optimizer.step()\n            optimizer.zero_grad()\n            global_step += 1\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(train_loader)-1):\n            print('Epoch: [{0}][{1}/{2}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  'Grad: {grad_norm:.4f} '\n                  'LR: {lr:.6f}  '\n                  .format(epoch+1, step, len(train_loader), \n                          remain=timeSince(start, float(step+1)/len(train_loader)),\n                          loss=losses,\n                          grad_norm=grad_norm,\n                          lr=scheduler.get_lr()[0]))\n        wandb.log({f\"[fold{fold}] loss\": losses.val,\n                   f\"[fold{fold}] lr\": scheduler.get_lr()[0]})\n    return losses.avg\n\n\ndef valid_fn(valid_loader, model, criterion, device):\n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    # switch to evaluation mode\n    model.eval()\n    preds = []\n    start = end = time.time()\n    for step, (images,csv_data ,labels) in enumerate(valid_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        \n        images = images.to(device).float()\n        labels = labels.to(device).long()\n        csv_data = csv_data.to(device).float()\n        \n        batch_size = labels.size(0)\n        # compute loss\n        with torch.no_grad():\n            y_preds = model(images,csv_data)\n        preds.append(y_preds.softmax(1).to('cpu').numpy())\n        loss = criterion(y_preds, labels)\n        losses.update(loss.item(), batch_size)\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(valid_loader)-1):\n            print('EVAL: [{0}/{1}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  .format(step, len(valid_loader),\n                          loss=losses,\n                          remain=timeSince(start, float(step+1)/len(valid_loader))))\n    predictions = np.concatenate(preds)\n    return losses.avg, predictions","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:33.302942Z","iopub.execute_input":"2022-07-25T23:59:33.303438Z","iopub.status.idle":"2022-07-25T23:59:33.336697Z","shell.execute_reply.started":"2022-07-25T23:59:33.303407Z","shell.execute_reply":"2022-07-25T23:59:33.335789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train loop","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Train loop\n# ====================================================\n\ndef train_loop(folds, fold):\n    scaler = amp.GradScaler()\n    LOGGER.info(f\"========== fold: {fold} training ==========\")\n    \n    if CFG.debug:\n        train_folds = folds[folds['fold'] != fold].sample(100)\n        valid_folds = folds[folds['fold'] == fold].sample(100)\n        \n    else:\n        train_folds = folds[folds['fold'] != fold]\n        valid_folds = folds[folds['fold'] == fold]\n        \n        \n    train_images_path = train_folds.path_jpeg.values\n    train_targets = train_folds.label.values\n        \n    valid_images_path = valid_folds.path_jpeg.values\n    valid_targets = valid_folds.label.values\n        \n    train_dataset = classificationDataset(dataframe=train_folds,\n                                image_paths=train_images_path,\n                                targets=train_targets,\n                                augmentations=get_transforms(data='train'))\n        \n    valid_dataset = classificationDataset(dataframe=valid_folds,\n                                image_paths=valid_images_path,\n                                targets=valid_targets,\n                                augmentations=get_transforms(data='valid'))\n        \n    train_loader = torch.utils.data.DataLoader(\n        train_dataset, batch_size=CFG.batch_size, shuffle=True, num_workers=4\n    )\n        \n    valid_loader = torch.utils.data.DataLoader(\n        valid_dataset, batch_size=CFG.batch_size, shuffle=False, num_workers=4\n    )\n    \n    # ====================================================\n    # scheduler \n    # ====================================================\n    def get_scheduler(optimizer):\n        if CFG.scheduler=='ReduceLROnPlateau':\n            scheduler = ReduceLROnPlateau(optimizer, **CFG.reduce_params)\n        elif CFG.scheduler=='CosineAnnealingLR':\n            scheduler = CosineAnnealingLR(optimizer, **CFG.cosanneal_params)\n        elif CFG.scheduler=='CosineAnnealingWarmRestarts':\n            scheduler = CosineAnnealingWarmRestarts(optimizer, **CFG.reduce_params)\n        return scheduler\n    \n    # ====================================================\n    # model & optimizer\n    # ====================================================\n        \n    model = CustomModel()\n    model.to(device)\n\n    optimizer = Adam(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\n    scheduler = get_scheduler(optimizer)\n\n    # ====================================================\n    # loop\n    # ====================================================\n    criterion = nn.CrossEntropyLoss()\n    best_loss = np.inf\n    best_acc = 0.0\n    \n    for epoch in range(CFG.epochs):\n        \n        start_time = time.time()\n        \n        # train\n        avg_loss = train_fn(fold, train_loader, model, criterion, optimizer, epoch, scheduler, device)\n\n        # eval\n        avg_val_loss, preds = valid_fn(valid_loader, model, criterion, device)\n        \n        if isinstance(scheduler, ReduceLROnPlateau):\n            scheduler.step(avg_val_loss)\n        elif isinstance(scheduler, CosineAnnealingLR):\n            scheduler.step()\n        elif isinstance(scheduler, CosineAnnealingWarmRestarts):\n            scheduler.step()\n\n        # scoring\n        \n        preds_label = np.argmax(preds, axis=1)\n        Accuracy = accuracy_score(preds_label, valid_targets)\n        score = get_score(valid_targets, preds)\n\n        elapsed = time.time() - start_time\n\n        LOGGER.info(f'Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  avg_val_loss: {avg_val_loss:.4f}  time: {elapsed:.0f}s')\n        LOGGER.info(f'Epoch {epoch+1} - Score: {score:.4f}')\n        LOGGER.info(f'Epoch {epoch+1} - Accuracy: {Accuracy:.4f}')\n        wandb.log({f\"[fold{fold}] epoch\": epoch+1, \n                   f\"[fold{fold}] avg_train_loss\": avg_loss, \n                   f\"[fold{fold}] avg_val_loss\": avg_val_loss,\n                   f\"[fold{fold}] Accuracy\": Accuracy,\n                   f\"[fold{fold}] score\": score})\n\n            \n        if Accuracy > best_acc:\n            best_acc = Accuracy\n            LOGGER.info(f'Epoch {epoch+1} - Save Best Accuracy: {best_acc:.4f} Model')\n            torch.save({'model': model.state_dict(), \n                        'preds_loss': preds},\n                        OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best_Accuracy.pth')\n        \n        \n    valid_folds[CFG.preds_col] = torch.load(OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best_Accuracy.pth', \n                                      map_location=torch.device('cpu'))['preds_loss']\n\n    return valid_folds","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:33.338258Z","iopub.execute_input":"2022-07-25T23:59:33.338633Z","iopub.status.idle":"2022-07-25T23:59:33.362010Z","shell.execute_reply.started":"2022-07-25T23:59:33.338597Z","shell.execute_reply":"2022-07-25T23:59:33.360985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# main\n# ====================================================\ndef main():\n\n    \"\"\"\n    Prepare: 1.train \n    \"\"\"\n\n    def get_result(result_df):\n        preds_loss = result_df[CFG.preds_col].values\n        labels = result_df[\"label\"].values\n        score_loss = get_score(labels, preds_loss)\n        LOGGER.info(f'Score with best loss weights: {score_loss:<.4f}')\n    \n    if CFG.train:\n        # train \n        oof_df = pd.DataFrame()\n        for fold in range(CFG.nfolds):\n            if fold in CFG.trn_folds:\n                _oof_df = train_loop(train_df, fold)\n                oof_df = pd.concat([oof_df, _oof_df])\n                LOGGER.info(f\"========== fold: {fold} result ==========\")\n                get_result(_oof_df)\n        # CV result\n        LOGGER.info(f\"========== CV ==========\")\n        get_result(oof_df)\n        # save result\n        oof_df[['bacterial_leaf_blight', 'bacterial_leaf_streak',\n       'bacterial_panicle_blight', 'blast', 'brown_spot', 'dead_heart',\n       'downy_mildew', 'hispa', 'normal', 'tungro']].to_csv(OUTPUT_DIR+f'{CFG.model_name}_oof_rgb_df_version{VERSION}.csv', index=False)\n        \n    wandb.finish()","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:33.364934Z","iopub.execute_input":"2022-07-25T23:59:33.365603Z","iopub.status.idle":"2022-07-25T23:59:33.376065Z","shell.execute_reply.started":"2022-07-25T23:59:33.365566Z","shell.execute_reply":"2022-07-25T23:59:33.374911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:33.378166Z","iopub.execute_input":"2022-07-25T23:59:33.378446Z","iopub.status.idle":"2022-07-25T23:59:33.580139Z","shell.execute_reply.started":"2022-07-25T23:59:33.378419Z","shell.execute_reply":"2022-07-25T23:59:33.578436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    main()","metadata":{"execution":{"iopub.status.busy":"2022-07-25T23:59:33.581835Z","iopub.execute_input":"2022-07-25T23:59:33.582222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}