{"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 -q --upgrade wandb \n!pip install  timm\n!pip install -U git+https://github.com/albu/albumentations > /dev/null && echo  \n!pip install opencv-python==4.5.5.64","metadata":{"id":"0KF62LrhX1YC","execution":{"iopub.status.busy":"2022-08-11T02:51:50.742357Z","iopub.execute_input":"2022-08-11T02:51:50.743215Z","iopub.status.idle":"2022-08-11T02:52:45.089875Z","shell.execute_reply.started":"2022-08-11T02:51:50.742719Z","shell.execute_reply":"2022-08-11T02:52:45.088684Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!conda install -y gdown\n\n!gdown --id 1OglNeDGuipJ4wcPFVIRy_eqxHGePIJ55","metadata":{"execution":{"iopub.status.busy":"2022-08-11T02:52:45.093285Z","iopub.execute_input":"2022-08-11T02:52:45.094121Z","iopub.status.idle":"2022-08-11T02:53:54.683317Z","shell.execute_reply.started":"2022-08-11T02:52:45.094080Z","shell.execute_reply":"2022-08-11T02:53:54.682082Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport wandb\nfrom skimage import filters, transform\nfrom skimage.io import imread\nfrom skimage import img_as_ubyte\nfrom typing import Tuple\n\n# ====================================================\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, OrderedDict\n\n\nimport scipy as sp\n\nfrom sklearn import preprocessing\nfrom sklearn.metrics import roc_auc_score, roc_curve, f1_score, accuracy_score\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold, train_test_split\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\n\nimport torch\nimport torchvision\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD\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, OneCycleLR\n\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\nimport timm\n\nfrom torch.cuda.amp import autocast, GradScaler\n\n# Functions for plotting:\nimport matplotlib.pyplot as plt\n%matplotlib inline\nplt.rcParams['image.cmap'] = 'Greys'\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nVERSION=1","metadata":{"id":"WE10PBb5YJmA","execution":{"iopub.status.busy":"2022-08-11T02:53:54.685204Z","iopub.execute_input":"2022-08-11T02:53:54.686851Z","iopub.status.idle":"2022-08-11T02:53:59.053536Z","shell.execute_reply.started":"2022-08-11T02:53:54.686803Z","shell.execute_reply":"2022-08-11T02:53:59.052432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv('../input/paddy-disease-classification/train.csv')\n\nsubmission = pd.read_csv('../input/paddy-disease-classification/sample_submission.csv')\ntrain_dir = '../input/paddy-disease-classification/train_images'","metadata":{"id":"3lteQt2VYSm6","execution":{"iopub.status.busy":"2022-08-11T02:53:59.056772Z","iopub.execute_input":"2022-08-11T02:53:59.058054Z","iopub.status.idle":"2022-08-11T02:53:59.091226Z","shell.execute_reply.started":"2022-08-11T02:53:59.058010Z","shell.execute_reply":"2022-08-11T02:53:59.090342Z"},"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":{"id":"Hr6aIQKiYpoN","execution":{"iopub.status.busy":"2022-08-11T02:53:59.092511Z","iopub.execute_input":"2022-08-11T02:53:59.092970Z","iopub.status.idle":"2022-08-11T02:53:59.234855Z","shell.execute_reply.started":"2022-08-11T02:53:59.092931Z","shell.execute_reply":"2022-08-11T02:53:59.233895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.head()","metadata":{"id":"HxCvKemOYp5D","outputId":"1c48e875-d517-4426-fa29-596e0ab39682","execution":{"iopub.status.busy":"2022-08-11T02:53:59.236361Z","iopub.execute_input":"2022-08-11T02:53:59.236741Z","iopub.status.idle":"2022-08-11T02:53:59.255004Z","shell.execute_reply.started":"2022-08-11T02:53:59.236702Z","shell.execute_reply":"2022-08-11T02:53:59.254163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.listdir('../input/paddy-disease-classification/train_images')","metadata":{"execution":{"iopub.status.busy":"2022-08-11T02:53:59.256414Z","iopub.execute_input":"2022-08-11T02:53:59.256757Z","iopub.status.idle":"2022-08-11T02:53:59.268092Z","shell.execute_reply.started":"2022-08-11T02:53:59.256723Z","shell.execute_reply":"2022-08-11T02:53:59.267165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\nimport shutil\nimport os\n\n!mkdir \"./Paddy_Images\"\n\nbase_path = '../input/paddy-disease-classification/train_images/'\nsrc_list = os.listdir(base_path)\n\ndst_dir = \"./Paddy_Images\"\nfor _,src_dir in tqdm(enumerate(src_list)):\n    scr_dir_path = os.path.join(base_path,src_dir)\n    for jpgfile in glob.iglob(os.path.join(scr_dir_path, \"*.jpg\")):\n        shutil.copy(jpgfile, dst_dir)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T02:53:59.271292Z","iopub.execute_input":"2022-08-11T02:53:59.271565Z","iopub.status.idle":"2022-08-11T02:54:47.085118Z","shell.execute_reply.started":"2022-08-11T02:53:59.271524Z","shell.execute_reply":"2022-08-11T02:54:47.083656Z"},"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":{"id":"T3sAX5DCYrv7","execution":{"iopub.status.busy":"2022-08-11T02:54:47.088389Z","iopub.execute_input":"2022-08-11T02:54:47.088868Z","iopub.status.idle":"2022-08-11T02:54:47.113843Z","shell.execute_reply.started":"2022-08-11T02:54:47.088817Z","shell.execute_reply":"2022-08-11T02:54:47.112448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.random.seed(2022)","metadata":{"id":"ycbVLGmOYzrH","execution":{"iopub.status.busy":"2022-08-11T02:54:47.129620Z","iopub.execute_input":"2022-08-11T02:54:47.132855Z","iopub.status.idle":"2022-08-11T02:54:47.143369Z","shell.execute_reply.started":"2022-08-11T02:54:47.132805Z","shell.execute_reply":"2022-08-11T02:54:47.142101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"siamese_data = pd.read_csv('./siamese_data.csv')\nsiamese_data","metadata":{"id":"xsbwpnJRZGIN","outputId":"5e3498ee-85ad-434c-d5ac-134101b8b117","execution":{"iopub.status.busy":"2022-08-11T02:54:47.149525Z","iopub.execute_input":"2022-08-11T02:54:47.152979Z","iopub.status.idle":"2022-08-11T02:54:47.427658Z","shell.execute_reply.started":"2022-08-11T02:54:47.152926Z","shell.execute_reply":"2022-08-11T02:54:47.426149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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='OneCycleLR' # ['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts','OneCycleLR']\n    epochs=8\n    # CosineAnnealingLR params\n    cosanneal_params={\n        'T_max':4,\n        'eta_min':1e-5,\n        'last_epoch':-1\n    }\n    #ReduceLROnPlateau params\n    reduce_params={\n        'mode':'min',\n        'factor':0.2,\n        'patience':4,\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    onecycle_params={\n        'pct_start':0.1,\n        'div_factor':1e2,\n        'max_lr':1e-3,\n        'steps_per_epoch':7, \n        'epochs':7\n    }\n    batch_size=32\n    lr=1e-2\n    weight_decay=1e-3\n    gradient_accumulation_steps=1\n    max_grad_norm=1000\n    target_size=siamese_data[\"label\"].shape[0]\n    nfolds=5\n    trn_folds=[0]\n    model_name='resnet50d'     #'vit_base_patch32_224_in21k' 'tf_efficientnetv2_b0' 'resnext50_32x4d' 'tresnet_m'\n    train=True\n    early_stop=True\n    target_col=\"label\"\n    fc_dim=512\n    early_stopping_steps=5\n    grad_cam=False\n    seed=42\n    \nif CFG.debug:\n    CFG.epochs=1\n    train=siamese_data.sample(n=1000, random_state=CFG.seed).reset_index(drop=True)","metadata":{"id":"dOtuOg_DZkqC","execution":{"iopub.status.busy":"2022-08-11T02:54:47.435192Z","iopub.execute_input":"2022-08-11T02:54:47.438342Z","iopub.status.idle":"2022-08-11T02:54:47.464759Z","shell.execute_reply.started":"2022-08-11T02:54:47.438284Z","shell.execute_reply":"2022-08-11T02:54:47.463500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nOUTPUT_DIR = f'./'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)","metadata":{"id":"DDCN2Y2IaAaQ","execution":{"iopub.status.busy":"2022-08-11T02:54:47.467842Z","iopub.execute_input":"2022-08-11T02:54:47.468241Z","iopub.status.idle":"2022-08-11T02:54:47.521014Z","shell.execute_reply.started":"2022-08-11T02:54:47.468196Z","shell.execute_reply":"2022-08-11T02:54:47.519748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"id":"_m4U3CTxcG6i","outputId":"1d886e61-aee6-43c7-ea53-3895bd9d6a5a","execution":{"iopub.status.busy":"2022-08-11T02:54:47.523101Z","iopub.execute_input":"2022-08-11T02:54:47.523498Z","iopub.status.idle":"2022-08-11T02:54:52.549441Z","shell.execute_reply.started":"2022-08-11T02:54:47.523461Z","shell.execute_reply":"2022-08-11T02:54:52.548389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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 seed_torch(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_torch(seed=CFG.seed)","metadata":{"id":"jkShLuE6cVov","execution":{"iopub.status.busy":"2022-08-11T02:54:52.551178Z","iopub.execute_input":"2022-08-11T02:54:52.551890Z","iopub.status.idle":"2022-08-11T02:54:52.593661Z","shell.execute_reply.started":"2022-08-11T02:54:52.551845Z","shell.execute_reply":"2022-08-11T02:54:52.584257Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RiceSiameseDataset(Dataset):\n    \n    def __init__(self, df=train_df, transform=None):\n        self.df  = df\n        self.path = './Paddy_Images/'\n        self.labels = df[\"label\"].values\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n        \n    def __getitem__(self,index):\n        # getting the image path\n        img1_id = self.df.iloc[index]['image1']\n        img2_id = self.df.iloc[index]['image2']\n\n        try:\n\n            image1 = cv2.imread(self.path+img1_id)\n            image1 = cv2.cvtColor(image1, cv2.COLOR_BGR2RGB)\n            image2 = cv2.imread(self.path+img2_id)\n            image2 = cv2.cvtColor(image2, cv2.COLOR_BGR2RGB)\n\n        except:\n          \n            image1 = Image.open(self.path+img1_id)\n            image1 = image1.convert(\"RGB\")\n            image1 = np.array(image1)\n\n            image2 = Image.open(self.path+img2_id)\n            image2 = image2.convert(\"RGB\")\n            image2 = np.array(image2)\n\n        if self.transform:\n            image1 = self.transform(image=image1)['image']\n            image2 = self.transform(image=image2)['image']\n        \n        label = torch.tensor(self.labels[index], dtype=torch.long)\n        \n        return image1, image2, label","metadata":{"id":"uGVb6HcNceMF","execution":{"iopub.status.busy":"2022-08-11T02:54:52.596229Z","iopub.execute_input":"2022-08-11T02:54:52.596605Z","iopub.status.idle":"2022-08-11T02:54:52.619379Z","shell.execute_reply.started":"2022-08-11T02:54:52.596570Z","shell.execute_reply":"2022-08-11T02:54:52.618518Z"},"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\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":{"id":"LfBc2AoScsOu","execution":{"iopub.status.busy":"2022-08-11T02:54:52.620358Z","iopub.execute_input":"2022-08-11T02:54:52.620679Z","iopub.status.idle":"2022-08-11T02:54:52.667709Z","shell.execute_reply.started":"2022-08-11T02:54:52.620646Z","shell.execute_reply":"2022-08-11T02:54:52.666396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = RiceSiameseDataset(siamese_data, transform=get_transforms(data='train'))\nimg1, img2, _ = train_dataset[120]\n\nf, ax = plt.subplots(2,1, figsize=(15,10))\nax[0].imshow(img1[0])\nax[1].imshow(img2[0])","metadata":{"id":"dgJ9gpZBc1ZQ","outputId":"4c567124-1866-4f25-f2c1-9b5a112b75dd","execution":{"iopub.status.busy":"2022-08-11T02:54:52.672454Z","iopub.execute_input":"2022-08-11T02:54:52.675080Z","iopub.status.idle":"2022-08-11T02:54:53.168674Z","shell.execute_reply.started":"2022-08-11T02:54:52.675039Z","shell.execute_reply":"2022-08-11T02:54:53.167683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img1.shape","metadata":{"id":"qls1OwGSc26g","outputId":"edd08ed0-114c-43f6-a0a8-aa7e1b33fad5","execution":{"iopub.status.busy":"2022-08-11T02:54:53.169787Z","iopub.execute_input":"2022-08-11T02:54:53.171186Z","iopub.status.idle":"2022-08-11T02:54:53.185687Z","shell.execute_reply.started":"2022-08-11T02:54:53.171132Z","shell.execute_reply":"2022-08-11T02:54:53.179037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img2.shape","metadata":{"id":"fvAsDLKkfyJY","outputId":"b19dcf96-9ed3-4628-9cba-d307882f9901","execution":{"iopub.status.busy":"2022-08-11T02:54:53.189445Z","iopub.execute_input":"2022-08-11T02:54:53.192107Z","iopub.status.idle":"2022-08-11T02:54:53.202769Z","shell.execute_reply.started":"2022-08-11T02:54:53.192065Z","shell.execute_reply":"2022-08-11T02:54:53.201887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# MODEL\n# ====================================================\nclass CustomModel(nn.Module):\n    def __init__(self, cfg, pretrained=False):\n        super().__init__()\n        self.cfg = cfg\n        self.model = timm.create_model(self.cfg.model_name, pretrained=pretrained, in_chans=3)\n        \n        if cfg.model_name == 'tf_efficientnetv2_b0':\n            self.n_features = self.model.classifier.in_features\n            self.model.classifier = nn.Linear(self.n_features, self.cfg.fc_dim)\n        \n        if cfg.model_name == \"efficientnet\":\n            self.n_features = self.model.classifier.in_features\n            self.model.classifier = nn.Linear(self.n_features, self.cfg.fc_dim)\n            \n        if cfg.model_name in ['resnext50_32x4d', 'resnet50d']:\n            self.in_features = self.model.fc.in_features\n            #self.model.fc = nn.Linear(self.in_features, self.cfg.fc_dim)\n            self.model.fc = nn.Identity()\n            self.model.global_pool = nn.Identity()\n            \n        if cfg.model_name == 'tresnet_m':\n            self.in_features = self.model.head.fc.in_features\n            self.model.head.fc = nn.Linear(self.in_features, self.cfg.fc_dim)\n            \n        elif cfg.model_name.split('_')[0] == 'vit':\n            self.n_features = self.model.head.in_features\n            self.model.head = nn.Linear(self.n_features, self.cfg.fc_dim)\n        \n        \n        #self.pooling = GeM()\n        self.pooling =  nn.AdaptiveAvgPool2d(1) # GAP\n        self.fc = nn.Linear(self.in_features, self.cfg.fc_dim)\n        \n    def forward(self, x1, x2):\n        batch_size = x1.shape[0]\n        # model backbone shape: torch.Size([4, 2048, 8, 8])\n        features1 = self.model(x1)\n        # gap shape: torch.Size([4, 2048])\n        features1 = self.pooling(features1).view(batch_size, -1)\n        fc1   = self.fc(features1)\n\n        # model backbone shape: torch.Size([4, 2048, 8, 8])\n        features2 = self.model(x2)\n        # gap shape: torch.Size([4, 2048])\n        features2 = self.pooling(features2).view(batch_size, -1)\n        fc2   = self.fc(features2)\n        return fc1, fc2","metadata":{"id":"HAwmOa50fzbb","execution":{"iopub.status.busy":"2022-08-11T02:54:53.205325Z","iopub.execute_input":"2022-08-11T02:54:53.205663Z","iopub.status.idle":"2022-08-11T02:54:53.231020Z","shell.execute_reply.started":"2022-08-11T02:54:53.205629Z","shell.execute_reply":"2022-08-11T02:54:53.229951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x1 = torch.rand((4, 3, 256, 256))\nx2 = torch.rand((4, 3, 256, 256))\n\nmodel = CustomModel(CFG)\ny1, y2 = model(x1, x2)","metadata":{"id":"UZexdCa4f1X_","execution":{"iopub.status.busy":"2022-08-11T02:54:53.235372Z","iopub.execute_input":"2022-08-11T02:54:53.235880Z","iopub.status.idle":"2022-08-11T02:54:56.244589Z","shell.execute_reply.started":"2022-08-11T02:54:53.235836Z","shell.execute_reply":"2022-08-11T02:54:56.243486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y1.shape","metadata":{"id":"nM9D7cK6f25N","outputId":"ce3cfcd3-d863-4384-b816-620436363781","execution":{"iopub.status.busy":"2022-08-11T02:54:56.247947Z","iopub.execute_input":"2022-08-11T02:54:56.248654Z","iopub.status.idle":"2022-08-11T02:54:56.259740Z","shell.execute_reply.started":"2022-08-11T02:54:56.248607Z","shell.execute_reply":"2022-08-11T02:54:56.258868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ContrastiveLoss(torch.nn.Module):\n    \"\"\"\n    Contrastive loss function.\n    Based on: http://yann.lecun.com/exdb/publis/pdf/hadsell-chopra-lecun-06.pdf\n    \"\"\"\n\n    def __init__(self, margin=1.0):\n        super(ContrastiveLoss, self).__init__()\n        self.margin = margin\n\n    def forward(self, output1, output2, label):\n        euclidean_distance = F.cosine_similarity(output1, output2)\n        loss_contrastive = torch.mean((1-label) * torch.pow(euclidean_distance, 2) +\n                                      (label) * torch.pow(torch.clamp(self.margin - euclidean_distance, min=0.0), 2))\n\n\n        return loss_contrastive","metadata":{"id":"gU-OCLJIf4gB","execution":{"iopub.status.busy":"2022-08-11T02:54:56.261086Z","iopub.execute_input":"2022-08-11T02:54:56.261966Z","iopub.status.idle":"2022-08-11T02:54:56.278190Z","shell.execute_reply.started":"2022-08-11T02:54:56.261929Z","shell.execute_reply":"2022-08-11T02:54:56.277112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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, layer):\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    tk0 = tqdm(train_loader, total=len(train_loader))\n    for step, (img1, img2, labels) in enumerate(tk0):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        img1 = img1.to(device).float()\n        img2 = img2.to(device).float()\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        if CFG.apex:\n            with autocast():\n                out1, out2 = model(img1, img2)\n                layer1 = layer(out1)\n                layer2 = layer(out2)\n                loss = criterion(layer1, layer2, labels)\n        else:\n            out1, out2 = model(img1, img2)\n            loss = criterion(out1, out2, labels)\n            \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_last_lr()[0]))\n        wandb.log({f\"[fold{fold}] loss\": losses.val,\n                   f\"[fold{fold}] lr\": scheduler.get_last_lr()[0]})\n    return losses.avg","metadata":{"id":"JkSnh1t5f7R2","execution":{"iopub.status.busy":"2022-08-11T02:54:56.283418Z","iopub.execute_input":"2022-08-11T02:54:56.283979Z","iopub.status.idle":"2022-08-11T02:54:56.308369Z","shell.execute_reply.started":"2022-08-11T02:54:56.283943Z","shell.execute_reply":"2022-08-11T02:54:56.307192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Train loop\n# ====================================================\ndef train_loop():\n    \n\n    # ====================================================\n    # loader\n    # ====================================================\n    \n    train_dataset = RiceSiameseDataset(siamese_data, transform=get_transforms(data='train'))\n\n    train_loader = DataLoader(train_dataset,\n                              batch_size=CFG.batch_size, \n                              shuffle=True, \n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=True)\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        elif CFG.scheduler=='OneCycleLR':\n            scheduler = OneCycleLR(optimizer, **CFG.onecycle_params)\n        return scheduler\n\n    # ====================================================\n    # model & optimizer\n    # ====================================================\n    model = CustomModel(CFG, pretrained=True)\n    model.to(device)\n    \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    criterions = ContrastiveLoss()\n    contrastive_layer = nn.Linear(512, 2)\n    best_score = 0.\n    best_loss = np.inf\n    \n    for epoch in range(CFG.epochs):\n        \n        start_time = time.time()\n        \n        # train\n        avg_loss = train_fn(train_df, train_loader, model, criterions, optimizer, epoch, scheduler, device, contrastive_layer)\n\n        \n        if isinstance(scheduler, ReduceLROnPlateau):\n            scheduler.step(avg_loss)\n        elif isinstance(scheduler, CosineAnnealingLR):\n            scheduler.step()\n        elif isinstance(scheduler, CosineAnnealingWarmRestarts):\n            scheduler.step()\n        elif isinstance(scheduler, OneCycleLR):\n            scheduler.step()\n\n\n        elapsed = time.time() - start_time\n\n        LOGGER.info(f'Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  time: {elapsed:.0f}s')\n        wandb.log({f\"epoch\": epoch+1, \n                   f\"avg_train_loss\": avg_loss})\n                 \n            \n        if avg_loss < best_loss:\n            best_loss = avg_loss\n            LOGGER.info(f'Epoch {epoch+1} - Save Best Loss: {best_loss:.4f} Model')\n            torch.save({'model': model.state_dict()}, \n                        OUTPUT_DIR+f'siamese_{CFG.model_name}_best_loss.pt')\n   \n   \n    wandb.finish()    \n    return model","metadata":{"id":"lBpTiomBf8Lk","execution":{"iopub.status.busy":"2022-08-11T02:54:56.313522Z","iopub.execute_input":"2022-08-11T02:54:56.313872Z","iopub.status.idle":"2022-08-11T02:54:56.330063Z","shell.execute_reply.started":"2022-08-11T02:54:56.313844Z","shell.execute_reply":"2022-08-11T02:54:56.329185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = train_loop()","metadata":{"id":"IHu_50BAf96A","outputId":"be50fe7b-11f4-498d-f338-fb277ae8ab32","execution":{"iopub.status.busy":"2022-08-11T02:54:56.331843Z","iopub.execute_input":"2022-08-11T02:54:56.332572Z"},"trusted":true},"execution_count":null,"outputs":[]}]}