{"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":"# Library","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append('../input/pytorch-image-models/pytorch-image-models-master')\n\nimport IPython.display\n\nimport os\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\nfrom PIL import Image\nfrom glob import glob\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\n#import Pyvips\n\nfrom sklearn import preprocessing\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold, train_test_split\nfrom skimage.filters import sobel\nfrom skimage import segmentation\nfrom skimage.color import label2rgb\nfrom skimage.color import rgb2hed, hed2rgb\nfrom skimage.exposure import rescale_intensity\nfrom skimage.measure import regionprops, regionprops_table\nfrom sklearn.preprocessing import StandardScaler\nfrom scipy import ndimage as ndi\nfrom matplotlib.patches import Rectangle\n\nimport torchvision\n\nfrom tqdm.auto import tqdm\nfrom tqdm import trange\nfrom time import sleep\nfrom functools import partial\nimport tifffile as tiff\n\nimport cv2 as cv\nfrom openslide import OpenSlide\nimport seaborn as sns\nfrom matplotlib import pyplot as plt\nfrom pprint import pprint\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD\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\nimport torchvision.transforms as transforms\nimport torch.optim as optim\nimport gc\n\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)\n\nfrom torch.cuda.amp import autocast, GradScaler\nImage.MAX_IMAGE_PIXELS = None\nimport warnings\nwarnings.filterwarnings('ignore')\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-17T19:17:09.910937Z","iopub.execute_input":"2022-08-17T19:17:09.911636Z","iopub.status.idle":"2022-08-17T19:17:09.929215Z","shell.execute_reply.started":"2022-08-17T19:17:09.911598Z","shell.execute_reply":"2022-08-17T19:17:09.927907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loading data","metadata":{}},{"cell_type":"code","source":"train_data = pd.read_csv('../input/mayo-clinic-strip-ai/train.csv')\ntransformed_train = pd.read_csv('../input/mayo-clinic-output/new_train.csv')\n#print(train_data.head())\nprint(train_data[train_data[\"image_id\"]==\"1a2e9e_0\"])\ntrain_data['enc_label'] = np.where(train_data['label']== 'CE', 1, 0)\ntrain, vaild = train_test_split(train_data, test_size=0.204)\ntest = pd.read_csv('../input/mayo-clinic-strip-ai/test.csv')\nsample_sub = pd.read_csv('../input/mayo-clinic-strip-ai/sample_submission.csv')\n","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:38.597053Z","iopub.execute_input":"2022-08-17T18:16:38.598091Z","iopub.status.idle":"2022-08-17T18:16:38.655436Z","shell.execute_reply.started":"2022-08-17T18:16:38.598051Z","shell.execute_reply":"2022-08-17T18:16:38.654333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EDA","metadata":{}},{"cell_type":"code","source":"print(train.head())\ntrain[train[\"label\"]==\"LAA\"]","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:38.915242Z","iopub.execute_input":"2022-08-17T18:16:38.915681Z","iopub.status.idle":"2022-08-17T18:16:38.945171Z","shell.execute_reply.started":"2022-08-17T18:16:38.915647Z","shell.execute_reply":"2022-08-17T18:16:38.94407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"patients_train = train['patient_id'].nunique()\npatients_test = test['patient_id'].nunique()\nprint(f\"Number of unique patient {patients_train}\")\n","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:39.14171Z","iopub.execute_input":"2022-08-17T18:16:39.14313Z","iopub.status.idle":"2022-08-17T18:16:39.152602Z","shell.execute_reply.started":"2022-08-17T18:16:39.143078Z","shell.execute_reply":"2022-08-17T18:16:39.151445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sns.countplot(train.label, palette=\"Reds_r\")\nplt.title(\"Label Count\");","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:39.314927Z","iopub.execute_input":"2022-08-17T18:16:39.316123Z","iopub.status.idle":"2022-08-17T18:16:39.529051Z","shell.execute_reply.started":"2022-08-17T18:16:39.316079Z","shell.execute_reply":"2022-08-17T18:16:39.527904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10,5))\nsns.countplot(train_data.groupby(\"patient_id\").image_num.size(), palette=\"Greens_r\")\nplt.xlabel(\"Number of images per patient\")\nplt.title(\"Max image number per patient in train\");\n","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:39.531258Z","iopub.execute_input":"2022-08-17T18:16:39.531965Z","iopub.status.idle":"2022-08-17T18:16:39.742964Z","shell.execute_reply.started":"2022-08-17T18:16:39.531922Z","shell.execute_reply":"2022-08-17T18:16:39.741898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Image Transformation ","metadata":{}},{"cell_type":"code","source":"def rezie_image(image):\n    resized_image = cv.resize(image,(int(image.shape[1]/33),int(image.shape[0]/33)),interpolation= cv.INTER_LINEAR)\n    return resized_image","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:39.832378Z","iopub.execute_input":"2022-08-17T18:16:39.833336Z","iopub.status.idle":"2022-08-17T18:16:39.839219Z","shell.execute_reply.started":"2022-08-17T18:16:39.833286Z","shell.execute_reply":"2022-08-17T18:16:39.837982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def grey_resize(image):\n    gray_resized_image = cv.cvtColor(image, cv.COLOR_RGB2GRAY)    \n    return gray_resized_image","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:39.991792Z","iopub.execute_input":"2022-08-17T18:16:39.992635Z","iopub.status.idle":"2022-08-17T18:16:39.998365Z","shell.execute_reply.started":"2022-08-17T18:16:39.992596Z","shell.execute_reply":"2022-08-17T18:16:39.997087Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def labeled_segment(grey_resized_image):\n    elevation_map = sobel(grey_resized_image)\n    markers = np.zeros_like(grey_resized_image)\n    markers[grey_resized_image >= grey_resized_image.mean()] = 1\n    markers[grey_resized_image < grey_resized_image.mean()] = 2\n    segmented_img = segmentation.watershed(elevation_map, markers)\n    filled_segments = ndi.binary_fill_holes(segmented_img - 1)\n    labeled_segments, _ = ndi.label(filled_segments)\n    return labeled_segments\n","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:40.152694Z","iopub.execute_input":"2022-08-17T18:16:40.153111Z","iopub.status.idle":"2022-08-17T18:16:40.160073Z","shell.execute_reply.started":"2022-08-17T18:16:40.153077Z","shell.execute_reply":"2022-08-17T18:16:40.158554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_labeled_segments(labeled_segments, resized_gray_img):\n    image_label_overlay = label2rgb(labeled_segments, image=resized_gray_img, bg_label=0)\n    fig, ax = plt.subplots(figsize=(10, 8))\n    ax.imshow(image_label_overlay, cmap=plt.cm.gray)\n    ax.set_title('segmentation')\n    ax.axis('off')\n","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:40.377919Z","iopub.execute_input":"2022-08-17T18:16:40.378644Z","iopub.status.idle":"2022-08-17T18:16:40.38499Z","shell.execute_reply.started":"2022-08-17T18:16:40.378608Z","shell.execute_reply":"2022-08-17T18:16:40.38357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_object_coordinates(labeled_segments):\n    properties =['area','bbox','convex_area','bbox_area', 'major_axis_length', 'minor_axis_length', 'eccentricity']\n    df = pd.DataFrame(regionprops_table(labeled_segments, properties=properties))\n    standard_scaler = StandardScaler()\n    scaled_area = standard_scaler.fit_transform(df.area.values.reshape(-1,1))\n    df['scaled_area'] = scaled_area\n    df.sort_values(by=\"scaled_area\", ascending=False, inplace=True)\n    objects = df[df['scaled_area']>=.75]\n    object_coordinates = [(row['bbox-0'],row['bbox-1'],row['bbox-2'],row['bbox-3'] )for index, row in objects.iterrows()]\n    return object_coordinates\n","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:40.572956Z","iopub.execute_input":"2022-08-17T18:16:40.573968Z","iopub.status.idle":"2022-08-17T18:16:40.582651Z","shell.execute_reply.started":"2022-08-17T18:16:40.573923Z","shell.execute_reply":"2022-08-17T18:16:40.581573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_object_coordinates(object_coordinates, resized_image):\n    fig, ax = plt.subplots(1,1, figsize=(18, 16), dpi = 80)\n    for blob in object_coordinates:\n        width = blob[3] - blob[1]\n        height = blob[2] - blob[0]\n        patch = Rectangle((blob[1],blob[0]), width, height, edgecolor='r', facecolor='none')\n","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:40.760879Z","iopub.execute_input":"2022-08-17T18:16:40.761654Z","iopub.status.idle":"2022-08-17T18:16:40.768535Z","shell.execute_reply.started":"2022-08-17T18:16:40.761616Z","shell.execute_reply":"2022-08-17T18:16:40.767326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_objects(object_coordinates, image, image_name, label, count):\n    plt.figure(figsize=(10,18))\n    for i in range(len(object_coordinates)):\n        coordinates = object_coordinates[i]\n        object_image = image[int(coordinates[0]):int(coordinates[2]), int(coordinates[1]):int(coordinates[3])]\n        #plt.imshow(object_image)\n        image_new_name = image_name + \"_\" + str(i)\n        new_train[\"image_name\"].append(image_new_name)\n        new_train[\"label\"].append(label)\n        new_train[\"image_count\"].append(count)\n        cv.imwrite(os.path.join(\"./\", f\"{image_new_name}.jpg\"), object_image)\n        \n\n","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:40.964228Z","iopub.execute_input":"2022-08-17T18:16:40.964609Z","iopub.status.idle":"2022-08-17T18:16:40.972318Z","shell.execute_reply.started":"2022-08-17T18:16:40.964576Z","shell.execute_reply":"2022-08-17T18:16:40.970984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"train_path = \"../input/mayo-clinic-strip-ai/train\"\nImage_names = train['image_id'].values\nnew_train={\"image_count\":[],\"image_name\":[],\"label\":[]}\ncount=1\nscale = 4\nfor image_name in Image_names:\n    Image_label = train.loc[train['image_id'] == image_name, 'enc_label'].iloc[0]\n    image = tiff.imread(os.path.join(train_path, f\"{image_name}.tif\"))\n    resized_image=rezie_image(image)\n    del image\n    gc.collect()\n    grey_resized_image = grey_resize(resized_image)\n    labeled_segments = labeled_segment(grey_resized_image)\n    object_coordinates = get_object_coordinates(labeled_segments)\n    save_objects(object_coordinates, resized_image, image_name, Image_label,count)\n    count+=1\n    #if count == 80:\n    #    break\nnew_train=pd.DataFrame.from_dict(new_train)\"\"\"","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:41.151031Z","iopub.execute_input":"2022-08-17T18:16:41.151755Z","iopub.status.idle":"2022-08-17T18:16:41.159796Z","shell.execute_reply.started":"2022-08-17T18:16:41.151718Z","shell.execute_reply":"2022-08-17T18:16:41.158605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loader","metadata":{}},{"cell_type":"code","source":"class TrainDataset(Dataset):\n    def __init__(self, path, df, transform=None):\n        self.df = df\n        self.path = path\n        self.Image_names = df['image_name'].values\n        self.labels = df['label'].values\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_name = self.Image_names[idx]\n        img= Image.open(os.path.join(self.path, f\"{file_name}.jpg\"))\n        if self.transform:\n            image=self.transform(img)\n\n        label = self.labels[idx]\n\n        return image, torch.tensor(label)","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:41.524477Z","iopub.execute_input":"2022-08-17T18:16:41.525277Z","iopub.status.idle":"2022-08-17T18:16:41.534492Z","shell.execute_reply.started":"2022-08-17T18:16:41.525233Z","shell.execute_reply":"2022-08-17T18:16:41.533299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transforms","metadata":{}},{"cell_type":"code","source":"batch_size=64\ndata_transform = transforms.Compose([\n        transforms.Resize((256,256)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                             std=[0.229, 0.224, 0.225])\n    ])\ntrain_dataset = TrainDataset(\"../input/mayo-clinic-output/\", transformed_train, transform = data_transform)\n#vaild_dataset = TrainDataset(\"../input/mayo-clinic-strip-ai/train\", vaild, transform = data_transform)\n#test_dataset = TrainDataset(\"../input/mayo-clinic-strip-ai/train\", test, transform = data_transform)\n\ndataset_loader = torch.utils.data.DataLoader(train_dataset,\n                                             batch_size=batch_size, shuffle=True,\n                                             num_workers=0)\n#dataset_loader_vaild = torch.utils.data.DataLoader(vaild_dataset,\n#                                             batch_size=batch_size, shuffle=True,\n#                                             num_workers=0)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:41.913835Z","iopub.execute_input":"2022-08-17T18:16:41.914584Z","iopub.status.idle":"2022-08-17T18:16:41.922601Z","shell.execute_reply.started":"2022-08-17T18:16:41.914545Z","shell.execute_reply":"2022-08-17T18:16:41.921411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Binary accuracy metric","metadata":{}},{"cell_type":"code","source":"def binary_acc(y_pred, y_test):\n    y_pred_tag = torch.round(torch.sigmoid(y_pred))\n\n    correct_results_sum = (y_pred_tag == y_test).sum().float()\n    acc = correct_results_sum/y_test.shape[0]\n    acc = torch.round(acc * 100)\n    return acc\n","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:42.284704Z","iopub.execute_input":"2022-08-17T18:16:42.285453Z","iopub.status.idle":"2022-08-17T18:16:42.29138Z","shell.execute_reply.started":"2022-08-17T18:16:42.285416Z","shell.execute_reply":"2022-08-17T18:16:42.290068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# AlexNet Model","metadata":{}},{"cell_type":"code","source":"class AlexNet(nn.Module):\n    def __init__(self):\n        super(AlexNet, self).__init__()\n        self.conv1 = nn.Conv2d(in_channels=3, out_channels= 96, kernel_size= 11, stride=4, padding=0 )\n        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2)\n        self.conv2 = nn.Conv2d(in_channels=96, out_channels=256, kernel_size=5, stride= 1, padding= 2)\n        self.conv3 = nn.Conv2d(in_channels=256, out_channels=384, kernel_size=3, stride= 1, padding= 1)\n        self.conv4 = nn.Conv2d(in_channels=384, out_channels=384, kernel_size=3, stride=1, padding=1)\n        self.conv5 = nn.Conv2d(in_channels=384, out_channels=256, kernel_size=3, stride=1, padding=1)\n        self.fc1  = nn.Linear(in_features= 9216, out_features= 4096)\n        self.fc2  = nn.Linear(in_features= 4096, out_features= 4096)\n        self.fc3 = nn.Linear(in_features=4096 , out_features=1)\n\n\n    def forward(self,x):\n        x = F.relu(self.conv1(x))\n        x = self.maxpool(x)\n        x = F.relu(self.conv2(x))\n        x = self.maxpool(x)\n        x = F.relu(self.conv3(x))\n        x = F.relu(self.conv4(x))\n        x = F.relu(self.conv5(x))\n        x = self.maxpool(x)\n        x = x.reshape(x.shape[0], -1)\n        x = F.relu(self.fc1(x))\n        x = F.relu(self.fc2(x))\n        x = self.fc3(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:42.70098Z","iopub.execute_input":"2022-08-17T18:16:42.702116Z","iopub.status.idle":"2022-08-17T18:16:42.714685Z","shell.execute_reply.started":"2022-08-17T18:16:42.702079Z","shell.execute_reply":"2022-08-17T18:16:42.71327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training ","metadata":{}},{"cell_type":"code","source":"\"\"\"model = AlexNet()\n#model = VGG16()\nmodel = model.to(device=device)\nprint(device)\n## Loss and optimizer\nlearning_rate = 1e-3 #I picked this because it seems to be the most used by experts\nload_model = True\ncriterion = nn.BCELoss()\noptimizer = optim.Adam(model.parameters(), lr= learning_rate) #Adam seems to be the most popular for deep learning\nfor epoch in range(10): #I decided to train the model for 5 epochs\n    with tqdm(dataset_loader, unit=\"batch\") as tepoch:\n        e=epoch+1\n        loss_ep = 0\n        epoch_acc = 0\n\n        for (data, targets) in (tepoch):\n            tepoch.set_description(f\"Epoch {epoch}\")\n            \n            data = data.to(device=device).requires_grad_(True)\n            targets = targets.type(torch.FloatTensor).to(device=device).requires_grad_(True)\n            optimizer.zero_grad()\n            y_pred = model(data)\n            sco = torch.clamp(y_pred,0,1)\n            loss = criterion(sco.squeeze(),targets)\n            acc = binary_acc(sco.squeeze(),targets)\n            loss.backward()\n            optimizer.step()\n            loss_ep += loss.item()\n            epoch_acc += acc.item()\n            tepoch.set_postfix(loss=loss.item(), accuracy=100. * acc.item())\n            sleep(0.1)\n    print(f'Epoch {e+0:03}: | Loss: {loss_ep/len(dataset_loader):.5f} | Acc: {epoch_acc/len(dataset_loader):.3f}')\n\n        \n   \"\"\"","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:43.151319Z","iopub.execute_input":"2022-08-17T18:16:43.15173Z","iopub.status.idle":"2022-08-17T18:16:43.160359Z","shell.execute_reply.started":"2022-08-17T18:16:43.151696Z","shell.execute_reply":"2022-08-17T18:16:43.159094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"with torch.no_grad():\n        num_correct = 0\n        num_samples = 0\n        \n        for batch_idx, (data,targets) in enumerate(dataset_loader_vaild):\n            data = data.to(device=device)\n            targets = targets.to(device=device)\n            ## Forward Pass\n            scores = model(data)\n            _, predictions = scores.max(1)\n            num_correct += (predictions == targets).sum()\n            num_samples += predictions.size(0)\n        print(\n            f\"Got {num_correct} / {num_samples} with accuracy {float(num_correct) / float(num_samples) * 100:.2f}\"\n        )\"\"\"","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:16:43.333049Z","iopub.execute_input":"2022-08-17T18:16:43.334191Z","iopub.status.idle":"2022-08-17T18:16:43.343502Z","shell.execute_reply.started":"2022-08-17T18:16:43.334149Z","shell.execute_reply":"2022-08-17T18:16:43.342377Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchvision.models as models\nimport copy\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel_ft = models.vgg16(pretrained=True)\nmodel_ft = models.densenet121(pretrained=True)\nmodel_ft=model_ft.to(device)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:19:39.718939Z","iopub.execute_input":"2022-08-17T18:19:39.71934Z","iopub.status.idle":"2022-08-17T18:19:43.262657Z","shell.execute_reply.started":"2022-08-17T18:19:39.719306Z","shell.execute_reply":"2022-08-17T18:19:43.261605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#num_ftrs = model_ft.classifier[0].in_features \n# Here the size of each output sample is set to 2.\n# Alternatively, it can be generalized to nn.Linear(num_ftrs, len(class_names)).\nmodel_ft.fc = nn.Linear(1024, 2)\n#model_ft = model_ft.to(device)\ncriterion = nn.CrossEntropyLoss()\n# Observe that all parameters are being optimized\noptimizer_ft = optim.SGD(model_ft.parameters(), lr=0.01, momentum=0.9)\n# Decay LR by a factor of 0.1 every 7 epochs\nexp_lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer_ft, step_size=7, gamma=0.1)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:20:03.109365Z","iopub.execute_input":"2022-08-17T18:20:03.109734Z","iopub.status.idle":"2022-08-17T18:20:03.121284Z","shell.execute_reply.started":"2022-08-17T18:20:03.109704Z","shell.execute_reply":"2022-08-17T18:20:03.119907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(model, criterion, optimizer, scheduler, num_epochs=25):\n    since = time.time()\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_acc = 0.0\n    for epoch in range(num_epochs):\n        with tqdm(dataset_loader, unit=\"batch\") as tepoch:\n            e=epoch+1\n            print('Epoch {}/{}'.format(epoch, num_epochs - 1))\n            print('-' * 10)\n            # Each epoch has a training and validation phase\n            for phase in ['train', 'val']:\n                if phase == 'train':\n                    model.train()  # Set model to training mode\n                else:\n                    model.eval()   # Set model to evaluate mode\n                running_loss = 0.0\n                running_corrects = 0\n                # Iterate over data.\n                for inputs, labels in tepoch:\n                    inputs = inputs.to(device)\n                    labels = labels.to(device)\n                    # zero the parameter gradients\n                    optimizer.zero_grad()\n                    # forward\n                    # track history if only in train\n                    with torch.set_grad_enabled(phase == 'train'):\n                        outputs = model(inputs)\n                        _, preds = torch.max(outputs, 1)\n                        loss = criterion(outputs, labels)\n                        # bacoptimizerkward + optimize only if in training phase\n                        if phase == 'train':\n                            loss.backward()\n                            optimizer.step()\n                    # statistics\n                    running_loss += loss.item() * inputs.size(0)\n                    running_corrects += torch.sum(preds == labels.data)\n                if phase == 'train':\n                    scheduler.step()\n                epoch_loss = running_loss / len(dataset_loader)\n                epoch_acc = running_corrects.double() / len(dataset_loader)\n                print('{} Loss: {:.4f} Acc: {:.4f}'.format(\n                    phase, epoch_loss, epoch_acc))\n                # deep copy the model\n            print()\n            tepoch.set_postfix(loss=loss.item())\n            sleep(0.1)\n    time_elapsed = time.time() - since\n    print('Training complete in {:.0f}m {:.0f}s'.format(\n        time_elapsed // 60, time_elapsed % 60))\n    # load best model weights\n    model.load_state_dict(best_model_wts)\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:20:05.115653Z","iopub.execute_input":"2022-08-17T18:20:05.116099Z","iopub.status.idle":"2022-08-17T18:20:05.130902Z","shell.execute_reply.started":"2022-08-17T18:20:05.116056Z","shell.execute_reply":"2022-08-17T18:20:05.129623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_ft = train_model(model_ft, criterion, optimizer_ft, exp_lr_scheduler,\n                       num_epochs=25)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-17T18:20:06.557505Z","iopub.execute_input":"2022-08-17T18:20:06.559169Z","iopub.status.idle":"2022-08-17T19:04:51.534624Z","shell.execute_reply.started":"2022-08-17T18:20:06.559121Z","shell.execute_reply":"2022-08-17T19:04:51.532934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def imshow(inp, title=None):\n    \"\"\"Imshow for Tensor.\"\"\"\n    inp = inp.numpy().transpose((1, 2, 0))\n    mean = np.array([0.485, 0.456, 0.406])\n    std = np.array([0.229, 0.224, 0.225])\n    inp = std * inp + mean\n    inp = np.clip(inp, 0, 1)\n    plt.imshow(inp)\n    if title is not None:\n        plt.title(title)\n    plt.pause(0.001)  # pause a bit so that plots are updated\n\n\n# Get a batch of training data\ninputs, classes = next(iter(dataset_loader))\n\n# Make a grid from batch\nout = torchvision.utils.make_grid(inputs)\nz=train['enc_label'].values\nprint(z)\nimshow(out, title=[z[x] for x in classes])","metadata":{"execution":{"iopub.status.busy":"2022-08-17T19:20:43.798908Z","iopub.execute_input":"2022-08-17T19:20:43.799445Z","iopub.status.idle":"2022-08-17T19:20:46.296482Z","shell.execute_reply.started":"2022-08-17T19:20:43.799407Z","shell.execute_reply":"2022-08-17T19:20:46.295517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_model(model, num_images=6):\n    was_training = model.training\n    model.eval()\n    images_so_far = 0\n    fig = plt.figure()\n\n    with torch.no_grad():\n        for i, (inputs, labels) in enumerate(dataset_loader):\n            inputs = inputs.to(device)\n            labels = labels.to(device)\n\n            outputs = model(inputs)\n            _, preds = torch.max(outputs, 1)\n\n            for j in range(inputs.size()[0]):\n                images_so_far += 1\n                ax = plt.subplot(num_images//2, 2, images_so_far)\n                ax.axis('off')\n                z=train['enc_label'].values\n                ax.set_title(f'predicted: {z[preds[j]]}')\n                imshow(inputs.cpu().data[j])\n\n                if images_so_far == num_images:\n                    model.train(mode=was_training)\n                    return\n        model.train(mode=was_training)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-17T19:20:06.805398Z","iopub.execute_input":"2022-08-17T19:20:06.805824Z","iopub.status.idle":"2022-08-17T19:20:06.816743Z","shell.execute_reply.started":"2022-08-17T19:20:06.80578Z","shell.execute_reply":"2022-08-17T19:20:06.815341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"visualize_model(model_ft)","metadata":{"execution":{"iopub.status.busy":"2022-08-17T19:20:06.818417Z","iopub.execute_input":"2022-08-17T19:20:06.818893Z","iopub.status.idle":"2022-08-17T19:20:07.508942Z","shell.execute_reply.started":"2022-08-17T19:20:06.818833Z","shell.execute_reply":"2022-08-17T19:20:07.504797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}