{"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":"#### ![](https://storage.googleapis.com/kaggle-competitions/kaggle/37333/logos/header.png)\n\n<h1 style=\"text-align: center; font-family: Verdana; font-size: 32px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; font-variant: small-caps; letter-spacing: 3px; color: #6d3e75; background-color: #ffffff;\">Mayo Clinic - STRIP AI</h1>\n\n### Goal of the Competition\nThe goal of this competition is to classify the blood clot origins in ischemic stroke. Using whole slide digital pathology images, you'll build a model that differentiates between the two major acute ischemic stroke (AIS) etiology subtypes: cardiac and large artery atherosclerosis.\n\nYour work will enable healthcare providers to better identify the origins of blood clots in deadly strokes, making it easier for physicians to prescribe the best post-stroke therapeutic management and reducing the likelihood of a second stroke.","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"# Importing Libraries & Data","metadata":{}},{"cell_type":"code","source":"!conda install ../input/how-to-use-pyvips-offline/*.tar.bz2 #installing pyviz\n# !pip download torchvision torch -d ./packages/torch/\n# !pip install -U torch torchvision","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-09-27T05:49:05.40081Z","iopub.execute_input":"2022-09-27T05:49:05.40118Z","iopub.status.idle":"2022-09-27T05:49:43.880118Z","shell.execute_reply.started":"2022-09-27T05:49:05.401098Z","shell.execute_reply":"2022-09-27T05:49:43.878914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import copy\nimport datetime\nimport gc\nimport os\nimport sys\nimport time\nimport warnings\n\nimport cv2\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport plotly.express as px\nimport pyvips\nimport seaborn as sns\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom PIL import Image\nfrom tqdm import tqdm\nfrom torchvision import models, transforms\nfrom easydict import EasyDict\nfrom IPython.core.display import HTML, display\nfrom sklearn.model_selection import train_test_split\nfrom torch.optim import lr_scheduler\nfrom torch.nn.functional import softmax\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\n\n# Suppress warnings\nwarnings.filterwarnings(\"ignore\")\n# For descriptive error messages\nos.environ['CUDA_LAUNCH_BLOCKING'] = \"1\"\nplt.style.use('ggplot')\n","metadata":{"execution":{"iopub.status.busy":"2022-09-27T05:49:52.700348Z","iopub.execute_input":"2022-09-27T05:49:52.701044Z","iopub.status.idle":"2022-09-27T05:49:52.710159Z","shell.execute_reply.started":"2022-09-27T05:49:52.701007Z","shell.execute_reply":"2022-09-27T05:49:52.709038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_PATH = \"../input/mayo-clinic-strip-ai\"\nTRAIN_PATH = \"../input/jpg-images-strip-ai/train\"\nTEST_PATH = \"../input/jpg-images-strip-ai/test\"","metadata":{"execution":{"iopub.status.busy":"2022-09-27T05:49:55.141276Z","iopub.execute_input":"2022-09-27T05:49:55.141663Z","iopub.status.idle":"2022-09-27T05:49:55.148608Z","shell.execute_reply.started":"2022-09-27T05:49:55.141619Z","shell.execute_reply":"2022-09-27T05:49:55.1476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv('../input/mayo-clinic-strip-ai/train.csv')\ntest_df = pd.read_csv('../input/mayo-clinic-strip-ai/test.csv')\nother_df = pd.read_csv('../input/mayo-clinic-strip-ai/other.csv')","metadata":{"execution":{"iopub.status.busy":"2022-09-27T05:49:55.930938Z","iopub.execute_input":"2022-09-27T05:49:55.931471Z","iopub.status.idle":"2022-09-27T05:49:55.972968Z","shell.execute_reply.started":"2022-09-27T05:49:55.931427Z","shell.execute_reply":"2022-09-27T05:49:55.971964Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(train_df.shape)\ntrain_df.head(10)","metadata":{"execution":{"iopub.status.busy":"2022-09-27T05:49:58.143648Z","iopub.execute_input":"2022-09-27T05:49:58.144679Z","iopub.status.idle":"2022-09-27T05:49:58.165742Z","shell.execute_reply.started":"2022-09-27T05:49:58.144606Z","shell.execute_reply":"2022-09-27T05:49:58.164719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(test_df.shape)\ntest_df.head(5)","metadata":{"execution":{"iopub.status.busy":"2022-09-27T05:50:00.255296Z","iopub.execute_input":"2022-09-27T05:50:00.255865Z","iopub.status.idle":"2022-09-27T05:50:00.277025Z","shell.execute_reply.started":"2022-09-27T05:50:00.255812Z","shell.execute_reply":"2022-09-27T05:50:00.276066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(other_df.shape)\nother_df.head(10)","metadata":{"execution":{"iopub.status.busy":"2022-09-27T05:50:03.073428Z","iopub.execute_input":"2022-09-27T05:50:03.073813Z","iopub.status.idle":"2022-09-27T05:50:03.090744Z","shell.execute_reply.started":"2022-09-27T05:50:03.07378Z","shell.execute_reply":"2022-09-27T05:50:03.089524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EDA","metadata":{}},{"cell_type":"code","source":"train_df.describe()","metadata":{"execution":{"iopub.status.busy":"2022-09-27T05:50:07.211584Z","iopub.execute_input":"2022-09-27T05:50:07.212574Z","iopub.status.idle":"2022-09-27T05:50:07.242574Z","shell.execute_reply.started":"2022-09-27T05:50:07.212536Z","shell.execute_reply":"2022-09-27T05:50:07.241549Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.groupby('label')['label'].count()","metadata":{"execution":{"iopub.status.busy":"2022-09-27T05:50:13.888779Z","iopub.execute_input":"2022-09-27T05:50:13.889139Z","iopub.status.idle":"2022-09-27T05:50:13.901238Z","shell.execute_reply.started":"2022-09-27T05:50:13.88911Z","shell.execute_reply":"2022-09-27T05:50:13.90024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.groupby(\"center_id\")['center_id'].count()","metadata":{"execution":{"iopub.status.busy":"2022-09-27T05:50:16.850974Z","iopub.execute_input":"2022-09-27T05:50:16.851368Z","iopub.status.idle":"2022-09-27T05:50:16.862595Z","shell.execute_reply.started":"2022-09-27T05:50:16.851333Z","shell.execute_reply":"2022-09-27T05:50:16.861249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = train_df.groupby(\"label\")[\"label\"].count().div(len(train_df)).mul(100)\ncenters = train_df.groupby(\"center_id\")[\"center_id\"].count().div(len(train_df)).mul(100)\n\nfig, ax = plt.subplots(1, 2, figsize=(18, 6))\nsns.barplot(x=labels.index, y=labels.values, ax=ax[0])\nax[0].set_title(\"Distribution of a target variable\"), ax[0].set_ylabel(\"%\")\nsns.barplot(x=centers.index, y=centers.values, ax=ax[1])\nax[1].set_title(\"Patients per clinic center\"), ax[1].set_ylabel(\"%\")\nplt.show()\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-09-27T05:50:21.199162Z","iopub.execute_input":"2022-09-27T05:50:21.199872Z","iopub.status.idle":"2022-09-27T05:50:21.628951Z","shell.execute_reply.started":"2022-09-27T05:50:21.199824Z","shell.execute_reply":"2022-09-27T05:50:21.626857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = (train_df.groupby([\"patient_id\",\"label\"])[\"image_num\"].count().reset_index(name='image_count'))\ndf[df[\"image_count\"]>2].set_index(\"patient_id\").T.style.background_gradient(cmap='Reds')","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-09-27T05:50:24.997482Z","iopub.execute_input":"2022-09-27T05:50:24.998265Z","iopub.status.idle":"2022-09-27T05:50:25.064566Z","shell.execute_reply.started":"2022-09-27T05:50:24.998217Z","shell.execute_reply":"2022-09-27T05:50:25.063632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Images","metadata":{}},{"cell_type":"code","source":"# Utility functions\nformat_to_dtype = {\n    \"uchar\": np.uint8,\n    \"char\": np.int8,\n    \"ushort\": np.uint16,\n    \"short\": np.int16,\n    \"uint\": np.uint32,\n    \"int\": np.int32,\n    \"float\": np.float32,\n    \"double\": np.float64,\n    \"complex\": np.complex64,\n    \"dpcomplex\": np.complex128,\n}\n\n\ndef vips2numpy(vi):\n    return np.ndarray(\n        buffer=vi.write_to_memory(),\n        dtype=format_to_dtype[vi.format],\n        shape=[vi.height, vi.width, vi.bands],\n    )\n\n\ndef read_image(image_id, dset, scale=None, verbose=1):\n    image = pyvips.Image.new_from_file(os.path.join(BASE_PATH, dset, f\"{image_id}.tif\"))\n    image = vips2numpy(image)\n    if verbose:\n        print(f\"[{image_id}] Image shape: {image.shape}\")\n\n    if scale:\n        new_size = (image.shape[1] // scale, image.shape[0] // scale)\n        image = cv2.resize(image, new_size, interpolation=cv2.INTER_AREA)\n        if image.shape[1] > 1.5 * image.shape[0]:\n            out = cv2.transpose(image)\n            image = cv2.flip(out, flipCode=0)\n\n        if verbose:\n            print(f\"[{image_id}] Resized Image shape: {image.shape}\")\n\n    return image\n\n\ndef plot_image(image, image_id):\n    plt.figure(figsize=(16, 10))\n    plt.imshow(image)\n    plt.title(f\"Image {image_id}\", fontsize=18)\n    plt.axis(\"off\")\n    plt.show()\n\n\ndef plot_list_img(sample_ids, dset, scale=20):\n    sample_images = []\n    for sample_id in sample_ids:\n        sample_images.append(read_image(sample_id, dset, scale=scale, verbose=0))\n    plt.figure(figsize=(16, 16))\n    for ind, (tmp_id, tmp_image) in enumerate(zip(sample_ids, sample_images)):\n        plt.subplot(2, 5, ind + 1)\n        plt.imshow(tmp_image)\n        plt.title(f\"{tmp_id}\", fontsize=10)\n        plt.axis(\"off\")\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-09-27T05:50:28.310922Z","iopub.execute_input":"2022-09-27T05:50:28.31345Z","iopub.status.idle":"2022-09-27T05:50:28.338402Z","shell.execute_reply.started":"2022-09-27T05:50:28.313409Z","shell.execute_reply":"2022-09-27T05:50:28.336504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_id = \"026c97_0\"\nimg = read_image(img_id, \"train\", scale=2)\nfig = px.imshow(img)\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-27T05:50:31.197512Z","iopub.execute_input":"2022-09-27T05:50:31.198025Z","iopub.status.idle":"2022-09-27T05:50:36.492303Z","shell.execute_reply.started":"2022-09-27T05:50:31.197988Z","shell.execute_reply":"2022-09-27T05:50:36.49095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Config","metadata":{}},{"cell_type":"code","source":"class TrainingConfigurations(EasyDict):\n    seed = 69\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    number_classes = 2\n    bs = 8\n    num_workers = 2\n    pin_memory = True\n    epochs = 10\n    optimizer = \"sgd\"\n    learning_rate = 1e-3\n    momentum = 0.92\n    weight_decay = 0\n    lr_scheduler = None\n    min_lr = 1e-2\n    lr_drop = None\n    T_max = 500\n    patience = 6\n    dropout = 0.\n    freeze = False\n    model_save_name = \"resnet_b6_best\"\n\n\ncfg = TrainingConfigurations()\ncfg","metadata":{"execution":{"iopub.status.busy":"2022-09-27T05:51:54.352458Z","iopub.execute_input":"2022-09-27T05:51:54.35285Z","iopub.status.idle":"2022-09-27T05:51:54.436876Z","shell.execute_reply.started":"2022-09-27T05:51:54.352817Z","shell.execute_reply":"2022-09-27T05:51:54.435603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed):\n#   for REPRODUCIBILITY.\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed(cfg.seed)","metadata":{"execution":{"iopub.status.busy":"2022-09-27T05:51:54.945293Z","iopub.execute_input":"2022-09-27T05:51:54.94571Z","iopub.status.idle":"2022-09-27T05:51:54.95486Z","shell.execute_reply.started":"2022-09-27T05:51:54.94567Z","shell.execute_reply":"2022-09-27T05:51:54.953596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = models.resnet34(pretrained=True)\nmodel","metadata":{"execution":{"iopub.status.busy":"2022-09-27T05:51:59.130407Z","iopub.execute_input":"2022-09-27T05:51:59.130834Z","iopub.status.idle":"2022-09-27T05:52:03.821604Z","shell.execute_reply.started":"2022-09-27T05:51:59.130797Z","shell.execute_reply":"2022-09-27T05:52:03.820645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"import torch.nn.functional as F\nclass Model(nn.Module):\n    def __init__(self, nc, dropout, freeze):\n        super(Model, self).__init__()\n\n        self.model = models.resnet34(pretrained=True)\n\n        num_ftrs = self.model.fc.in_features \n#         self.model.fc = nn.Sequential(\n#                     #nn.Linear(num_ftrs,500),\n# #                     nn.Dropout(0.1),\n#                     nn.BatchNorm2d(num_ftrs),\n#                     nn.ReLU(),\n#                     nn.Dropout(0.1),\n#                     nn.Linear(num_ftrs, nc)\n#                 )\n        self.batchnorm = nn.BatchNorm1d(1000)\n        self.dropout = nn.Dropout(0.1)\n        self.linear = nn.Linear(1000,2)\n\n    def forward(self, x):\n        out = self.model(x)\n        out = self.batchnorm(out)\n        out = F.relu(out)\n        out = self.dropout(out)\n        out = self.linear(out)\n        return out\n","metadata":{"execution":{"iopub.status.busy":"2022-09-27T08:22:12.577783Z","iopub.execute_input":"2022-09-27T08:22:12.578164Z","iopub.status.idle":"2022-09-27T08:22:12.586929Z","shell.execute_reply.started":"2022-09-27T08:22:12.578132Z","shell.execute_reply":"2022-09-27T08:22:12.585626Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset Class","metadata":{}},{"cell_type":"code","source":"class StripDataset(Dataset):\n    def __init__(self, df, dataset_path, augmentation=None):\n        self.df = df\n        self.dataset_path = dataset_path\n#         self.dset = dset\n        self.image_id = df['image_id'].values\n        self.patient_id = df['patient_id'].values\n        self.targets = df['label'].values\n        self.augs = augmentation\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        image_id = self.image_id[idx]\n        target = self.targets[idx]\n#         image = read_image(image_id, self.dset, scale=None, verbose=0)\n        image = Image.open(os.path.join(self.dataset_path, image_id + \".jpg\"))\n        \n        if self.augs:\n            image = self.augs(image)\n        target_map = {\n            \"CE\": 0,\n            \"LAA\": 1\n        }\n        return image, target_map[target]","metadata":{"execution":{"iopub.status.busy":"2022-09-27T08:22:12.882625Z","iopub.execute_input":"2022-09-27T08:22:12.883007Z","iopub.status.idle":"2022-09-27T08:22:12.893588Z","shell.execute_reply.started":"2022-09-27T08:22:12.882976Z","shell.execute_reply":"2022-09-27T08:22:12.892485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"augs = transforms.Compose([\n#     transforms.ToPILImage(),\n    transforms.Resize((480, 480)),\n    transforms.ToTensor(), \n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])","metadata":{"execution":{"iopub.status.busy":"2022-09-27T08:22:13.072413Z","iopub.execute_input":"2022-09-27T08:22:13.072731Z","iopub.status.idle":"2022-09-27T08:22:13.080074Z","shell.execute_reply.started":"2022-09-27T08:22:13.072702Z","shell.execute_reply":"2022-09-27T08:22:13.077145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataloader","metadata":{}},{"cell_type":"code","source":"df_train, df_test = train_test_split(\n    train_df, test_size=0.2, random_state=42, stratify=train_df.label)\n\ndf_train.shape, df_test.shape","metadata":{"execution":{"iopub.status.busy":"2022-09-27T08:22:13.791785Z","iopub.execute_input":"2022-09-27T08:22:13.792133Z","iopub.status.idle":"2022-09-27T08:22:13.805619Z","shell.execute_reply.started":"2022-09-27T08:22:13.792104Z","shell.execute_reply":"2022-09-27T08:22:13.804272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = StripDataset(df_train, TRAIN_PATH, augs)\ntest_dataset = StripDataset(df_test, TRAIN_PATH, augs)\n\ntrain_dataloader = DataLoader(\n    train_dataset,\n    batch_size=8,\n    num_workers=2,\n    shuffle=True,\n    pin_memory=True,\n)\ntest_dataloader = DataLoader(\n    test_dataset,\n    batch_size=8,\n    num_workers=2,\n    shuffle=False,\n    pin_memory=True,\n)","metadata":{"execution":{"iopub.status.busy":"2022-09-27T08:22:14.393076Z","iopub.execute_input":"2022-09-27T08:22:14.393472Z","iopub.status.idle":"2022-09-27T08:22:14.400942Z","shell.execute_reply.started":"2022-09-27T08:22:14.39344Z","shell.execute_reply":"2022-09-27T08:22:14.399871Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Begin Training","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, optimizer, scheduler, criterion, dataloader, device):\n    model.train()\n\n    total = 0\n    running_loss = 0\n    correct = 0\n    bar = tqdm(dataloader)\n    for batch, target in bar:\n        batch, target = batch.to(device), target.to(device)\n        batch_size = batch.size(0)\n\n        optimizer.zero_grad()\n        output = model(batch)\n        loss = criterion(output, target)\n\n        loss.backward()\n        optimizer.step()\n\n        if scheduler:\n            scheduler.step()\n\n        total += batch_size\n        running_loss += loss.item() * batch_size\n        pred = output.argmax(\n            dim=1, keepdim=True\n        )  # get the index of the max log-probability\n        correct += pred.eq(target.view_as(pred)).sum().item()\n\n        epoch_loss = running_loss / total\n        acc = correct / total\n\n        bar.set_postfix(\n            Loss=epoch_loss, Accuracy=acc * 100, LR=optimizer.param_groups[0][\"lr\"]\n        )\n\n    return {\"loss\": epoch_loss, \"accuracy\": acc}\n\n\n@torch.no_grad()\ndef evaluate_one_epoch(model, criterion, dataloader, device):\n    model.eval()\n\n    total = 0\n    running_loss = 0.0\n    correct = 0\n\n    for batch, target in dataloader:\n        batch, target = batch.to(device), target.to(device)\n        batch_size = batch.size(0)\n\n        output = model(batch)\n        loss = criterion(output, target)\n\n        total += batch_size\n        running_loss += loss.item() * batch_size\n        pred = output.argmax(\n            dim=1, keepdim=True\n        )  # get the index of the max log-probability\n        correct += pred.eq(target.view_as(pred)).sum().item()\n\n        epoch_loss = running_loss / total\n        acc = correct / total\n\n    epoch_loss = running_loss / total\n    acc = correct / total\n\n    print(\"Validation Loss: {:.4f} Accuracy: {:.2f}%\".format(epoch_loss, acc * 100))\n    return {\"loss\": epoch_loss, \"accuracy\": acc}","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-09-27T08:22:15.487847Z","iopub.execute_input":"2022-09-27T08:22:15.488206Z","iopub.status.idle":"2022-09-27T08:22:15.500892Z","shell.execute_reply.started":"2022-09-27T08:22:15.488175Z","shell.execute_reply":"2022-09-27T08:22:15.499897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def start_training(model, optimizer, scheduler, device, num_epochs):\n    st = time.time()\n    best_loss = np.inf\n    history = {\"train\": [], \"test\": []}\n    for epoch in range(1, num_epochs+ 1):\n        print(\"Epoch #{}\".format(epoch))\n\n        train_stats = train_one_epoch(\n            model,\n            optimizer,\n            lr_scheduler if cfg.lr_scheduler == \"cosine_annealing\" else None,\n            criterion,\n            train_dataloader,\n            device,\n        )\n\n        test_stats = evaluate_one_epoch(\n            model, criterion, test_dataloader, device)\n\n        if cfg.lr_scheduler and not cfg.lr_scheduler == \"cosine_annealing\":\n            lr_scheduler.step(\n                test_stats[\"loss\"]\n            ) if cfg.lr_scheduler == \"reducelronpleatue\" else lr_scheduler.step()\n\n        # saving best model\n        if test_stats[\"loss\"] < best_loss:\n            print(\n                \"Validation Loss Improved (%g ---> %g)\"\n                % (best_loss, test_stats[\"loss\"])\n            )\n            best_loss = test_stats[\"loss\"]\n            #torch.save(model.state_dict(), cfg.model_save_name + \".pt\")\n            traced = torch.jit.trace(model, torch.rand(1, 3, 512, 512).to(device))\n            traced.save('model.pth')\n            print(\"Model Saved\")\n        else:\n            print(\"Validation loss did not improve from- \", best_loss)\n\n        history[\"train\"].append(\n            {\n                \"epoch\": epoch,\n                \"accuracy\": train_stats[\"accuracy\"],\n                \"loss\": train_stats[\"loss\"],\n            }\n        )\n        history[\"test\"].append(\n            {\n                \"epoch\": epoch,\n                \"accuracy\": test_stats[\"accuracy\"],\n                \"loss\": test_stats[\"loss\"],\n            }\n        )\n\n        # end epoch ----------------------------------------------------------------------------------------------------\n\n    print(\"\\nTraining finished in\")\n    total_time = time.time() - st\n    total_time_str = str(datetime.timedelta(seconds=int(total_time)))\n    print(\"Training time {}\".format(total_time_str))\n    \n    return history ","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-09-27T08:22:16.104512Z","iopub.execute_input":"2022-09-27T08:22:16.105211Z","iopub.status.idle":"2022-09-27T08:22:16.116786Z","shell.execute_reply.started":"2022-09-27T08:22:16.105174Z","shell.execute_reply":"2022-09-27T08:22:16.115724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Model(cfg.number_classes, cfg.dropout, cfg.freeze)\nmodel.to(cfg.device)\n\nmodel_parameters = filter(lambda parameter: parameter.requires_grad, model.parameters())\n#optimizer = optim.Adam(model.parameters(),lr = 1e-5)\noptimizer = optim.SGD(model.parameters(), lr=1e-4, momentum=0.9) \nscheduler = None\ncriterion = nn.CrossEntropyLoss()\n\nhistory = start_training(\n    model, optimizer, scheduler, device=cfg['device'], num_epochs=20)\n","metadata":{"execution":{"iopub.status.busy":"2022-09-27T08:22:16.826177Z","iopub.execute_input":"2022-09-27T08:22:16.827293Z","iopub.status.idle":"2022-09-27T08:30:13.497285Z","shell.execute_reply.started":"2022-09-27T08:22:16.827248Z","shell.execute_reply":"2022-09-27T08:30:13.495533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_history(history):\n\n    plt.figure(figsize=(20,6))\n#     fig, axes = plt.subplots(1, 2, sharey=True, figsize=(22, 6))\n    \n    plt.subplot(1,2,1)\n    for k in [\"train\", \"test\"]:\n        plt.plot([i[\"loss\"] for i in history[k]])\n\n    plt.title('Loss')\n    plt.xlabel('epochs')\n    plt.ylabel('loss')\n    plt.legend(['train', 'valid'], loc='upper left')\n\n    plt.subplot(1,2,2)\n    for k in [\"train\", \"test\"]:\n        plt.plot([i[\"accuracy\"] for i in history[k]])\n\n    plt.title('Accuracy')\n    plt.xlabel('epochs')\n    plt.ylabel('accuracy')\n    plt.legend(['train', 'valid'], loc='upper left')\n    \n    plt.show()\n\n\n\nplot_history(history)","metadata":{"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp ~/.cache/torch/hub/checkpoints/* .","metadata":{"execution":{"iopub.status.busy":"2022-09-27T08:12:06.411344Z","iopub.status.idle":"2022-09-27T08:12:06.412092Z","shell.execute_reply.started":"2022-09-27T08:12:06.411832Z","shell.execute_reply":"2022-09-27T08:12:06.411856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}