{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":6927,"databundleVersionId":45059,"sourceType":"competition"},{"sourceId":234720,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":200494,"modelId":222308}],"dockerImageVersionId":30822,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## **1. Import libaries and device setup**","metadata":{}},{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \n\nimport os\nimport zipfile\nfrom tqdm import tqdm\nimport random\n\nimport torch as t \nfrom torch import nn, optim\nfrom torch.utils.data import DataLoader, random_split\nfrom torchvision.utils import make_grid, draw_segmentation_masks\nfrom torchvision import transforms\n\nimport matplotlib.pyplot as plt\n\nimport requests\n\nfrom PIL import Image","metadata":{"_uuid":"1fcdcb8b-1748-40b9-82d3-a9598a35161a","_cell_guid":"f3e3e951-9670-4e50-847e-88f2299ec63e","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-01-17T19:10:05.380501Z","iopub.execute_input":"2025-01-17T19:10:05.380902Z","iopub.status.idle":"2025-01-17T19:10:10.329275Z","shell.execute_reply.started":"2025-01-17T19:10:05.380863Z","shell.execute_reply":"2025-01-17T19:10:10.328247Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEVICE = 'cuda' if t.cuda.is_available() else 'cpu'\nprint(f\"Current available device = {DEVICE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-17T19:10:10.330603Z","iopub.execute_input":"2025-01-17T19:10:10.331105Z","iopub.status.idle":"2025-01-17T19:10:10.389691Z","shell.execute_reply.started":"2025-01-17T19:10:10.331067Z","shell.execute_reply":"2025-01-17T19:10:10.388834Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"gpu_counts = t.cuda.device_count()\nfor i in range(gpu_counts):\n    print(f\"Device name = {t.cuda.get_device_name(i)}\")\n    print(f\"Device properties = {t.cuda.get_device_properties(i)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-17T19:10:10.391557Z","iopub.execute_input":"2025-01-17T19:10:10.391794Z","iopub.status.idle":"2025-01-17T19:10:10.435969Z","shell.execute_reply.started":"2025-01-17T19:10:10.391773Z","shell.execute_reply":"2025-01-17T19:10:10.434952Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **2. Data preprations and importing libaries**","metadata":{}},{"cell_type":"code","source":"for dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))","metadata":{"_uuid":"7601f8f4-0340-4531-88d8-9d1b5c2eddde","_cell_guid":"62b64cc6-cc6c-468d-9bfb-96c5bd696bfb","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-01-17T19:10:15.425577Z","iopub.execute_input":"2025-01-17T19:10:15.425950Z","iopub.status.idle":"2025-01-17T19:10:15.436085Z","shell.execute_reply.started":"2025-01-17T19:10:15.425917Z","shell.execute_reply":"2025-01-17T19:10:15.434923Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if not os.path.exists('/kaggle/working/dataset/train_masks'):\n    print(\"Train mask does not exists... Extracting the data\")\n    with zipfile.ZipFile('/kaggle/input/carvana-image-masking-challenge/train_masks.zip') as z:\n        z.extractall('/kaggle/working/dataset/')\n\nif not os.path.exists('/kaggle/working/dataset/train'):\n    print(\"Train does not exists... Extracting the data\")\n    with zipfile.ZipFile('/kaggle/input/carvana-image-masking-challenge/train.zip') as z:\n        z.extractall('/kaggle/working/dataset/')\n\nif not os.path.exists('/kaggle/working/dataset/test'):\n    print(\"Test does not exists... Extracting the data\")\n    with zipfile.ZipFile('/kaggle/input/carvana-image-masking-challenge/test.zip') as z:\n        z.extractall('/kaggle/working/dataset/')","metadata":{"_uuid":"114c03cc-fc07-4d9e-aa79-e7b5fb066732","_cell_guid":"1ac37b5d-8f70-4e23-90b6-b560e546f99a","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-01-17T19:10:15.701429Z","iopub.execute_input":"2025-01-17T19:10:15.701756Z","iopub.status.idle":"2025-01-17T19:13:10.656360Z","shell.execute_reply.started":"2025-01-17T19:10:15.701732Z","shell.execute_reply":"2025-01-17T19:13:10.655167Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"unet_parts_git = 'https://raw.githubusercontent.com/Dhamu785/U-Nets/refs/heads/main/01_U-net/PyTorch/01_Uygar%20Kurt/unet_parts.py'\nunet_git = 'https://raw.githubusercontent.com/Dhamu785/U-Nets/refs/heads/main/01_U-net/PyTorch/01_Uygar%20Kurt/unet.py'\ndataset_git = 'https://raw.githubusercontent.com/Dhamu785/U-Nets/refs/heads/main/01_U-net/PyTorch/01_Uygar%20Kurt/dataset.py'\nunet_a = 'https://raw.githubusercontent.com/Dhamu785/U-Nets/refs/heads/main/01_U-net/PyTorch/01_Uygar%20Kurt/unet_a.py'\n\nif os.path.exists('unet_parts.py'):\n    print('Removing the file...')\n    os.remove('unet_parts.py')\n    print(\"Downloading unet_parts.py....\")\n    r = requests.get(unet_parts_git).content\n    with open(\"unet_parts.py\", 'wb') as f:\n        f.write(r)\nelse:\n    print(\"Downloading unet_parts.py....\")\n    r = requests.get(unet_parts_git).content\n    with open(\"unet_parts.py\", 'wb') as f:\n        f.write(r)\n\nif os.path.exists('unet_a.py'):\n    print('Removing the file...')\n    os.remove('unet_a.py')\n    print(\"Downloading unet_a.py....\")\n    r = requests.get(unet_a).content\n    with open(\"unet_a.py\", 'wb') as f:\n        f.write(r)\nelse:\n    print(\"Downloading unet_a.py....\")\n    r = requests.get(unet_a).content\n    with open(\"unet_a.py\", 'wb') as f:\n        f.write(r)\n\n\n\nif os.path.exists('unet.py'):\n    print('Removing the file...')\n    os.remove('unet.py')\n    print(\"Downloading unet.py....\")\n    r = requests.get(unet_git).content\n    with open('unet.py', 'wb') as f:\n        f.write(r)\nelse:\n    print(\"Downloading unet.py....\")\n    r = requests.get(unet_git).content\n    with open('unet.py', 'wb') as f:\n        f.write(r)\n\nif os.path.exists('dataset.py'):\n    print('Removing the file...')\n    os.remove('dataset.py')\n    print(\"Downloading dataset.py....\")\n    r = requests.get(dataset_git).content\n    with open('dataset.py', 'wb') as f:\n        f.write(r)\nelse:\n    print(\"Downloading dataset.py....\")\n    r = requests.get(dataset_git).content\n    with open('dataset.py', 'wb') as f:\n        f.write(r)","metadata":{"_uuid":"efcfa9c2-8b58-4a53-af92-50a4bc72a457","_cell_guid":"e7fd9b0e-8f15-4b4a-9b74-1f3c95c88561","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-01-17T19:13:10.657938Z","iopub.execute_input":"2025-01-17T19:13:10.658309Z","iopub.status.idle":"2025-01-17T19:13:11.285019Z","shell.execute_reply.started":"2025-01-17T19:13:10.658274Z","shell.execute_reply":"2025-01-17T19:13:11.284089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from dataset import seg_dataset\nfrom unet_a import unet","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-17T19:13:33.041992Z","iopub.execute_input":"2025-01-17T19:13:33.042347Z","iopub.status.idle":"2025-01-17T19:13:33.049862Z","shell.execute_reply.started":"2025-01-17T19:13:33.042317Z","shell.execute_reply":"2025-01-17T19:13:33.049004Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **3. Loss and accuracy**","metadata":{}},{"cell_type":"code","source":"import torch.nn.functional as F","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-17T19:13:36.604179Z","iopub.execute_input":"2025-01-17T19:13:36.604505Z","iopub.status.idle":"2025-01-17T19:13:36.608563Z","shell.execute_reply.started":"2025-01-17T19:13:36.604478Z","shell.execute_reply":"2025-01-17T19:13:36.607554Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def loss_iou(y_pred, y_true, inf):\n    if not inf:\n        if not y_pred.requires_grad:\n            raise ValueError(\"y_pred should have gradient tracking\")\n    \n    device = y_pred.device\n    # binary_pred = t.where(y_pred <= 0, t.zeros_like(y_pred, device=device, requires_grad=True), t.ones_like(y_pred, device=device, requires_grad=True))\n    y_true = t.where(y_true <= 0, t.zeros_like(y_pred, device=device), t.ones_like(y_pred, device=device))\n    \n    # y_pred = F.sigmoid(y_pred)\n    \n    intersection = t.abs((y_pred.view((-1)) * y_true.view((-1))).sum().float())\n    union = t.abs((y_pred.sum() + y_true.sum()).float())\n    # print(f\"intersection = {t.abs(intersection)}, union={union}\")\n\n    iou = (t.abs(intersection) + 1e-5) / ((union + 1e-5) - t.abs(intersection))\n    iou_loss = 1 - iou\n    return iou_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-17T19:13:37.288163Z","iopub.execute_input":"2025-01-17T19:13:37.288521Z","iopub.status.idle":"2025-01-17T19:13:37.294918Z","shell.execute_reply.started":"2025-01-17T19:13:37.288496Z","shell.execute_reply":"2025-01-17T19:13:37.293765Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def acc_iou(y_pred, y_true, inf):\n    ls = loss_iou(y_pred, y_true, inf)\n    return 1-ls","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-17T19:13:38.222100Z","iopub.execute_input":"2025-01-17T19:13:38.222419Z","iopub.status.idle":"2025-01-17T19:13:38.226973Z","shell.execute_reply.started":"2025-01-17T19:13:38.222395Z","shell.execute_reply":"2025-01-17T19:13:38.225738Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **4. Train test loop**","metadata":{"_uuid":"7f208690-136c-4a7c-91f0-a84c8c1500e9","_cell_guid":"39ddd028-43ca-434f-9aad-eaa29e6d596e","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-12-31T10:33:58.428724Z","iopub.execute_input":"2024-12-31T10:33:58.429146Z","iopub.status.idle":"2024-12-31T10:33:58.436249Z","shell.execute_reply.started":"2024-12-31T10:33:58.429115Z","shell.execute_reply":"2024-12-31T10:33:58.434920Z"}}},{"cell_type":"code","source":"img_path = '/kaggle/working/dataset/train'\nimg_msk_path = '/kaggle/working/dataset/train_masks'\n\nimg_name = os.listdir(img_path)[100]\nmsk_list = os.listdir(img_msk_path)\nmsk_idx = msk_list.index(img_name.split('.')[0]+'_mask.gif')\nmsk_name = msk_list[msk_idx]\n# print(img_name, msk_name)\nsingle_img_path = os.path.join(img_path, img_name)\nsingle_msk_path = os.path.join(img_msk_path, msk_name)\n\ntransforms_pipe = transforms.Compose([\n            transforms.Resize((512,512)),\n            transforms.ToTensor()\n        ])\nimg = Image.open(single_img_path).convert('RGB')\nmsk = Image.open(single_msk_path).convert('L')\n\nimg_transformed = transforms_pipe(img)\nmsk_transformed = transforms_pipe(msk)\n# print(img_transformed.shape, msk_transformed.shape)\nmsk_transformed = t.where(msk_transformed <= 0, t.zeros_like(msk_transformed, device='cpu'), t.ones_like(msk_transformed, device='cpu'))\nprint(t.unique(msk_transformed))\nplt.subplot(1,2,1)\nplt.imshow(img_transformed.permute(1, 2, 0).to('cpu'))\nplt.subplot(1,2,2)\nplt.imshow(msk_transformed[0].to('cpu'), cmap='gray');","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-17T19:13:43.356373Z","iopub.execute_input":"2025-01-17T19:13:43.356730Z","iopub.status.idle":"2025-01-17T19:13:44.071062Z","shell.execute_reply.started":"2025-01-17T19:13:43.356700Z","shell.execute_reply":"2025-01-17T19:13:44.070038Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LEARNING_RATE = 1e-4\nBATCH_SIZE = 4\nEPOCHS = 20\nDATA_PATH = \"/kaggle/working/dataset/\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-17T19:13:51.321165Z","iopub.execute_input":"2025-01-17T19:13:51.321494Z","iopub.status.idle":"2025-01-17T19:13:51.326034Z","shell.execute_reply.started":"2025-01-17T19:13:51.321470Z","shell.execute_reply":"2025-01-17T19:13:51.324662Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_img(model):\n    img = img_transformed.unsqueeze(0).to(DEVICE)\n    # print(img.shape)\n    pred = model(img)\n    plt.figure(figsize=(15,5))\n    plt.subplot(1,3,1)\n    plt.imshow(img_transformed.permute(1, 2, 0).to('cpu'))\n    plt.axis('off')\n    plt.title('Original img')\n    plt.subplot(1,3,2)\n    plt.imshow(msk_transformed[0].to('cpu'), cmap='gray')\n    plt.title('Original mask')\n    plt.axis('off')\n    plt.subplot(1,3,3)\n    plt.imshow(pred[0][0].detach().to('cpu').numpy(), cmap='gray')\n    plt.title('Predicted mask')\n    plt.axis('off')\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-17T19:13:51.911271Z","iopub.execute_input":"2025-01-17T19:13:51.911600Z","iopub.status.idle":"2025-01-17T19:13:51.918065Z","shell.execute_reply.started":"2025-01-17T19:13:51.911577Z","shell.execute_reply":"2025-01-17T19:13:51.916935Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data = seg_dataset(DATA_PATH)\ngenerator = t.Generator().manual_seed(42)\n\ntrain_dataset, val_dataset = random_split(dataset=data, lengths=(0.8, 0.2), generator=generator)\n\ntrain_dataloader = DataLoader(train_dataset, BATCH_SIZE, True)\nval_dataloader = DataLoader(val_dataset, BATCH_SIZE, True)\n\nmodel = unet(in_channel=3, num_classes=1).to(DEVICE)\noptimizer = optim.Adam(params = model.parameters(), lr=LEARNING_RATE)\n# loss = nn.BCEWithLogitsLoss()\n\nhistory = {'loss_per_epoch_train':list(), 'loss_per_epoch_test':list(), \n           'acc_per_epoch_train':list(), 'acc_per_epoch_test':list()}\n\nfor epoch in range(EPOCHS):\n    model.train()\n    train_loss_per_batch = 0\n    train_acc_per_batch = 0\n    epoch_pbar = tqdm(range(len(train_dataloader)), desc=\"Batch processing\",unit=\"batchs\")\n    for idx,batch in enumerate(train_dataloader):\n        img = batch[0].float().to(DEVICE)\n        mask = batch[1].float().to(DEVICE)\n\n        # 1. Forward pass\n        y_pred = model(img)\n        # print(y_pred.min(), y_pred.max())\n        # 2. Calculate the loss\n        ls = loss_iou(y_pred, mask, False)\n        # print(\"Loss = \", ls)\n\n        acc = acc_iou(y_pred, mask, False)\n        train_loss_per_batch += ls.item()\n        train_acc_per_batch += acc.item()\n\n        optimizer.zero_grad()\n        ls.backward()\n        optimizer.step()\n        epoch_pbar.update(1)\n    epoch_pbar.close()\n    train_loss_per_batch /= idx+1\n    train_acc_per_batch /= idx+1\n    history['loss_per_epoch_train'].append(train_loss_per_batch)\n    history['acc_per_epoch_train'].append(train_acc_per_batch)\n    \n    model.eval()\n    test_loss_per_batch = 0\n    test_acc_per_batch = 0\n    with t.inference_mode():\n        for idx, batch in enumerate(val_dataloader):\n            img = batch[0].float().to(DEVICE)\n            mask = batch[1].float().to(DEVICE)\n\n            y_pred_test = model(img)\n            test_ls = loss_iou(y_pred_test, mask, True)\n            test_acc = acc_iou(y_pred_test, mask, True)\n\n            test_loss_per_batch += test_ls.item()\n            test_acc_per_batch += test_acc.item()\n        \n        test_loss_per_batch /= idx+1\n        test_acc_per_batch /= idx+1\n        plot_img(model)\n\n    history['loss_per_epoch_test'].append(test_loss_per_batch)\n    history['acc_per_epoch_test'].append(test_acc_per_batch)\n\n    print(f\"{epoch+1} / {EPOCHS} | train_loss = {train_loss_per_batch:.4f} | train_acc = {train_acc_per_batch:.4f} | test_loss = {test_loss_per_batch:.4f} | test_acc = {test_acc_per_batch:.4f}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **5. Model saving**","metadata":{}},{"cell_type":"code","source":"t.save(model.state_dict(), 'model_wt-iou.pt')\nt.save(model, 'entire-model-iou.pt')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **6. Model analysis**","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(13,5))\nplt.subplot(1,2,1)\nplt.plot(history['acc_per_epoch_train'], color = 'orange')\nplt.plot(history['acc_per_epoch_test'], color = 'green')\nplt.legend(['train', 'test'])\nplt.title(\"Accuracy\")\nplt.subplot(1,2,2)\nplt.plot(history['loss_per_epoch_train'], color = 'orange')\nplt.plot(history['loss_per_epoch_test'], color = 'green')\nplt.legend(['train', 'test'])\nplt.title(\"Loss\")\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **7. Model predictions**","metadata":{}},{"cell_type":"code","source":"# Model loading\n\n## Type-1\n## model_from_saved_1 = t.load('/kaggle/input/instance-segmentation-carvana-image-challenge/pytorch/default/1/Carvana Image_mdl.pt')\n## Type-2\nmodel_from_saved_2 = unet(in_channel=3, num_classes=1).to(DEVICE)\nmodel_from_saved_2.load_state_dict(t.load('model_wt-iou.pt', weights_only=True, map_location=t.device(DEVICE)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-17T19:14:00.777896Z","iopub.execute_input":"2025-01-17T19:14:00.778231Z","iopub.status.idle":"2025-01-17T19:14:02.867515Z","shell.execute_reply.started":"2025-01-17T19:14:00.778193Z","shell.execute_reply":"2025-01-17T19:14:02.866543Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load the random test data\n\nclass get_data:\n    def __init__(self, no_of_sample):\n        self.train_img_path = '/kaggle/working/dataset/train'\n        self.train_msk_path = '/kaggle/working/dataset/train_masks'\n        test_files = os.listdir(self.train_img_path)\n        random_100 = random.choices(test_files, k=no_of_sample)\n        \n        transformed_img = []\n        transformed_msk = []\n        \n        for i in random_100:\n            transformed_img.append(transforms_pipe(Image.open(os.path.join(self.train_img_path, i)).convert('RGB')))\n            transformed_msk.append(transforms_pipe(Image.open(os.path.join(self.train_msk_path, i.split('.')[0]+\"_mask.gif\")).convert('L')))\n            \n        transformed_img = np.array(transformed_img)\n        transformed_msk = np.array(transformed_msk)\n        self.img = t.from_numpy(transformed_img)\n        self.msk = t.from_numpy(transformed_msk)\n\n    def mask_overlay(self, preds):\n        preds = preds >= 0.5\n        overlays = []\n        for i in range(len(self.img)):\n            overlays.append(draw_segmentation_masks(self.img[i], preds[i], colors='blue'))\n        overlays = t.from_numpy(np.array(overlays))\n        return overlays\n\n    def make_grids(self, preds, overlay):\n        predictions = make_grid(preds, 10, 1, pad_value = 2).moveaxis(0, 2)\n        seg_overlay = make_grid(overlay, 10, 1, pad_value = 2).moveaxis(0, 2)\n        images = make_grid(self.img, 10, 1, pad_value = 2).moveaxis(0, 2)\n        masks = make_grid(self.msk, 10, 1, pad_value = 2).moveaxis(0, 2)\n\n        return images, masks, predictions, seg_overlay","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-17T19:14:57.151904Z","iopub.execute_input":"2025-01-17T19:14:57.152288Z","iopub.status.idle":"2025-01-17T19:14:57.161248Z","shell.execute_reply.started":"2025-01-17T19:14:57.152253Z","shell.execute_reply":"2025-01-17T19:14:57.159995Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data = get_data(8)\n\nwith t.inference_mode():\n    preds = model_from_saved_2(data.img.to(DEVICE))\n\nseg_overlay = data.mask_overlay(preds)\nimages, masks, predictions, seg_overlay = data.make_grids(preds, seg_overlay)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-17T19:15:04.891800Z","iopub.execute_input":"2025-01-17T19:15:04.892128Z","iopub.status.idle":"2025-01-17T19:15:05.788220Z","shell.execute_reply.started":"2025-01-17T19:15:04.892104Z","shell.execute_reply":"2025-01-17T19:15:05.787181Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_data = [images, masks, predictions, seg_overlay]\ntitles = ['Original images', 'Labels', 'Predictions', 'Prediction mask applied']\nfor i in range(len(plot_data)):\n    plt.figure(figsize=(20, 30))\n    plt.subplot(4,1,i+1)\n    plt.imshow(plot_data[i].cpu().numpy())\n    plt.title(label = titles[i], fontweight=10, fontstyle='normal')\n    plt.axis('off')\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-17T19:16:52.166033Z","iopub.execute_input":"2025-01-17T19:16:52.166398Z","iopub.status.idle":"2025-01-17T19:16:54.516218Z","shell.execute_reply.started":"2025-01-17T19:16:52.166371Z","shell.execute_reply":"2025-01-17T19:16:54.515155Z"}},"outputs":[],"execution_count":null}]}