{"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":"## Initialization","metadata":{}},{"cell_type":"code","source":"!pip install -qU python-gdcm pydicom pylibjpeg","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-02-07T11:00:36.201073Z","iopub.execute_input":"2023-02-07T11:00:36.201629Z","iopub.status.idle":"2023-02-07T11:00:52.888708Z","shell.execute_reply.started":"2023-02-07T11:00:36.201516Z","shell.execute_reply":"2023-02-07T11:00:52.887552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport cv2\nimport glob\nimport gdcm\nimport pydicom\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\nfrom tqdm.notebook import tqdm\nfrom joblib import Parallel, delayed\n\nfrom plotly.subplots import make_subplots\nimport plotly.graph_objs as go\nimport copy\nimport torch\nfrom PIL import Image\nfrom PIL import Image, ImageDraw\nfrom torch.utils.data import Dataset\nimport torchvision.transforms as transforms\nfrom torch.utils.data import random_split\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport torch.nn as nn\nfrom torchvision import utils\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:00:52.894382Z","iopub.execute_input":"2023-02-07T11:00:52.896647Z","iopub.status.idle":"2023-02-07T11:00:54.736357Z","shell.execute_reply.started":"2023-02-07T11:00:52.896601Z","shell.execute_reply":"2023-02-07T11:00:54.735268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Examples","metadata":{}},{"cell_type":"markdown","source":"## Save the processed data\n**Images are quite big so resizing them is necessary.**\n  - Use 256 to train your first models, or if you don't have a lot of compute\n  - use 512 to have competitive models\n  - Check if 768/1024 is better, if you have the compute power\n\n**I advise using the `png` format because the jpg compression can be annoying during inference.**","metadata":{}},{"cell_type":"code","source":"SAVE_FOLDER = \"output/\"\nSIZE = 512\nEXTENSION = \"png\"\n\nos.makedirs(SAVE_FOLDER, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:00:54.741463Z","iopub.execute_input":"2023-02-07T11:00:54.742344Z","iopub.status.idle":"2023-02-07T11:00:54.751266Z","shell.execute_reply.started":"2023-02-07T11:00:54.742305Z","shell.execute_reply":"2023-02-07T11:00:54.75027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process(f, size=512, save_folder=\"\", extension=\"png\"):\n    patient = f.split('/')[-2]\n    image = f.split('/')[-1][:-4]\n\n    dicom = pydicom.dcmread(f)\n    img = dicom.pixel_array\n\n    img = (img - img.min()) / (img.max() - img.min())\n\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        img = 1 - img\n    \n    img = cv2.resize(img, (size, size))\n    return Image.fromarray(img)\n    #cv2.imwrite(save_folder + where + '/' + str(cancer) + '/' + f\"{patient}_{image}.{extension}\", (img * 255).astype(np.uint8))","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:00:54.756717Z","iopub.execute_input":"2023-02-07T11:00:54.759588Z","iopub.status.idle":"2023-02-07T11:00:54.768491Z","shell.execute_reply.started":"2023-02-07T11:00:54.759552Z","shell.execute_reply":"2023-02-07T11:00:54.767659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport csv\n\ntrain_set = []\ntest_set = []\nval_set = []\n\ndf = pd.read_csv('/kaggle/input/new-csv/new_test.csv')\nfor i, row in df.iterrows():\n    test_set.append(f\"{row['patient_id']}/{row['image_id']}.dcm\")\nprint('test_done')\n\ndf = pd.read_csv('/kaggle/input/new-csv/new_train.csv')\nfor i, row in df.iterrows():\n    train_set.append(f\"{row['patient_id']}/{row['image_id']}.dcm\")\nprint('train_done')\n\ndf = pd.read_csv('/kaggle/input/new-csv/new_validation.csv')\nfor i, row in df.iterrows():\n    val_set.append(f\"{row['patient_id']}/{row['image_id']}.dcm\")\nprint('val_done')","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:00:54.772875Z","iopub.execute_input":"2023-02-07T11:00:54.77561Z","iopub.status.idle":"2023-02-07T11:00:57.859684Z","shell.execute_reply.started":"2023-02-07T11:00:54.775576Z","shell.execute_reply":"2023-02-07T11:00:57.858551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(train_set[0:3])\nprint(test_set[0:3])\nprint(val_set[0:3])","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:00:57.861619Z","iopub.execute_input":"2023-02-07T11:00:57.862013Z","iopub.status.idle":"2023-02-07T11:00:57.868281Z","shell.execute_reply.started":"2023-02-07T11:00:57.861974Z","shell.execute_reply":"2023-02-07T11:00:57.867204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Main loop","metadata":{}},{"cell_type":"code","source":"torch.manual_seed(0)\n\nclass pytorch_data(Dataset):\n    \n    def __init__(self,data_dir,transform,data_type=\"train\"):      \n        # Get Labels\n        if data_type == \"train\":\n            file_names=train_set\n            labels_data=\"/kaggle/input/new-csv/new_train.csv\"\n        elif data_type == \"val\":\n            file_names=val_set\n            labels_data=\"/kaggle/input/new-csv/new_validation.csv\"\n        else :\n            file_names=test_set\n            labels_data=\"/kaggle/input/new-csv/new_test.csv\"\n        \n        self.full_filenames = [os.path.join(data_dir, f) for f in file_names]\n        labels_df=pd.read_csv(labels_data)\n        labels_df.set_index(\"cancer\", inplace=True) # set data frame index to id\n        print(labels_df)\n        self.labels = [labels_df.loc[int(filename[:-4].split('/')[1])].values[0] for filename in file_names]  # obtained labels from df\n        self.transform = transform\n      \n    def __len__(self):\n        return len(self.full_filenames) # size of dataset\n      \n    def __getitem__(self, idx):\n        # open image, apply transforms and return with label\n        image = process(self.full_filenames[idx])  # Open Image with PIL\n        image = self.transform(image) # Apply Specific Transformation to Image\n        return image, self.labels[idx]\n","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:02:13.384321Z","iopub.execute_input":"2023-02-07T11:02:13.384688Z","iopub.status.idle":"2023-02-07T11:02:13.396039Z","shell.execute_reply.started":"2023-02-07T11:02:13.384656Z","shell.execute_reply":"2023-02-07T11:02:13.393714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **사이즈 고치기**","metadata":{}},{"cell_type":"code","source":"# define transformation that converts a PIL image into PyTorch tensors\nimport torchvision.transforms as transforms\ndata_transformer = transforms.Compose([transforms.ToTensor(),\n                                       transforms.Resize((512,512))])","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:02:25.225115Z","iopub.execute_input":"2023-02-07T11:02:25.225688Z","iopub.status.idle":"2023-02-07T11:02:25.2311Z","shell.execute_reply.started":"2023-02-07T11:02:25.225651Z","shell.execute_reply":"2023-02-07T11:02:25.230129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define an object of the custom dataset for the train folder.\ndata_dir = '/kaggle/input/rsna-breast-cancer-detection/train_images/'\nimg_dataset = pytorch_data(data_dir, data_transformer, \"train\") # Histopathalogic images","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:55.052574Z","iopub.execute_input":"2023-02-07T11:01:55.052957Z","iopub.status.idle":"2023-02-07T11:01:55.232164Z","shell.execute_reply.started":"2023-02-07T11:01:55.052923Z","shell.execute_reply":"2023-02-07T11:01:55.230641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# load an example tensor\nimg,label=img_dataset[10]\nprint(img.shape,torch.min(img),torch.max(img))","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:02.007387Z","iopub.execute_input":"2023-02-07T11:01:02.007762Z","iopub.status.idle":"2023-02-07T11:01:03.068558Z","shell.execute_reply.started":"2023-02-07T11:01:02.007727Z","shell.execute_reply":"2023-02-07T11:01:03.067527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Split","metadata":{}},{"cell_type":"code","source":"# Split Pytorch tensor\ntrain_ts=img_dataset\nval_ts= pytorch_data(data_dir, data_transformer, \"val\")\ntest_ts= pytorch_data(data_dir, data_transformer, \"test\")\n\nprint(\"train dataset size:\", len(train_ts))\nprint(\"validation dataset size:\", len(val_ts))\nprint(\"test dataset size:\", len(test_ts))","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:03.070837Z","iopub.execute_input":"2023-02-07T11:01:03.071534Z","iopub.status.idle":"2023-02-07T11:01:04.991163Z","shell.execute_reply.started":"2023-02-07T11:01:03.071494Z","shell.execute_reply":"2023-02-07T11:01:04.989997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import plotly.express as px\n\ndef plot_img(x,y,title=None):\n\n    npimg = x.numpy() # convert tensor to numpy array\n    npimg_tr=np.transpose(npimg, (1,2,0)) # Convert to H*W*C shape\n    fig = px.imshow(npimg_tr)\n    fig.update_layout(template='plotly_white')\n    fig.update_layout(title=title,height=300,margin={'l':10,'r':20,'b':10})\n    fig.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:04.992854Z","iopub.execute_input":"2023-02-07T11:01:04.993231Z","iopub.status.idle":"2023-02-07T11:01:06.489134Z","shell.execute_reply.started":"2023-02-07T11:01:04.993195Z","shell.execute_reply":"2023-02-07T11:01:06.48818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define the following transformations for the training dataset\ntr_transf = transforms.Compose([\n#     transforms.Resize((40,40)),\n    transforms.RandomHorizontalFlip(p=0.5), \n    transforms.RandomVerticalFlip(p=0.5),  \n    transforms.RandomRotation(45),         \n#     transforms.RandomResizedCrop(50,scale=(0.8,1.0),ratio=(1.0,1.0)),\n    transforms.ToTensor()])","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:06.490453Z","iopub.execute_input":"2023-02-07T11:01:06.491129Z","iopub.status.idle":"2023-02-07T11:01:06.49696Z","shell.execute_reply.started":"2023-02-07T11:01:06.49109Z","shell.execute_reply":"2023-02-07T11:01:06.495778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# For the validation dataset, we don't need any augmentation; simply convert images into tensors\nval_transf = transforms.Compose([\n    transforms.ToTensor()])\n\n# After defining the transformations, overwrite the transform functions of train_ts, val_ts\ntrain_ts.transform=tr_transf\nval_ts.transform=val_transf","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:06.498459Z","iopub.execute_input":"2023-02-07T11:01:06.499061Z","iopub.status.idle":"2023-02-07T11:01:06.506906Z","shell.execute_reply.started":"2023-02-07T11:01:06.499027Z","shell.execute_reply":"2023-02-07T11:01:06.505978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# The subset can also have transform attribute (if we asign)\ntrain_ts.transform","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:06.508743Z","iopub.execute_input":"2023-02-07T11:01:06.509411Z","iopub.status.idle":"2023-02-07T11:01:06.519803Z","shell.execute_reply.started":"2023-02-07T11:01:06.509375Z","shell.execute_reply":"2023-02-07T11:01:06.51872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Creating Dataloaders","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\n# Training DataLoader\ntrain_dl = DataLoader(train_ts,\n                      batch_size=32, \n                      shuffle=True)\n\n# Validation DataLoader\nval_dl = DataLoader(val_ts,\n                    batch_size=32,\n                    shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:06.521301Z","iopub.execute_input":"2023-02-07T11:01:06.521755Z","iopub.status.idle":"2023-02-07T11:01:06.528269Z","shell.execute_reply.started":"2023-02-07T11:01:06.521719Z","shell.execute_reply":"2023-02-07T11:01:06.527297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define Binary Classifier","metadata":{}},{"cell_type":"code","source":"def findConv2dOutShape(hin,win,conv,pool=2):\n    # get conv arguments\n    kernel_size=conv.kernel_size\n    stride=conv.stride\n    padding=conv.padding\n    dilation=conv.dilation\n\n    hout=np.floor((hin+2*padding[0]-dilation[0]*(kernel_size[0]-1)-1)/stride[0]+1)\n    wout=np.floor((win+2*padding[1]-dilation[1]*(kernel_size[1]-1)-1)/stride[1]+1)\n\n    if pool:\n        hout/=pool\n        wout/=pool\n    return int(hout),int(wout)\n\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n# Neural Network\nclass Network(nn.Module):\n    \n    # Network Initialisation\n    def __init__(self, params):\n        \n        super(Network, self).__init__()\n    \n        Cin,Hin,Win=params[\"shape_in\"]\n        init_f=params[\"initial_filters\"] \n        num_fc1=params[\"num_fc1\"]  \n        num_classes=params[\"num_classes\"] \n        self.dropout_rate=params[\"dropout_rate\"] \n        \n        # Convolution Layers\n        self.conv1 = nn.Conv2d(Cin, init_f, kernel_size=3)\n        h,w=findConv2dOutShape(Hin,Win,self.conv1)\n        self.conv2 = nn.Conv2d(init_f, 2*init_f, kernel_size=3)\n        h,w=findConv2dOutShape(h,w,self.conv2)\n        self.conv3 = nn.Conv2d(2*init_f, 4*init_f, kernel_size=3)\n        h,w=findConv2dOutShape(h,w,self.conv3)\n        self.conv4 = nn.Conv2d(4*init_f, 8*init_f, kernel_size=3)\n        h,w=findConv2dOutShape(h,w,self.conv4)\n        \n        # compute the flatten size\n        self.num_flatten=h*w*8*init_f\n        self.fc1 = nn.Linear(self.num_flatten, num_fc1)\n        self.fc2 = nn.Linear(num_fc1, num_classes)\n\n    def forward(self,X):\n        \n        # Convolution & Pool Layers\n        X = F.relu(self.conv1(X)); \n        X = F.max_pool2d(X, 2, 2)\n        X = F.relu(self.conv2(X))\n        X = F.max_pool2d(X, 2, 2)\n        X = F.relu(self.conv3(X))\n        X = F.max_pool2d(X, 2, 2)\n        X = F.relu(self.conv4(X))\n        X = F.max_pool2d(X, 2, 2)\n\n        X = X.view(-1, self.num_flatten)\n        \n        X = F.relu(self.fc1(X))\n        X=F.dropout(X, self.dropout_rate)\n        X = self.fc2(X)\n        return F.log_softmax(X, dim=1)","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:06.52983Z","iopub.execute_input":"2023-02-07T11:01:06.530226Z","iopub.status.idle":"2023-02-07T11:01:06.546231Z","shell.execute_reply.started":"2023-02-07T11:01:06.530152Z","shell.execute_reply":"2023-02-07T11:01:06.545301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Neural Network Predefined Parameters\nparams_model={\n        \"shape_in\": (1,512,512), \n        \"initial_filters\": 8,    \n        \"num_fc1\": 100,\n        \"dropout_rate\": 0.25,\n        \"num_classes\": 1}\n\n# Create instantiation of Network class\ncnn_model = Network(params_model)\n\n# define computation hardware approach (GPU/CPU)\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel = cnn_model.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:06.549098Z","iopub.execute_input":"2023-02-07T11:01:06.549355Z","iopub.status.idle":"2023-02-07T11:01:09.445877Z","shell.execute_reply.started":"2023-02-07T11:01:06.549332Z","shell.execute_reply":"2023-02-07T11:01:09.444686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%pip install torchsummary","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:09.447863Z","iopub.execute_input":"2023-02-07T11:01:09.448568Z","iopub.status.idle":"2023-02-07T11:01:19.275094Z","shell.execute_reply.started":"2023-02-07T11:01:09.44853Z","shell.execute_reply":"2023-02-07T11:01:19.273939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchsummary import summary\nsummary(cnn_model, input_size=(1, 512, 512),device=device.type)","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:19.276711Z","iopub.execute_input":"2023-02-07T11:01:19.277095Z","iopub.status.idle":"2023-02-07T11:01:25.10936Z","shell.execute_reply.started":"2023-02-07T11:01:19.277062Z","shell.execute_reply":"2023-02-07T11:01:25.108318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_func = nn.NLLLoss(reduction=\"sum\")","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:25.111687Z","iopub.execute_input":"2023-02-07T11:01:25.112805Z","iopub.status.idle":"2023-02-07T11:01:25.118811Z","shell.execute_reply.started":"2023-02-07T11:01:25.112735Z","shell.execute_reply":"2023-02-07T11:01:25.117837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import optim\nopt = optim.Adam(cnn_model.parameters(), lr=3e-4)\nlr_scheduler = ReduceLROnPlateau(opt, mode='min',factor=0.5, patience=20,verbose=1)","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:25.121953Z","iopub.execute_input":"2023-02-07T11:01:25.12303Z","iopub.status.idle":"2023-02-07T11:01:25.335413Z","shell.execute_reply.started":"2023-02-07T11:01:25.122949Z","shell.execute_reply":"2023-02-07T11:01:25.334341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Model","metadata":{}},{"cell_type":"code","source":"''' Helper Functions'''\n\n# Function to get the learning rate\ndef get_lr(opt):\n    for param_group in opt.param_groups:\n        return param_group['lr']\n\n# Function to compute the loss value per batch of data\ndef loss_batch(loss_func, output, target, opt=None):\n    \n    loss = loss_func(output, target) # get loss\n    pred = output.argmax(dim=1, keepdim=True) # Get Output Class\n    metric_b=pred.eq(target.view_as(pred)).sum().item() # get performance metric\n    \n    if opt is not None:\n        opt.zero_grad()\n        loss.backward()\n        opt.step()\n\n    return loss.item(), metric_b\n\n# Compute the loss value & performance metric for the entire dataset (epoch)\ndef loss_epoch(model,loss_func,dataset_dl,opt=None):\n    \n    run_loss=0.0 \n    t_metric=0.0\n    len_data=len(dataset_dl.dataset)\n\n    # internal loop over dataset\n    for xb, yb in dataset_dl:\n        # move batch to device\n        xb=xb.to(device)\n        yb=yb.to(device)\n        output=model(xb) # get model output\n        loss_b,metric_b=loss_batch(loss_func, output, yb, opt) # get loss per batch\n        run_loss+=loss_b        # update running loss\n\n        if metric_b is not None: # update running metric\n            t_metric+=metric_b    \n    \n    loss=run_loss/float(len_data)  # average loss value\n    metric=t_metric/float(len_data) # average metric value\n    \n    return loss, metric","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:25.33709Z","iopub.execute_input":"2023-02-07T11:01:25.337834Z","iopub.status.idle":"2023-02-07T11:01:25.34876Z","shell.execute_reply.started":"2023-02-07T11:01:25.337781Z","shell.execute_reply":"2023-02-07T11:01:25.347836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"params_train={\n \"train\": train_dl,\"val\": val_dl,\n \"epochs\": 50,\n \"optimiser\": optim.Adam(cnn_model.parameters(),\n                         lr=3e-4),\n \"lr_change\": ReduceLROnPlateau(opt,\n                                mode='min',\n                                factor=0.5,\n                                patience=20,\n                                verbose=0),\n \"f_loss\": nn.NLLLoss(reduction=\"sum\"),\n \"weight_path\": \"weights.pt\",\n \"check\": False, \n}\n","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:25.350216Z","iopub.execute_input":"2023-02-07T11:01:25.350627Z","iopub.status.idle":"2023-02-07T11:01:25.367826Z","shell.execute_reply.started":"2023-02-07T11:01:25.350591Z","shell.execute_reply":"2023-02-07T11:01:25.364772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.notebook import trange, tqdm\n\ndef train_val(model, params,verbose=False):\n    \n    # Get the parameters\n    epochs=params[\"epochs\"]\n    loss_func=params[\"f_loss\"]\n    opt=params[\"optimiser\"]\n    train_dl=params[\"train\"]\n    val_dl=params[\"val\"]\n    lr_scheduler=params[\"lr_change\"]\n    weight_path=params[\"weight_path\"]\n    \n    loss_history={\"train\": [],\"val\": []} # history of loss values in each epoch\n    metric_history={\"train\": [],\"val\": []} # histroy of metric values in each epoch\n    best_model_wts = copy.deepcopy(model.state_dict()) # a deep copy of weights for the best performing model\n    best_loss=float('inf') # initialize best loss to a large value\n    \n    ''' Train Model n_epochs '''\n    \n    for epoch in tqdm(range(epochs)):\n        \n        ''' Get the Learning Rate '''\n        current_lr=get_lr(opt)\n        if(verbose):\n            print('Epoch {}/{}, current lr={}'.format(epoch, epochs - 1, current_lr))\n        \n        '''\n        \n        Train Model Process\n        \n        '''\n        \n        model.train()\n        train_loss, train_metric = loss_epoch(model,loss_func,train_dl,opt)\n\n        # collect losses\n        loss_history[\"train\"].append(train_loss)\n        metric_history[\"train\"].append(train_metric)\n        \n        '''\n        \n        Evaluate Model Process\n        \n        '''\n        \n        model.eval()\n        with torch.no_grad():\n            val_loss, val_metric = loss_epoch(model,loss_func,val_dl)\n        \n        # store best model\n        if(val_loss < best_loss):\n            best_loss = val_loss\n            best_model_wts = copy.deepcopy(model.state_dict())\n            \n            # store weights into a local file\n            torch.save(model.state_dict(), weight_path)\n            if(verbose):\n                print(\"Copied best model weights!\")\n        \n        # collect loss and metric for validation dataset\n        loss_history[\"val\"].append(val_loss)\n        metric_history[\"val\"].append(val_metric)\n        \n        # learning rate schedule\n        lr_scheduler.step(val_loss)\n        if current_lr != get_lr(opt):\n            if(verbose):\n                print(\"Loading best model weights!\")\n            model.load_state_dict(best_model_wts) \n\n        if(verbose):\n            print(f\"train loss: {train_loss:.6f}, dev loss: {val_loss:.6f}, accuracy: {100*val_metric:.2f}\")\n            print(\"-\"*10) \n\n    # load best model weights\n    model.load_state_dict(best_model_wts)\n        \n    return model, loss_history, metric_history","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:25.370371Z","iopub.execute_input":"2023-02-07T11:01:25.371151Z","iopub.status.idle":"2023-02-07T11:01:25.388177Z","shell.execute_reply.started":"2023-02-07T11:01:25.371099Z","shell.execute_reply":"2023-02-07T11:01:25.387117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Process","metadata":{}},{"cell_type":"code","source":"params_train={\n \"train\": train_dl,\"val\": val_dl,\n \"epochs\": 50,\n \"optimiser\": optim.Adam(cnn_model.parameters(),lr=3e-4),\n \"lr_change\": ReduceLROnPlateau(opt,\n                                mode='min',\n                                factor=0.5,\n                                patience=20,\n                                verbose=0),\n \"f_loss\": nn.NLLLoss(reduction=\"sum\"),\n \"weight_path\": \"weights.pt\",\n}\n\n''' Actual Train / Evaluation of CNN Model '''\n# train and validate the model\n\ncnn_model,loss_hist,metric_hist=train_val(cnn_model,params_train)","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:25.390042Z","iopub.execute_input":"2023-02-07T11:01:25.390542Z","iopub.status.idle":"2023-02-07T11:01:45.904498Z","shell.execute_reply.started":"2023-02-07T11:01:25.390508Z","shell.execute_reply":"2023-02-07T11:01:45.90123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import seaborn as sns; sns.set(style='whitegrid')\n\nepochs=params_train[\"epochs\"]\n\nfig,ax = plt.subplots(1,2,figsize=(12,5))\n\nsns.lineplot(x=[*range(1,epochs+1)],y=loss_hist[\"train\"],ax=ax[0],label='loss_hist[\"train\"]')\nsns.lineplot(x=[*range(1,epochs+1)],y=loss_hist[\"val\"],ax=ax[0],label='loss_hist[\"val\"]')\nsns.lineplot(x=[*range(1,epochs+1)],y=metric_hist[\"train\"],ax=ax[1],label='metric_hist[\"train\"]')\nsns.lineplot(x=[*range(1,epochs+1)],y=metric_hist[\"val\"],ax=ax[1],label='metric_hist[\"val\"]')\nplt.title('Convergence History')","metadata":{"execution":{"iopub.status.busy":"2023-02-07T11:01:45.90773Z","iopub.status.idle":"2023-02-07T11:01:45.910199Z","shell.execute_reply.started":"2023-02-07T11:01:45.909929Z","shell.execute_reply":"2023-02-07T11:01:45.909955Z"},"trusted":true},"execution_count":null,"outputs":[]}]}