{"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- Pytorch timm resnet50d + GAP pooling\n- Added arcface loss to the crossentropy loss\n- MultiLabelStratificationKFold instead of StratificationKFold as stated [here](http://www.kaggle.com/competitions/paddy-disease-classification/discussion/321731#1801785).","metadata":{}},{"cell_type":"markdown","source":"# Install Specific library","metadata":{}},{"cell_type":"code","source":"!pip install -q --upgrade wandb\n!pip install  timm\n!pip install iterative-stratification","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:56:37.462475Z","iopub.execute_input":"2022-08-07T01:56:37.463530Z","iopub.status.idle":"2022-08-07T01:57:14.001357Z","shell.execute_reply.started":"2022-08-07T01:56:37.462925Z","shell.execute_reply":"2022-08-07T01:57:14.000216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{}},{"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\n\nTRAIN_DIR = '../input/paddy-disease-classification/train_images/'","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:14.004798Z","iopub.execute_input":"2022-08-07T01:57:14.005404Z","iopub.status.idle":"2022-08-07T01:57:14.895441Z","shell.execute_reply.started":"2022-08-07T01:57:14.005367Z","shell.execute_reply":"2022-08-07T01:57:14.894449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('../input/paddy-disease-classification/train.csv')\ntrain[\"image_paths\"] = train.apply(lambda row: TRAIN_DIR + row['label'] + '/' + row['image_id'], axis=1)","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:14.896839Z","iopub.execute_input":"2022-08-07T01:57:14.897203Z","iopub.status.idle":"2022-08-07T01:57:15.063412Z","shell.execute_reply.started":"2022-08-07T01:57:14.897161Z","shell.execute_reply":"2022-08-07T01:57:15.062519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train[\"label\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:15.066086Z","iopub.execute_input":"2022-08-07T01:57:15.066440Z","iopub.status.idle":"2022-08-07T01:57:15.082507Z","shell.execute_reply.started":"2022-08-07T01:57:15.066405Z","shell.execute_reply":"2022-08-07T01:57:15.081576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train[\"image_paths\"]","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:15.085126Z","iopub.execute_input":"2022-08-07T01:57:15.085475Z","iopub.status.idle":"2022-08-07T01:57:15.095690Z","shell.execute_reply.started":"2022-08-07T01:57:15.085442Z","shell.execute_reply":"2022-08-07T01:57:15.094747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Quick EDA","metadata":{}},{"cell_type":"markdown","source":"# Data Preprocessing","metadata":{}},{"cell_type":"code","source":"from sklearn import preprocessing\n\nle = preprocessing.LabelEncoder()\nle.fit(train['label'])\ntrain['label'] = le.transform(train['label'])","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:15.097075Z","iopub.execute_input":"2022-08-07T01:57:15.098022Z","iopub.status.idle":"2022-08-07T01:57:15.158073Z","shell.execute_reply.started":"2022-08-07T01:57:15.097988Z","shell.execute_reply":"2022-08-07T01:57:15.157215Z"},"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-08-07T01:57:15.159441Z","iopub.execute_input":"2022-08-07T01:57:15.159796Z","iopub.status.idle":"2022-08-07T01:57:15.166131Z","shell.execute_reply.started":"2022-08-07T01:57:15.159763Z","shell.execute_reply":"2022-08-07T01:57:15.165191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Library","metadata":{}},{"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\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn import preprocessing\nfrom sklearn.metrics import roc_auc_score, roc_curve, f1_score, accuracy_score, log_loss\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold\nfrom iterstrat.ml_stratifiers import MultilabelStratifiedKFold\n\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 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, OneCycleLR\n\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\n\n\nimport timm\n\nfrom torch.cuda.amp import autocast, GradScaler\n\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nVERSION = 1","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:15.167607Z","iopub.execute_input":"2022-08-07T01:57:15.168156Z","iopub.status.idle":"2022-08-07T01:57:18.526126Z","shell.execute_reply.started":"2022-08-07T01:57:15.168120Z","shell.execute_reply":"2022-08-07T01:57:18.524976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"class CFG:\n    debug=False\n    apex=False\n    print_freq=100\n    size=256\n    num_workers=4\n    scheduler='CosineAnnealingLR' # ['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts','OneCycleLR']\n    epochs=60\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    # 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':1e2,\n        'max_lr':1e-4,\n        'steps_per_epoch':20, \n        'epochs':20\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    target_size=len(le.classes_)\n    fold=0\n    nfolds=5\n    trn_folds=[0, 1, 2, 3, 4]\n    model_name='tf_efficientnet_b0'     #'vit_base_patch32_224_in21k' 'tf_efficientnetv2_b0' 'resnext50_32x4d' 'resnet50d'\n    preds_col = train[\"label\"].value_counts().index.values.tolist()\n    train=True\n    early_stop=True\n    target_col=\"label\"\n    scale=25.0\n    margin=0.60\n    easy_margin=False\n    ls_eps=0.0\n    fc_dim=512\n    early_stopping_steps=5\n    seed=42\n    \nif CFG.debug:\n    CFG.epochs=1\n    train=train.sample(n=1000, random_state=CFG.seed).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:18.528306Z","iopub.execute_input":"2022-08-07T01:57:18.529301Z","iopub.status.idle":"2022-08-07T01:57:18.545477Z","shell.execute_reply.started":"2022-08-07T01:57:18.529255Z","shell.execute_reply":"2022-08-07T01:57:18.544501Z"},"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=\"Microsoft Rice 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-08-07T01:57:18.551467Z","iopub.execute_input":"2022-08-07T01:57:18.551955Z","iopub.status.idle":"2022-08-07T01:57:23.534716Z","shell.execute_reply.started":"2022-08-07T01:57:18.551929Z","shell.execute_reply":"2022-08-07T01:57:23.533590Z"},"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 = accuracy_score(y_true, y_pred.argmax(1))\n    return score\n\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 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":{"execution":{"iopub.status.busy":"2022-08-07T01:57:23.536854Z","iopub.execute_input":"2022-08-07T01:57:23.537450Z","iopub.status.idle":"2022-08-07T01:57:23.551311Z","shell.execute_reply.started":"2022-08-07T01:57:23.537410Z","shell.execute_reply":"2022-08-07T01:57:23.550292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV schem","metadata":{}},{"cell_type":"code","source":"%%time\nskf = MultilabelStratifiedKFold(n_splits=CFG.nfolds, shuffle=True, random_state=CFG.seed)\nfor fold, (trn_idx, vld_idx) in enumerate(skf.split(train, train[['label', 'age', 'variety']])):\n    train.loc[vld_idx, \"folds\"] = int(fold)\ntrain[\"folds\"] = train[\"folds\"].astype(int)","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:23.553150Z","iopub.execute_input":"2022-08-07T01:57:23.553978Z","iopub.status.idle":"2022-08-07T01:57:23.947071Z","shell.execute_reply.started":"2022-08-07T01:57:23.553931Z","shell.execute_reply":"2022-08-07T01:57:23.946096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Dataset\n# ====================================================\nclass TrainDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df['image_paths'].values\n        self.labels = df[CFG.target_col].values\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_path = self.file_names[idx]\n        image = cv2.imread(file_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n       \n        if self.transform:\n            image = self.transform(image=image)['image']\n        label = torch.tensor(self.labels[idx]).float()\n        return image, label","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:23.948508Z","iopub.execute_input":"2022-08-07T01:57:23.948860Z","iopub.status.idle":"2022-08-07T01:57:23.961368Z","shell.execute_reply.started":"2022-08-07T01:57:23.948825Z","shell.execute_reply":"2022-08-07T01:57:23.959718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transforms","metadata":{}},{"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.Transpose(p=0.5),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.ShiftScaleRotate(p=0.5),\n        A.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n        ToTensorV2(),\n        ]\n    )\n\n    elif data == 'valid':\n        return A.Compose([\n            A.Resize(CFG.size, CFG.size),\n            A.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n            ToTensorV2(),\n        ])","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:23.964418Z","iopub.execute_input":"2022-08-07T01:57:23.966001Z","iopub.status.idle":"2022-08-07T01:57:23.980008Z","shell.execute_reply.started":"2022-08-07T01:57:23.965521Z","shell.execute_reply":"2022-08-07T01:57:23.978767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = TrainDataset(train, transform=get_transforms(data='train'))\n\nimage_loader = DataLoader(train_dataset, \n                          batch_size  = 64, \n                          shuffle     = True, \n                          num_workers = 3,\n                          pin_memory  = True)\n\nfor images, _ in image_loader:\n    print('images.shape:', images.shape)\n    plt.figure(figsize=(16,8))\n    plt.axis('off')\n    plt.imshow(make_grid(images, nrow=16).permute((1, 2, 0)))\n    break","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:23.981216Z","iopub.execute_input":"2022-08-07T01:57:23.981511Z","iopub.status.idle":"2022-08-07T01:57:32.831642Z","shell.execute_reply.started":"2022-08-07T01:57:23.981481Z","shell.execute_reply":"2022-08-07T01:57:32.830766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class ArcMarginProduct(nn.Module):\n    def __init__(self, in_features, out_features, scale=30.0, margin=0.50, easy_margin=False, ls_eps=0.0):\n        super(ArcMarginProduct, self).__init__()\n        self.in_features = in_features\n        self.out_features = out_features\n        self.scale = scale\n        self.margin = margin\n        self.ls_eps = ls_eps  # label smoothing\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n\n        self.easy_margin = easy_margin\n        self.cos_m = math.cos(margin)\n        self.sin_m = math.sin(margin)\n        self.th = math.cos(math.pi - margin)\n        self.mm = math.sin(math.pi - margin) * margin\n\n    def forward(self, input, label):\n        # --------------------------- cos(theta) & phi(theta) ---------------------------\n        cosine = F.linear(F.normalize(input), F.normalize(self.weight))\n        sine = torch.sqrt(1.0 - torch.pow(cosine, 2))\n        phi = cosine * self.cos_m - sine * self.sin_m\n        if self.easy_margin:\n            phi = torch.where(cosine > 0, phi, cosine)\n        else:\n            phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n        # --------------------------- convert label to one-hot ---------------------------\n        # one_hot = torch.zeros(cosine.size(), requires_grad=True, device='cuda')\n        one_hot = torch.zeros(cosine.size(), device=device)\n        one_hot.scatter_(1, label.view(-1, 1).long(), 1)\n        if self.ls_eps > 0:\n            one_hot = (1 - self.ls_eps) * one_hot + self.ls_eps / self.out_features\n        # -------------torch.where(out_i = {x_i if condition_i else y_i) -------------\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        output *= self.scale\n\n        return output","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:32.832934Z","iopub.execute_input":"2022-08-07T01:57:32.834824Z","iopub.status.idle":"2022-08-07T01:57:32.850606Z","shell.execute_reply.started":"2022-08-07T01:57:32.834782Z","shell.execute_reply":"2022-08-07T01:57:32.849779Z"},"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 in ['tf_efficientnetv2_b0', 'tf_efficientnet_b5', 'tf_efficientnet_b2', 'tf_efficientnet_b0']:\n            self.in_features = self.model.classifier.in_features\n            self.model.classifier = nn.Identity()\n            self.model.global_pool = nn.Identity()\n            \n        if cfg.model_name in ['resnext50_32x4d', 'resnet50d', 'resnet34d']:\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.in_features = self.model.head.in_features\n            self.model.head = nn.Linear(self.in_features, self.cfg.fc_dim)\n        \n        \n        #self.pooling = GeM()\n        self.pooling =  nn.AdaptiveAvgPool2d(1) # GAP\n        self.probs = nn.Linear(self.in_features, self.cfg.target_size)\n        self.bn = nn.BatchNorm1d(self.in_features) # BNNeck\n        self.final = ArcMarginProduct(\n            self.in_features,\n            cfg.target_size,\n            scale = cfg.scale,\n            margin = cfg.margin,\n            easy_margin = False,\n            ls_eps = 0.0\n        )\n\n    def forward(self, x, label):\n        batch_size = x.shape[0]\n        # model backbone shape: torch.Size([4, 2048, 8, 8])\n        features = self.model(x)\n        # gap shape: torch.Size([4, 2048])\n        features = self.pooling(features).view(batch_size, -1)\n        arcface = self.final(features, label)\n        bn = self.bn(features)\n        probs   = self.probs(bn)\n        return probs, arcface","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:32.852853Z","iopub.execute_input":"2022-08-07T01:57:32.853807Z","iopub.status.idle":"2022-08-07T01:57:32.869955Z","shell.execute_reply.started":"2022-08-07T01:57:32.853770Z","shell.execute_reply":"2022-08-07T01:57:32.868961Z"},"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, labels) in enumerate(train_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        images = images.to(device).float()\n        labels = labels.to(device).long()\n        batch_size = labels.size(0)\n        if CFG.apex:\n            with autocast():\n                probs, arcface = model(images, labels)\n                arcface_loss = nn.CrossEntropyLoss()(arcface, labels)\n                loss = criterion(probs, labels)\n        else:\n            probs, arcface = model(images, labels)\n            arcface_loss = nn.CrossEntropyLoss()(arcface, labels)\n            loss = criterion(probs, labels)\n        # record loss\n        sum_loss = loss + arcface_loss\n        losses.update(sum_loss.item(), batch_size)\n        if CFG.gradient_accumulation_steps > 1:\n            sum_loss = sum_loss / CFG.gradient_accumulation_steps\n        if CFG.apex:\n            scaler.scale(sum_loss).backward()\n        else:\n            sum_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, labels) in enumerate(valid_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        images = images.to(device).float()\n        labels = labels.to(device).long()\n        batch_size = labels.size(0)\n        # compute loss\n        with torch.no_grad():\n            probs, _ = model(images, labels)\n        preds.append(probs.softmax(1).to('cpu').numpy())\n        loss = criterion(probs, 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-08-07T01:57:32.871740Z","iopub.execute_input":"2022-08-07T01:57:32.872167Z","iopub.status.idle":"2022-08-07T01:57:32.906683Z","shell.execute_reply.started":"2022-08-07T01:57:32.872053Z","shell.execute_reply":"2022-08-07T01:57:32.905658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train loop","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Train loop\n# ====================================================\ndef train_loop(folds, fold):\n    \n    LOGGER.info(f\"========== fold: {fold} training ==========\")\n\n    # ====================================================\n    # loader\n    # ====================================================\n    trn_idx = folds[folds['folds'] != fold].index\n    val_idx = folds[folds['folds'] == fold].index\n\n    train_folds = folds.loc[trn_idx].reset_index(drop=True)\n    valid_folds = folds.loc[val_idx].reset_index(drop=True)\n    valid_labels = valid_folds[\"label\"].values\n\n    train_dataset = TrainDataset(train_folds, transform=get_transforms(data='train'))\n    valid_dataset = TrainDataset(valid_folds, transform=get_transforms(data='valid'))\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    valid_loader = DataLoader(valid_dataset, \n                              batch_size=CFG.batch_size, \n                              shuffle=False, \n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\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    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_score = 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        score = get_score(valid_labels, 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        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}] score\": score})\n\n            \n        if avg_val_loss < best_loss:\n            best_loss = avg_val_loss\n            LOGGER.info(f'Epoch {epoch+1} - Save Best Loss: {best_loss:.4f} Model')\n            torch.save({'model': model.state_dict(), \n                        'preds_loss': preds},\n                        OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best_loss.pth')\n            \n        if best_score < score:\n            best_score = score\n            LOGGER.info(f'Epoch {epoch+1} - Save Best score: {score:.4f} Model')\n            torch.save({'model': model.state_dict(), \n                        'preds_loss': preds},\n                        OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best_score.pth')\n        \n        \n   \n    valid_folds[CFG.preds_col] = torch.load(OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best_loss.pth', \n                                      map_location=torch.device('cpu'))['preds_loss']\n   \n\n    return valid_folds","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:32.911883Z","iopub.execute_input":"2022-08-07T01:57:32.915109Z","iopub.status.idle":"2022-08-07T01:57:32.937409Z","shell.execute_reply.started":"2022-08-07T01:57:32.915069Z","shell.execute_reply":"2022-08-07T01:57:32.936225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Main","metadata":{}},{"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, 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.to_csv(OUTPUT_DIR+f'{CFG.model_name}_oof_version{VERSION}.csv', index=False)\n        \n    wandb.finish()","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:32.938840Z","iopub.execute_input":"2022-08-07T01:57:32.939264Z","iopub.status.idle":"2022-08-07T01:57:32.952708Z","shell.execute_reply.started":"2022-08-07T01:57:32.939229Z","shell.execute_reply":"2022-08-07T01:57:32.951481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:32.954133Z","iopub.execute_input":"2022-08-07T01:57:32.954636Z","iopub.status.idle":"2022-08-07T01:57:33.147364Z","shell.execute_reply.started":"2022-08-07T01:57:32.954602Z","shell.execute_reply":"2022-08-07T01:57:33.146275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    main()","metadata":{"execution":{"iopub.status.busy":"2022-08-07T01:57:33.148739Z","iopub.execute_input":"2022-08-07T01:57:33.149153Z"},"trusted":true},"execution_count":null,"outputs":[]}]}