{"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 torch==1.10.1+cu111 torchvision==0.11.2+cu111 torchaudio==0.10.1 -f https://download.pytorch.org/whl/torch_stable.html\n!pip install timm # install pytorch image models\n!pip install torchmetrics","metadata":{"papermill":{"duration":17.090323,"end_time":"2022-05-01T17:18:26.18714","exception":false,"start_time":"2022-05-01T17:18:09.096817","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:35:08.026615Z","iopub.execute_input":"2022-05-12T08:35:08.026954Z","iopub.status.idle":"2022-05-12T08:35:27.951744Z","shell.execute_reply.started":"2022-05-12T08:35:08.026877Z","shell.execute_reply":"2022-05-12T08:35:27.950927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport os\nimport pandas as pd\nimport numpy as np\nimport random \n\nimport albumentations as A\nimport cv2\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom sklearn import preprocessing\nfrom sklearn.model_selection import StratifiedKFold\nimport timm\n\nimport torchvision\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport torchvision.models as models\nimport torch.nn.functional as F\nfrom torch import nn\nimport torchmetrics \nfrom torch.nn.modules.loss import _Loss","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":4.951238,"end_time":"2022-05-01T17:18:31.209509","exception":false,"start_time":"2022-05-01T17:18:26.258271","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:35:27.955388Z","iopub.execute_input":"2022-05-12T08:35:27.955619Z","iopub.status.idle":"2022-05-12T08:35:37.112817Z","shell.execute_reply.started":"2022-05-12T08:35:27.955589Z","shell.execute_reply":"2022-05-12T08:35:37.111984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!nvidia-smi","metadata":{"papermill":{"duration":0.74824,"end_time":"2022-05-01T17:18:32.001737","exception":false,"start_time":"2022-05-01T17:18:31.253497","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:35:37.114703Z","iopub.execute_input":"2022-05-12T08:35:37.114982Z","iopub.status.idle":"2022-05-12T08:35:37.827064Z","shell.execute_reply.started":"2022-05-12T08:35:37.114944Z","shell.execute_reply":"2022-05-12T08:35:37.826280Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GlobalConstantsConfigure():\n    def __init__(self):\n        self.continue_training = False\n        self.last_model = '../input/sorghum-cultivar-identification-models/eff_v2_sgd_50_silu.pt' \n        self.num_epochs_done = 0\n        self.seed = 107\n        self.fold = 1\n        self.num_folds = 4\n        self.num_classes = 100\n        self.biggest_loss = 999\n        self.training_size_rate = 0.8\n        self.training_dir = '../input/sorghum-id-fgvc-9/train_images'\n        self.model_name = 'tf_efficientnetv2_m_in21k'\n        self.model_path = './efficientnetv2_b5_sgd_50.pt'\n        self.image_size = 512\n        self.batch_size = 8\n        self.val_batch_size = 32\n        self.lr = 3e-5 # 3e-5\n        self.num_epochs = 25\n        self.steps_per_decay = 5\n        self.device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\n        self.num_workers = 2  # if torch.cuda.is_available() else 4\ngcc = GlobalConstantsConfigure()","metadata":{"papermill":{"duration":0.117496,"end_time":"2022-05-01T17:18:32.165302","exception":false,"start_time":"2022-05-01T17:18:32.047806","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:35:37.829685Z","iopub.execute_input":"2022-05-12T08:35:37.829899Z","iopub.status.idle":"2022-05-12T08:35:37.902106Z","shell.execute_reply.started":"2022-05-12T08:35:37.829871Z","shell.execute_reply":"2022-05-12T08:35:37.901097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed) : \n    os.environ['PYTHONHASHSEED'] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\nset_seed(gcc.seed)","metadata":{"papermill":{"duration":0.053876,"end_time":"2022-05-01T17:18:32.263641","exception":false,"start_time":"2022-05-01T17:18:32.209765","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:35:37.903911Z","iopub.execute_input":"2022-05-12T08:35:37.904410Z","iopub.status.idle":"2022-05-12T08:35:37.914885Z","shell.execute_reply.started":"2022-05-12T08:35:37.904373Z","shell.execute_reply":"2022-05-12T08:35:37.914173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_all = pd.read_csv('../input/sorghum-id-fgvc-9/train_cultivar_mapping.csv')\nprint(len(df_all))\ndf_all.dropna(inplace=True)\nprint(len(df_all))\ndf_all.head()","metadata":{"papermill":{"duration":0.135029,"end_time":"2022-05-01T17:18:33.266092","exception":false,"start_time":"2022-05-01T17:18:33.131063","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:35:37.916184Z","iopub.execute_input":"2022-05-12T08:35:37.916654Z","iopub.status.idle":"2022-05-12T08:35:37.976260Z","shell.execute_reply.started":"2022-05-12T08:35:37.916618Z","shell.execute_reply":"2022-05-12T08:35:37.975559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"unique_cultivars = list(df_all[\"cultivar\"].unique())","metadata":{"papermill":{"duration":0.054091,"end_time":"2022-05-01T17:18:33.361686","exception":false,"start_time":"2022-05-01T17:18:33.307595","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:35:37.977598Z","iopub.execute_input":"2022-05-12T08:35:37.978052Z","iopub.status.idle":"2022-05-12T08:35:37.989171Z","shell.execute_reply.started":"2022-05-12T08:35:37.978017Z","shell.execute_reply":"2022-05-12T08:35:37.988515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_all[\"file_path\"] = df_all[\"image\"].apply(lambda image: '../input/sorghum-id-fgvc-9/train_images/' + image)\ndf_all[\"cultivar_index\"] = df_all[\"cultivar\"].map(lambda item: unique_cultivars.index(item))\ndf_all[\"is_exist\"] = df_all[\"file_path\"].apply(lambda file_path: os.path.exists(file_path))\ndf_all = df_all[df_all.is_exist==True]\ndf_all.head()","metadata":{"papermill":{"duration":13.655239,"end_time":"2022-05-01T17:18:47.059218","exception":false,"start_time":"2022-05-01T17:18:33.403979","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:35:37.990460Z","iopub.execute_input":"2022-05-12T08:35:37.990730Z","iopub.status.idle":"2022-05-12T08:35:53.881636Z","shell.execute_reply.started":"2022-05-12T08:35:37.990694Z","shell.execute_reply":"2022-05-12T08:35:53.880914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"skf = StratifiedKFold(n_splits=gcc.num_folds, shuffle=True, random_state=gcc.seed)\n\n\ntrain_folds = []\nval_folds = []\n\nfor train_idx, valid_idx in skf.split(df_all['image'], df_all[\"cultivar_index\"]):\n    train_folds.append(train_idx)\n    val_folds.append(valid_idx)\n#     df_train = df_all.iloc[train_idx]\n#     df_valid = df_all.iloc[valid_idx]\n\n# print(train_folds)\n# print(val_folds)\ndf_train = df_all.iloc[train_folds[gcc.fold]]\ndf_valid = df_all.iloc[val_folds[gcc.fold]]\n\n\n\n\nprint(f\"train size: {len(df_train)}\")\nprint(f\"valid size: {len(df_valid)}\")\n\nprint(df_train.cultivar.value_counts())\nprint(df_valid.cultivar.value_counts())","metadata":{"papermill":{"duration":0.074622,"end_time":"2022-05-01T17:18:47.176438","exception":false,"start_time":"2022-05-01T17:18:47.101816","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:35:53.882755Z","iopub.execute_input":"2022-05-12T08:35:53.883585Z","iopub.status.idle":"2022-05-12T08:35:53.913452Z","shell.execute_reply.started":"2022-05-12T08:35:53.883545Z","shell.execute_reply":"2022-05-12T08:35:53.912800Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# a = ","metadata":{"execution":{"iopub.status.busy":"2022-05-12T08:35:53.916585Z","iopub.execute_input":"2022-05-12T08:35:53.916788Z","iopub.status.idle":"2022-05-12T08:35:53.922583Z","shell.execute_reply.started":"2022-05-12T08:35:53.916758Z","shell.execute_reply":"2022-05-12T08:35:53.921875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SorghumDataset(Dataset):\n    def __init__(self, dirs, labels, transformation=None):\n        super(SorghumDataset,self).__init__()\n        self.dirs = dirs\n        self.labels = labels\n        self.transformation = transformation\n    def __len__(self):\n        return len(self.dirs)\n\n    def __getitem__(self, index):\n        image = cv2.imread(self.dirs[index])\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        label = self.labels[index] # need to one hot encoding here\n        \n        image = np.array(image)\n\n        if self.transformation:\n            aug_image = self.transformation(image=image)\n            image = aug_image['image']\n            \n        image = image / 255.\n        image = image.transpose((2, 0, 1))\n        \n        image = torch.from_numpy(image).type(torch.float32)\n        image = transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))(image)\n        \n        labels = torch.from_numpy(np.array(self.labels[index])).type(torch.float32)\n\n\n        return image, labels","metadata":{"papermill":{"duration":0.056726,"end_time":"2022-05-01T17:18:48.005809","exception":false,"start_time":"2022-05-01T17:18:47.949083","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:35:53.924963Z","iopub.execute_input":"2022-05-12T08:35:53.925639Z","iopub.status.idle":"2022-05-12T08:35:53.935162Z","shell.execute_reply.started":"2022-05-12T08:35:53.925597Z","shell.execute_reply":"2022-05-12T08:35:53.934387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_transformation = A.Compose([\n    A.Resize(width=gcc.image_size, height=gcc.image_size, p=1.0),\n    A.Flip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.ShiftScaleRotate(p=0.5),\n    A.HueSaturationValue(p=0.5),\n#     A.OneOf([\n#         A.RandomBrightnessContrast(p=0.5),\n#         A.RandomGamma(p=0.5),\n#     ], p=0.5),\n#     A.OneOf([\n#         A.Blur(p=0.1),\n#         A.GaussianBlur(p=0.1),\n#         A.MotionBlur(p=0.1),\n#     ], p=0.1),\n#     A.OneOf([\n#         A.GaussNoise(p=0.1),\n#         A.ISONoise(p=0.1),\n#         A.GridDropout(ratio=0.5, p=0.2),\n#         A.CoarseDropout(max_holes=16, min_holes=8, max_height=16, max_width=16, min_height=8, min_width=8, p=0.2)\n#     ], p=0.2),\n\n])\nvalidation_transformation = A.Compose([\n    A.Resize(width=gcc.image_size, height=gcc.image_size, p=1.0)\n])","metadata":{"papermill":{"duration":0.053572,"end_time":"2022-05-01T17:18:48.102106","exception":false,"start_time":"2022-05-01T17:18:48.048534","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:35:53.936168Z","iopub.execute_input":"2022-05-12T08:35:53.938260Z","iopub.status.idle":"2022-05-12T08:35:53.946764Z","shell.execute_reply.started":"2022-05-12T08:35:53.938216Z","shell.execute_reply":"2022-05-12T08:35:53.946028Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# training_set = SorghumDataset(training_dirs, training_labels, training_transformation)\n# validation_set = SorghumDataset(validation_dirs, validation_labels, validation_transformation)\n\ntraining_set = SorghumDataset(df_train.file_path.values, df_train.cultivar_index.values, training_transformation)\nvalidation_set = SorghumDataset(df_valid.file_path.values, df_valid.cultivar_index.values, validation_transformation)\n\n\ntraining_dataloader = DataLoader(\n    training_set,\n    batch_size = gcc.batch_size,\n    shuffle = True,\n    num_workers = gcc.num_workers,\n    pin_memory = True, \n    drop_last = True\n)\nvalidation_dataloader = DataLoader(\n    validation_set,\n    batch_size = gcc.val_batch_size,\n    shuffle = True,\n    num_workers = gcc.num_workers,\n    pin_memory = True,\n    drop_last = True\n)","metadata":{"papermill":{"duration":0.050942,"end_time":"2022-05-01T17:18:48.195564","exception":false,"start_time":"2022-05-01T17:18:48.144622","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:35:53.949059Z","iopub.execute_input":"2022-05-12T08:35:53.949380Z","iopub.status.idle":"2022-05-12T08:35:53.957347Z","shell.execute_reply.started":"2022-05-12T08:35:53.949348Z","shell.execute_reply":"2022-05-12T08:35:53.956617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomModel(torch.nn.Module): \n    def __init__(self, model_backbone):\n        super(CustomModel,self).__init__()\n        self.model = model_backbone\n        \n        self.model.classifier = nn.Sequential(\n            nn.BatchNorm1d(1280),\n            nn.Linear(1280, 512),\n            nn.Dropout(0.5),\n            nn.ReLU(inplace=True),\n            nn.Linear(512, gcc.num_classes),\n            \n#             nn.BatchNorm1d(1280),\n#             nn.Linear(1280, 512),\n#             nn.Dropout(0.5),\n#             nn.SiLU(inplace=True),\n\n#             nn.Linear(512, 256),\n#             nn.Dropout(0.5),\n#             nn.SiLU(inplace=True),\n#             nn.Linear(256, gcc.num_classes)\n        )\n    def forward(self,x):\n        x = self.model(x)\n        return x\n","metadata":{"papermill":{"duration":0.051162,"end_time":"2022-05-01T17:18:48.289325","exception":false,"start_time":"2022-05-01T17:18:48.238163","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:35:53.958968Z","iopub.execute_input":"2022-05-12T08:35:53.959355Z","iopub.status.idle":"2022-05-12T08:35:53.969984Z","shell.execute_reply.started":"2022-05-12T08:35:53.959298Z","shell.execute_reply":"2022-05-12T08:35:53.969301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def to_one_hot(labels, num_classes, dtype=torch.float, dim=1):\n    if labels.ndim < dim + 1:\n        shape = list(labels.shape) + [1] * (dim + 1 - len(labels.shape))\n        labels = torch.reshape(labels, shape)\n    sh = list(labels.shape)\n    sh[dim] = num_classes\n    o = torch.zeros(size=sh, dtype=dtype, device=labels.device)\n    labels = o.scatter_(dim=dim, index=labels.long(), value=1)\n    return labels\n\n\nclass PolyLoss(_Loss):\n    def __init__(self, softmax, ce_weight=None, reduction='mean', epsilon=1.0):\n        super().__init__()\n        self.softmax = softmax\n        self.reduction = reduction\n        self.epsilon = epsilon\n        self.cross_entropy = nn.CrossEntropyLoss(weight=ce_weight, reduction='none')\n\n    def forward(self, input, target):\n\n        if len(input.shape) - len(target.shape) == 1:\n            target = target.unsqueeze(1).long()\n        n_pred_ch, n_target_ch = input.shape[1], target.shape[1]\n        if n_pred_ch != n_target_ch:\n            self.ce_loss = self.cross_entropy(input, torch.squeeze(target, dim=1).long())\n            target = to_one_hot(target, num_classes=n_pred_ch)\n        else:\n            self.ce_loss = self.cross_entropy(input, torch.argmax(target, dim=1))\n\n        if self.softmax:\n            input = torch.softmax(input, 1)\n\n        pt = (input * target).sum(dim=1) \n        \n        poly_loss = self.ce_loss + self.epsilon * (1 - pt)\n\n        polyl = torch.mean(poly_loss)  # the batch and channel average\n        # polyl = torch.sum(poly_loss)  # sum over the batch and channel dims\n        return (polyl)","metadata":{"execution":{"iopub.status.busy":"2022-05-12T08:35:53.972857Z","iopub.execute_input":"2022-05-12T08:35:53.973446Z","iopub.status.idle":"2022-05-12T08:35:53.987349Z","shell.execute_reply.started":"2022-05-12T08:35:53.973415Z","shell.execute_reply":"2022-05-12T08:35:53.986600Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# backbone = models.efficientnet_b5(pretrained=True) \nbackbone = timm.create_model(gcc.model_name,pretrained=True)\n\n\n# print(index)\nmodel = CustomModel(backbone)\nce = torch.nn.CrossEntropyLoss()\nloss_func = PolyLoss(ce) # torch.nn.CrossEntropyLoss()\nmetrics_acc = torchmetrics.Accuracy(threshold=0.0, num_classes = gcc.num_classes)\nprint(model)\n\n\n# for index, child in enumerate(backbone.children()):\n#     print(index)\n#     if index <= 7:\n#         for param in child.parameters():\n#             param.requires_grad = False\n\n\ntrainable_parameters = [param for param in model.parameters() if param.requires_grad == True]\noptimizer = torch.optim.Adam(trainable_parameters, lr = gcc.lr)\n# optimizer = torch.optim.SGD(trainable_parameters, lr = gcc.lr, momentum = 0.9)\n# optimizer = torch.optim.SGD(trainable_parameters, lr = gcc.lr, momentum=0.9)\n# lr_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, gcc.steps_per_decay)\nlr_scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=5, T_mult=2)\n\nmodel.to(gcc.device)\n\nif gcc.continue_training == True:\n    # model.load_state_dict(torch.load(gcc.last_model))\n    checkpoint = torch.load(gcc.last_model)\n    \n    model.load_state_dict(checkpoint['model_state_dict'])\n    \n    \n    # optimizer = torch.optim.Adam(trainable_parameters, lr = gcc.lr)\n    # optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n    \n    # lr_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, gcc.steps_per_decay)\n    # lr_scheduler.load_state_dict(checkpoint['scheduler_state_dict'])\n    \n    # print(lr_scheduler.state_dict())\n\nprint('load model done')","metadata":{"papermill":{"duration":20.068797,"end_time":"2022-05-01T17:19:08.400854","exception":false,"start_time":"2022-05-01T17:18:48.332057","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:35:53.990245Z","iopub.execute_input":"2022-05-12T08:35:53.990775Z","iopub.status.idle":"2022-05-12T08:36:09.549761Z","shell.execute_reply.started":"2022-05-12T08:35:53.990736Z","shell.execute_reply":"2022-05-12T08:36:09.548845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def calc_accuracy(pred, true):\n#     # print(pred, true)\n#     true = true.type(torch.int64) # label\n#     pred = F.softmax(pred, dim = 1)\n#     true = torch.zeros(pred.shape[0], pred.shape[1]).scatter_(1, true.unsqueeze(1), 1.)\n#     acc = (true.argmax(-1) == pred.argmax(-1)).float().detach().numpy()\n#     acc = float(acc.sum() / len(acc))\n#     return round(acc, 4)","metadata":{"papermill":{"duration":0.053574,"end_time":"2022-05-01T17:19:08.500047","exception":false,"start_time":"2022-05-01T17:19:08.446473","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:36:09.551030Z","iopub.execute_input":"2022-05-12T08:36:09.551306Z","iopub.status.idle":"2022-05-12T08:36:09.556523Z","shell.execute_reply.started":"2022-05-12T08:36:09.551278Z","shell.execute_reply":"2022-05-12T08:36:09.555804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def training_progress(training_dataloader, loss_func, scheduler):\n    model.train()\n    training_loss = 0\n    training_acc = 0\n    cnt = 0 \n    print('Learning rate: ',scheduler.get_last_lr())\n    print(scheduler.state_dict())\n    training_loader = tqdm(training_dataloader, desc='Iterating through the training set')\n    for image, label in training_loader:\n        image = image.to(gcc.device)\n        label = label.to(gcc.device)\n        \n        output = model(image)\n        # output.to(gcc.device)\n\n        acc = metrics_acc(output.cpu().argmax(1), label.cpu().int())\n        loss = loss_func(output, label.long())\n        # calculate accuracy here\n\n        training_loss += loss.item()\n        training_acc += acc\n        cnt +=1 \n        \n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        scheduler.step()\n    \n    mean_training_loss = training_loss / cnt\n    mean_training_acc = training_acc / cnt\n    \n    return mean_training_loss, mean_training_acc\n    ","metadata":{"papermill":{"duration":0.054778,"end_time":"2022-05-01T17:19:08.599762","exception":false,"start_time":"2022-05-01T17:19:08.544984","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:36:09.558889Z","iopub.execute_input":"2022-05-12T08:36:09.559284Z","iopub.status.idle":"2022-05-12T08:36:10.566525Z","shell.execute_reply.started":"2022-05-12T08:36:09.559249Z","shell.execute_reply":"2022-05-12T08:36:10.565645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def validation_progress(validation_dataloader, loss_func):\n    model.eval()\n    validation_loss = 0\n    validation_acc = 0\n    cnt = 0 \n    validation_loader = tqdm(validation_dataloader, desc='Iterating through the validation set')\n    with torch.no_grad():\n        for image, label in validation_loader:\n            image = image.to(gcc.device)\n            label = label.to(gcc.device)\n\n            output = model(image)\n            loss = loss_func(output, label.long())\n            # acc = calc_accuracy(output.cpu(), label.cpu())\n            # output.to(gcc.device)\n            acc = metrics_acc(output.cpu().argmax(1), label.cpu().int())\n            # calculate accuracy here\n            validation_loss += loss.item()\n            validation_acc += acc\n            \n            cnt += 1\n\n    mean_validation_loss = validation_loss / cnt\n    mean_validation_acc = validation_acc / cnt\n    return mean_validation_loss, mean_validation_acc","metadata":{"papermill":{"duration":0.055002,"end_time":"2022-05-01T17:19:08.709353","exception":false,"start_time":"2022-05-01T17:19:08.654351","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:36:10.570904Z","iopub.execute_input":"2022-05-12T08:36:10.572492Z","iopub.status.idle":"2022-05-12T08:36:10.586370Z","shell.execute_reply.started":"2022-05-12T08:36:10.572449Z","shell.execute_reply":"2022-05-12T08:36:10.585424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def training_model(model, training_dataloader, validation_dataloader, loss_func, scheduler):\n    training_losses_history, validation_losses_history = [], []\n    training_acc_history, validation_acc_history = [], []\n    best_loss = gcc.biggest_loss\n    for epoch in range(gcc.num_epochs):\n        \n        training_loss, training_acc = training_progress(training_dataloader, loss_func, scheduler)\n        training_losses_history.append(training_loss)\n        training_acc_history.append(training_acc)\n        \n        validation_loss, validation_acc = validation_progress(validation_dataloader, loss_func)\n        validation_losses_history.append(validation_loss)\n        validation_acc_history.append(validation_acc)\n        \n        if validation_loss <= best_loss: # sussy baka\n            best_loss = validation_loss\n            torch.save({\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'scheduler_state_dict': scheduler.state_dict()\n            }, gcc.model_name + '_best.pt')\n            # torch.save(model.state_dict(), gcc.model_name + '_best.pt')\n        \n        if epoch == gcc.num_epochs - 1: # i believe my timing capability\n            torch.save({\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'scheduler_state_dict': scheduler.state_dict()\n            }, gcc.model_name + '_' + str(gcc.num_epochs_done + gcc.num_epochs) + '_last.pt')\n\n        print(f'Epoch {epoch + 1}/{gcc.num_epochs} | Training_loss : {training_loss:.3f} | Validation_loss : {validation_loss:.3f}' \n             + f' Training_acc : {training_acc:.3f} | Validation_acc : {validation_acc:.3f}'\n             )\n    return training_losses_history, validation_losses_history, training_acc_history, validation_acc_history\n","metadata":{"papermill":{"duration":0.056086,"end_time":"2022-05-01T17:19:08.810488","exception":false,"start_time":"2022-05-01T17:19:08.754402","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:36:10.591017Z","iopub.execute_input":"2022-05-12T08:36:10.591428Z","iopub.status.idle":"2022-05-12T08:36:10.607728Z","shell.execute_reply.started":"2022-05-12T08:36:10.591392Z","shell.execute_reply":"2022-05-12T08:36:10.606811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_losses_history, validation_losses_history, training_acc_history, validation_acc_history = training_model(model, training_dataloader, validation_dataloader, loss_func, lr_scheduler)","metadata":{"papermill":{"duration":19927.831381,"end_time":"2022-05-01T22:51:16.69062","exception":false,"start_time":"2022-05-01T17:19:08.859239","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-05-12T08:36:10.610407Z","iopub.execute_input":"2022-05-12T08:36:10.613058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# testing_progress(testing_dataloader, loss_func)","metadata":{"papermill":{"duration":2.266744,"end_time":"2022-05-01T22:51:21.448697","exception":false,"start_time":"2022-05-01T22:51:19.181953","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_loss_history(model_name, train_loss_history, val_loss_history, num_epochs):\n    \n    x = np.arange(num_epochs)\n    fig = plt.figure(figsize=(10, 6))\n    plt.plot(x, train_loss_history, label='Train Loss', lw=3)\n    plt.plot(x, val_loss_history, label='Validation Loss', lw=3)\n\n    plt.title(f\"{model_name}\", fontsize=20)\n    plt.legend(fontsize=12)\n    plt.xlabel(\"Epoch\", fontsize=15)\n    plt.ylabel(\"Loss\", fontsize=15)\n\n    plt.show()\n    \nplot_loss_history(gcc.model_name, training_losses_history, validation_losses_history, gcc.num_epochs)","metadata":{"papermill":{"duration":2.48219,"end_time":"2022-05-01T22:51:26.132614","exception":false,"start_time":"2022-05-01T22:51:23.650424","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_acc_history(model_name, train_acc_history, val_acc_history, num_epochs):\n    \n    x = np.arange(num_epochs)\n    fig = plt.figure(figsize=(10, 6))\n    plt.plot(x, train_acc_history, label='Training Accuracy', lw=3)\n    plt.plot(x, val_acc_history, label='Validation Accuracy', lw=3)\n\n    plt.title(f\"{model_name}\", fontsize=20)\n    plt.legend(fontsize=12)\n    plt.xlabel(\"Epoch\", fontsize=15)\n    plt.ylabel(\"Accuracy\", fontsize=15)\n\n    plt.show()\n    \nplot_acc_history(gcc.model_name, training_acc_history, validation_acc_history, gcc.num_epochs)","metadata":{"papermill":{"duration":2.941261,"end_time":"2022-05-01T22:51:31.331167","exception":false,"start_time":"2022-05-01T22:51:28.389906","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint = torch.load(gcc.model_name + '_best.pt')\nmodel.load_state_dict(checkpoint['model_state_dict'])","metadata":{"papermill":{"duration":2.650889,"end_time":"2022-05-01T22:51:36.32067","exception":false,"start_time":"2022-05-01T22:51:33.669781","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.read_csv('../input/sorghum-id-fgvc-9/sample_submission.csv')\nsub.head()","metadata":{"papermill":{"duration":2.519486,"end_time":"2022-05-01T22:51:41.159131","exception":false,"start_time":"2022-05-01T22:51:38.639645","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub[\"filename\"] = sub[\"filename\"].apply(lambda image: '../input/sorghum-id-fgvc-9/test/' + image)\nsub[\"cultivar\"] = 0\nsub.head()","metadata":{"papermill":{"duration":2.214654,"end_time":"2022-05-01T22:51:45.585691","exception":false,"start_time":"2022-05-01T22:51:43.371037","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testing_dataset = SorghumDataset(sub['filename'], sub['cultivar'], validation_transformation)\ntesting_dataloader = DataLoader(testing_dataset, \n                                batch_size=gcc.val_batch_size, \n                                shuffle=False, \n                                num_workers=gcc.num_workers)","metadata":{"papermill":{"duration":2.341576,"end_time":"2022-05-01T22:51:50.13553","exception":false,"start_time":"2022-05-01T22:51:47.793954","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# predictions = np.zeros(len(testing_dataloader))\npredictions = []\ncnt = 0 \nwith torch.no_grad():\n    for image, label in tqdm(testing_dataloader):\n        image = image.to(gcc.device)\n        outputs = model(image)\n        # print(outputs)\n        preds = outputs.detach().cpu()\n        predictions.append(preds.argmax(1)) # need optimize here\n        # print(predictions)","metadata":{"papermill":{"duration":1233.351127,"end_time":"2022-05-01T23:12:26.04408","exception":false,"start_time":"2022-05-01T22:51:52.692953","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tmp = predictions[0]\nfor i in range(len(predictions) - 1):\n    tmp = torch.cat((tmp, predictions[i+1]))","metadata":{"papermill":{"duration":2.393549,"end_time":"2022-05-01T23:12:30.884395","exception":false,"start_time":"2022-05-01T23:12:28.490846","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# predictions = label_encoder.inverse_transform(tmp)\npredictions = [unique_cultivars[pred] for pred in tmp]","metadata":{"papermill":{"duration":2.84231,"end_time":"2022-05-01T23:12:36.101424","exception":false,"start_time":"2022-05-01T23:12:33.259114","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.read_csv('../input/sorghum-id-fgvc-9/sample_submission.csv')\nsub['cultivar'] = predictions\nsub.to_csv('submission.csv', index=False)\nsub.head()","metadata":{"papermill":{"duration":2.514665,"end_time":"2022-05-01T23:12:41.029383","exception":false,"start_time":"2022-05-01T23:12:38.514718","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]}]}