{"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":"## Changelog\n\n#### v1\n- data prepration, eda\n\n#### v2\n- switching to pyvips for loading from tifffile, due to reduced ram and cpu usage\n- training config\n- base data class\n\n### v4\n- defining model efficientnet v2m\n- dataset class\n- first training \n\n### v5\n- downloading torch for installing offline from dataset\n\n### v8\n- first submission\n\n### v9\n- kernel died due to memory usage thats why will need to create another notebook for submission ","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-07T04:01:30.112267Z","iopub.execute_input":"2022-09-07T04:01:30.113287Z","iopub.status.idle":"2022-09-07T04:02:32.455864Z","shell.execute_reply.started":"2022-09-07T04:01:30.11324Z","shell.execute_reply":"2022-09-07T04:02:32.454482Z"},"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-07T04:02:32.45841Z","iopub.execute_input":"2022-09-07T04:02:32.458798Z","iopub.status.idle":"2022-09-07T04:02:37.609309Z","shell.execute_reply.started":"2022-09-07T04:02:32.458758Z","shell.execute_reply":"2022-09-07T04:02:37.608301Z"},"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-07T04:02:37.610578Z","iopub.execute_input":"2022-09-07T04:02:37.611311Z","iopub.status.idle":"2022-09-07T04:02:37.619639Z","shell.execute_reply.started":"2022-09-07T04:02:37.611268Z","shell.execute_reply":"2022-09-07T04:02:37.618733Z"},"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-07T04:02:37.622685Z","iopub.execute_input":"2022-09-07T04:02:37.623268Z","iopub.status.idle":"2022-09-07T04:02:37.671928Z","shell.execute_reply.started":"2022-09-07T04:02:37.623241Z","shell.execute_reply":"2022-09-07T04:02:37.671104Z"},"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-07T04:02:37.674428Z","iopub.execute_input":"2022-09-07T04:02:37.676421Z","iopub.status.idle":"2022-09-07T04:02:37.695181Z","shell.execute_reply.started":"2022-09-07T04:02:37.676394Z","shell.execute_reply":"2022-09-07T04:02:37.694123Z"},"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-07T04:02:37.697627Z","iopub.execute_input":"2022-09-07T04:02:37.698218Z","iopub.status.idle":"2022-09-07T04:02:37.70951Z","shell.execute_reply.started":"2022-09-07T04:02:37.698183Z","shell.execute_reply":"2022-09-07T04:02:37.708355Z"},"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-07T04:02:37.71088Z","iopub.execute_input":"2022-09-07T04:02:37.711215Z","iopub.status.idle":"2022-09-07T04:02:37.725858Z","shell.execute_reply.started":"2022-09-07T04:02:37.711183Z","shell.execute_reply":"2022-09-07T04:02:37.723938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EDA","metadata":{}},{"cell_type":"code","source":"train_df.describe()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.groupby('label')['label'].count()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.groupby(\"center_id\")['center_id'].count()","metadata":{"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,"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,"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-07T04:02:37.727803Z","iopub.execute_input":"2022-09-07T04:02:37.728264Z","iopub.status.idle":"2022-09-07T04:02:37.740876Z","shell.execute_reply.started":"2022-09-07T04:02:37.728229Z","shell.execute_reply":"2022-09-07T04:02:37.739968Z"},"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Images of patient id = 09644e (CE)","metadata":{}},{"cell_type":"code","source":"sample_ids = [\"09644e_0\",\"09644e_1\",\"09644e_2\",\"09644e_3\",\"09644e_4\"]\nplot_list_img(sample_ids, \"train\", scale=100)","metadata":{"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Images of patient id = 91b9d3 (LAA)","metadata":{}},{"cell_type":"code","source":"sample_ids = [\"91b9d3_0\",\"91b9d3_1\",\"91b9d3_2\",\"91b9d3_3\",\"91b9d3_4\"]\nplot_list_img(sample_ids, \"train\", scale=100)","metadata":{"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### CE class sample","metadata":{}},{"cell_type":"code","source":"sample_ids = train_df[train_df.label==\"CE\"].image_id[:10].values\nplot_list_img(sample_ids, \"train\", scale=100)","metadata":{"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### LAA class sample","metadata":{}},{"cell_type":"code","source":"sample_ids = train_df[train_df.label==\"LAA\"].image_id[:10].values\nplot_list_img(sample_ids, \"train\", scale=100)","metadata":{"_kg_hide-input":true,"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 = \"efficientnet_b6_best\"\n\n\ncfg = TrainingConfigurations()\ncfg","metadata":{"execution":{"iopub.status.busy":"2022-09-07T04:02:37.742079Z","iopub.execute_input":"2022-09-07T04:02:37.743063Z","iopub.status.idle":"2022-09-07T04:02:37.82623Z","shell.execute_reply.started":"2022-09-07T04:02:37.743026Z","shell.execute_reply":"2022-09-07T04:02:37.825198Z"},"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-07T04:02:37.83144Z","iopub.execute_input":"2022-09-07T04:02:37.831895Z","iopub.status.idle":"2022-09-07T04:02:37.841041Z","shell.execute_reply.started":"2022-09-07T04:02:37.83186Z","shell.execute_reply":"2022-09-07T04:02:37.840046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"\nclass Model(nn.Module):\n    def __init__(self, nc, dropout, freeze):\n        super(Model, self).__init__()\n\n        self.model = models.efficientnet_b6(pretrained=True)\n\n        num_ftrs = self.model.classifier[1].in_features\n        self.model.classifier[0] = nn.Dropout(p=dropout, inplace=True)\n        self.model.classifier[1] = nn.Linear(num_ftrs, nc)\n\n    def forward(self, x):\n        out = self.model(x)\n        return out","metadata":{"execution":{"iopub.status.busy":"2022-09-07T04:02:37.842627Z","iopub.execute_input":"2022-09-07T04:02:37.843156Z","iopub.status.idle":"2022-09-07T04:02:37.851479Z","shell.execute_reply.started":"2022-09-07T04:02:37.843101Z","shell.execute_reply":"2022-09-07T04:02:37.850479Z"},"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-07T04:02:37.85308Z","iopub.execute_input":"2022-09-07T04:02:37.853636Z","iopub.status.idle":"2022-09-07T04:02:37.862884Z","shell.execute_reply.started":"2022-09-07T04:02:37.853602Z","shell.execute_reply":"2022-09-07T04:02:37.861906Z"},"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-07T04:02:37.864436Z","iopub.execute_input":"2022-09-07T04:02:37.864967Z","iopub.status.idle":"2022-09-07T04:02:37.87698Z","shell.execute_reply.started":"2022-09-07T04:02:37.864933Z","shell.execute_reply":"2022-09-07T04:02:37.875963Z"},"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-07T04:02:37.878387Z","iopub.execute_input":"2022-09-07T04:02:37.879476Z","iopub.status.idle":"2022-09-07T04:02:37.898356Z","shell.execute_reply.started":"2022-09-07T04:02:37.879441Z","shell.execute_reply":"2022-09-07T04:02:37.897139Z"},"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=cfg.bs,\n    num_workers=2,\n    shuffle=True,\n    pin_memory=True,\n)\ntest_dataloader = DataLoader(\n    test_dataset,\n    batch_size=cfg.bs,\n    num_workers=2,\n    shuffle=False,\n    pin_memory=True,\n)","metadata":{"execution":{"iopub.status.busy":"2022-09-07T04:02:37.89953Z","iopub.execute_input":"2022-09-07T04:02:37.900069Z","iopub.status.idle":"2022-09-07T04:02:37.908376Z","shell.execute_reply.started":"2022-09-07T04:02:37.900035Z","shell.execute_reply":"2022-09-07T04:02:37.907388Z"},"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-07T04:02:37.910353Z","iopub.execute_input":"2022-09-07T04:02:37.911002Z","iopub.status.idle":"2022-09-07T04:02:37.924362Z","shell.execute_reply.started":"2022-09-07T04:02:37.910968Z","shell.execute_reply":"2022-09-07T04:02:37.923414Z"},"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, cfg.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            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-07T04:02:37.92577Z","iopub.execute_input":"2022-09-07T04:02:37.926298Z","iopub.status.idle":"2022-09-07T04:02:37.941769Z","shell.execute_reply.started":"2022-09-07T04:02:37.926263Z","shell.execute_reply":"2022-09-07T04:02:37.940774Z"},"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())\noptimizer = optim.SGD(\n    model_parameters, lr=cfg.learning_rate, momentum=cfg.momentum, weight_decay=cfg.weight_decay)\nscheduler = None\ncriterion = nn.CrossEntropyLoss()\n\nhistory = start_training(\n    model, optimizer, scheduler, device=cfg['device'], num_epochs=cfg.epochs)\n","metadata":{"execution":{"iopub.status.busy":"2022-09-07T04:02:37.957714Z","iopub.execute_input":"2022-09-07T04:02:37.958255Z","iopub.status.idle":"2022-09-07T04:04:02.653759Z","shell.execute_reply.started":"2022-09-07T04:02:37.958219Z","shell.execute_reply":"2022-09-07T04:04:02.652404Z"},"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,"execution":{"iopub.status.busy":"2022-09-07T04:04:47.214437Z","iopub.execute_input":"2022-09-07T04:04:47.214846Z","iopub.status.idle":"2022-09-07T04:04:47.570614Z","shell.execute_reply.started":"2022-09-07T04:04:47.214808Z","shell.execute_reply":"2022-09-07T04:04:47.569686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp ~/.cache/torch/hub/checkpoints/* .","metadata":{"execution":{"iopub.status.busy":"2022-09-07T04:04:47.895008Z","iopub.execute_input":"2022-09-07T04:04:47.897443Z","iopub.status.idle":"2022-09-07T04:04:49.108356Z","shell.execute_reply.started":"2022-09-07T04:04:47.897407Z","shell.execute_reply":"2022-09-07T04:04:49.107049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}