{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"}],"dockerImageVersionId":30699,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import cv2\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport numpy as np\nimport torch.nn as nn3\nfrom torchvision.models import resnet50\nimport pandas as pd\nimport os\nfrom torch.cuda.amp import GradScaler, autocast\nfrom glob import glob\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport pandas as pd\nimport torch.nn as nn\n\n\n\nprint(\"IMPORT OK !!!\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MyDataset(Dataset):\n    def __init__(self, meta, image_size=(512, 512)):\n        self.meta = meta\n        self.image_size = image_size\n\n    def __len__(self):\n        return len(self.meta)\n\n    def __getitem__(self, index):\n        f = self.meta[index]\n        if f is None:\n            raise ValueError(\"File path not found in metadata\")\n        \n        # Load image\n        image = cv2.imread(f, cv2.IMREAD_GRAYSCALE)\n        image = cv2.resize(image, self.image_size)\n        image = torch.from_numpy(image).float().unsqueeze(0)  # Convert to tensor and add channel dimension\n        image = norm_by_percentile(image)\n\n        \n        \n        # Load mask\n        maskfile = f.replace('images', 'labels')\n        if os.path.exists(maskfile):\n            mask = cv2.imread(maskfile, cv2.IMREAD_GRAYSCALE)\n            mask = cv2.resize(mask, self.image_size)\n            mask = torch.from_numpy(mask).float().unsqueeze(0)  # Convert to tensor and add channel dimension\n            mask = norm_by_percentile(mask)\n\n        else:\n            mask = torch.zeros_like(image)  # Create a zero mask with same shape as image\n        \n        return image, mask\n    \n   \n\n\nclass DotDict(dict):\n    \"\"\"dot.notation access to dictionary attributes\"\"\"\n    def __getattr__(self, attr):\n        return self.get(attr)\n    \n    __setattr__= dict.__setitem__\n    __delattr__= dict.__delitem__\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_images(images, masks, n=5):\n    \"\"\"\n    Display a batch of images and their corresponding masks.\n    \n    Args:\n    - images (torch.Tensor): Batch of images.\n    - masks (torch.Tensor): Batch of masks.\n    - n (int): Number of images to display. Default is 5.\n    \"\"\"\n    images = images.cpu().numpy()\n    masks = masks.cpu().numpy()\n    \n    batch_size = images.shape[0]\n    n = min(n, batch_size)  # Ensure n does not exceed batch size\n\n    plt.figure(figsize=(15, 10))\n    for i in range(n):\n        plt.subplot(n, 2, 2 * i + 1)\n        plt.imshow(images[i, 0], cmap='gray')\n        plt.title('Image')\n        plt.axis('off')\n        \n        plt.subplot(n, 2, 2 * i + 2)\n        plt.imshow(masks[i, 0], cmap='gray')\n        plt.title('Mask')\n        plt.axis('off')\n    plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def norm_by_percentile(x, low=10, high=99.8, alpha=0.01):\n    xmin = np.percentile(x, low)\n    xmax = np.percentile(x, high)\n    x = (x - xmin) / (xmax - xmin + 1e-5)  # Avoid division by zero\n    x = np.clip(x, 0, 1)  # Ensure values are in the range [0, 1]\n    if 1:\n        x[x > 1] = (x[x > 1] - 1) * alpha + 1\n        x[x < 0] = x[x < 0] * alpha\n    return x\n    #return x / 65536.0","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_checkpoint(epoch, model, optimizer, loss, file_path):\n    torch.save({\n        'epoch': epoch,\n        'model_state_dict': model.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'loss': loss,\n    }, file_path)\n\ndef load_checkpoint(file_path, model, optimizer):\n    checkpoint = torch.load(file_path)\n    model.load_state_dict(checkpoint['model_state_dict'])\n    optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n    epoch = checkpoint['epoch']\n    loss = checkpoint['loss']\n    return epoch, model, optimizer, loss","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def init_weights(m):\n    if isinstance(m, nn.Conv2d) or isinstance(m, nn.Linear):\n        nn.init.kaiming_normal_(m.weight, nonlinearity='relu')\n        if m.bias is not None:\n            nn.init.constant_(m.bias, 0)\n\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.models as models\n\nclass Net(nn.Module):\n    def __init__(self, weights=None):\n    #def __init__(self, weights=None):\n\n        super(Net, self).__init__()\n        \n        self.backbone = models.resnet50(weights=None)\n       \n        self.initialize_weights(self.backbone)\n        \n        self.backbone.conv1 = nn.Conv2d(1, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n        \n        self.encoder = nn.Sequential(*list(self.backbone.children())[:-2])\n        \n        self.upconv1 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.conv1 = nn.Conv2d(2048 + 1024, 1024, kernel_size=3, padding=1)\n        \n        self.upconv2 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.conv2 = nn.Conv2d(1024 + 512, 512, kernel_size=3, padding=1)\n        \n        self.upconv3 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.conv3 = nn.Conv2d(512 + 256, 256, kernel_size=3, padding=1)\n        \n        self.upconv4 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.conv4 = nn.Conv2d(256 + 64, 64, kernel_size=3, padding=1)\n        \n        self.upconv5 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.conv5 = nn.Conv2d(64, 32, kernel_size=3, padding=1)\n        \n        self.vessel_out = nn.Conv2d(32, 1, kernel_size=1)\n        self.kidney_out = nn.Conv2d(32, 1, kernel_size=1)\n    \n    def forward(self, x):\n        # Encoder\n        enc1 = self.backbone.conv1(x)\n        enc1 = self.backbone.bn1(enc1)\n        enc1 = self.backbone.relu(enc1)\n        enc1 = self.backbone.maxpool(enc1)\n        \n        enc2 = self.backbone.layer1(enc1)\n        enc3 = self.backbone.layer2(enc2)\n        enc4 = self.backbone.layer3(enc3)\n        enc5 = self.backbone.layer4(enc4)\n        \n        # Decoder with debugging prints\n        dec1 = self.upconv1(enc5)\n        if dec1.size()[2:] != enc4.size()[2:]:\n            dec1 = F.interpolate(dec1, size=enc4.shape[2:], mode='bilinear', align_corners=True)\n        dec1 = torch.cat((dec1, enc4), dim=1)\n        dec1 = self.conv1(dec1)\n        \n        dec2 = self.upconv2(dec1)\n        if dec2.size()[2:] != enc3.size()[2:]:\n            dec2 = F.interpolate(dec2, size=enc3.shape[2:], mode='bilinear', align_corners=True)\n        dec2 = torch.cat((dec2, enc3), dim=1)\n        dec2 = self.conv2(dec2)\n        \n        dec3 = self.upconv3(dec2)\n        if dec3.size()[2:] != enc2.size()[2:]:\n            dec3 = F.interpolate(dec3, size=enc2.shape[2:], mode='bilinear', align_corners=True)\n        dec3 = torch.cat((dec3, enc2), dim=1)\n        dec3 = self.conv3(dec3)\n        \n        dec4 = self.upconv4(dec3)\n        if dec4.size()[2:] != enc1.size()[2:]:\n            dec4 = F.interpolate(dec4, size=enc1.shape[2:], mode='bilinear', align_corners=True)\n        dec4 = torch.cat((dec4, enc1), dim=1)\n        dec4 = self.conv4(dec4)\n        \n        dec5 = self.upconv5(dec4)\n        if dec5.size()[2:] != x.size()[2:]:\n            dec5 = F.interpolate(dec5, size=x.shape[2:], mode='bilinear', align_corners=True)\n        dec5 = self.conv5(dec5)\n        \n        vessel = self.vessel_out(dec5)\n        kidney = self.kidney_out(dec5)\n        \n        return vessel, kidney\n    \n    def initialize_weights(self, model):\n        for m in model.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.BatchNorm2d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n\nprint(\"NET OK!!!\")\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(net, train_loader, val_loader, criterion, optimizer, start_epoch=0, total_epochs=10, accumulation_steps=4, gradient_clip=1.0):\n    net.train().cuda()\n    scaler = GradScaler()\n\n    best_val_loss = float('inf')\n\n    for epoch in range(start_epoch, total_epochs):\n        epoch_loss = 0\n        optimizer.zero_grad()\n        for i, (images, masks) in enumerate(train_loader):\n            images, masks = images.cuda(), masks.cuda()\n            \n            # Check for NaNs in input data\n            if torch.isnan(images).any() or torch.isnan(masks).any():\n                print(\"NaNs found in input data\")\n                continue\n            \n            # Show the first batch of images and masks at the start of each epoch\n            if i == 0:\n                show_images(images, masks, n=1)\n\n            with autocast():\n                vessel_pred, kidney_pred = net(images)\n                \n                if torch.isnan(vessel_pred).any() or torch.isnan(kidney_pred).any():\n                    print(\"NaNs found in model outputs\")\n                    continue\n\n                loss_vessel = criterion(vessel_pred, masks)\n                loss_kidney = criterion(kidney_pred, masks)\n                loss = loss_vessel + loss_kidney\n\n            scaler.scale(loss).backward()\n\n            if (i + 1) % accumulation_steps == 0:\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(net.parameters(), gradient_clip)\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n\n            epoch_loss += loss.item()\n        \n        print(f'Epoch [{epoch + 1}/{total_epochs}], Loss: {epoch_loss / len(train_loader)}')\n        val_loss = validate(net, val_loader, criterion)\n        print(f'Epoch [{epoch + 1}/{total_epochs}], Validation Loss: {val_loss}')\n\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            torch.save({'epoch': epoch, 'model_state_dict': net.state_dict(), 'optimizer_state_dict': optimizer.state_dict()}, 'best_model.pth')\n\n    print(\"Training Complete\")\n    return net\n\ndef validate(net, val_loader, criterion):\n    net.eval()\n    val_loss = 0\n    with torch.no_grad():\n        for images, masks in val_loader:\n            images, masks = images.cuda(), masks.cuda()\n            vessel_pred, kidney_pred = net(images)\n            loss_vessel = criterion(vessel_pred, masks)\n            loss_kidney = criterion(kidney_pred, masks)\n            loss = loss_vessel + loss_kidney\n            val_loss += loss.item()\n    return val_loss / len(val_loader)\n\nprint(\"TRAIN OK!!!\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\n\n\n# Load the model and optimizer\nnet = Net(weights=None).cuda()\n#net = Net(weights=None).cuda()\n\n\n\n#net.apply(init_weights)\noptimizer = torch.optim.Adam(net.parameters(), lr=1e-5, weight_decay=1e-5)\n#optimizer = torch.optim.RMSprop(net.parameters(), lr=1e-7, weight_decay=1e-5)\n\n\n# Collect all image files from the training directory\n\ntrain_files = []\n\n# List of directories or patterns you want to include\npatterns = [\n    '/kaggle/input/blood-vessel-segmentation/train/kidney_1_dense/images/*.tif'\n    #,'/kaggle/input/blood-vessel-segmentation/train/kidney_1_voi/images/*.tif',\n    #'/kaggle/input/blood-vessel-segmentation/train/kidney_3_sparse/images/*.tif'\n    \n]\n\nfor pattern in patterns:\n    train_files.extend(glob.glob(pattern, recursive=True))\n\n\n\n# Create metadata for all images\ntrain_meta = [\n    DotDict(\n        file=train_files,\n        shape=(len(train_files), 512, 512),\n        id=[os.path.splitext(os.path.basename(f))[0] for f in train_files]\n    )\n]\n\n# Prepare training DataLoader\ntrain_dataset = MyDataset(train_files)  \ntrain_loader = DataLoader(train_dataset, batch_size=16, shuffle=True)\n\n\n# Prepare validation DataLoader\nval_files = glob.glob('/kaggle/input/blood-vessel-segmentation/train/kidney_2/images/*.tif', recursive=True)\nval_dataset = MyDataset(val_files)\nval_loader = DataLoader(val_dataset, batch_size=16, shuffle=False)\n\n\ncriterion = nn.BCEWithLogitsLoss()\n\n# Define the number of epochs you want to train for\ntotal_epochs = 1\n#checkpoint_dir = './checkpoints'\n\n# Load latest checkpoint if exists\nstart_epoch = 0\n#checkpoint_path = os.path.join(checkpoint_dir, 'latest_checkpoint.pth')\n#if os.path.isfile(checkpoint_path):\n   # start_epoch, net, optimizer, _ = load_checkpoint(checkpoint_path, net, optimizer)\n    #print(f\"Resuming training from epoch {start_epoch}\")\n\n# Continue training for additional epochs\ntrain(net, train_loader, val_loader, criterion, optimizer, start_epoch=start_epoch, total_epochs=total_epochs)\n\n","metadata":{"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install connected-components-3d","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from timeit import default_timer as timer\nimport datetime\nimport cc3d\nimport glob\nimport torch\nimport cv2\nimport numpy as np\n\ndef choose_biggest_object(mask, threshold):\n    mask = ((mask > threshold) * 255).astype(np.uint8)\n    num_label, label, stats, centroid = cv2.connectedComponentsWithStats(mask, connectivity=8)\n    max_label = -1\n    max_area = -1\n    for l in range(1, num_label):\n        if stats[l, cv2.CC_STAT_AREA] >= max_area:\n            max_area = stats[l, cv2.CC_STAT_AREA]\n            max_label = l\n    processed = (label == max_label).astype(np.uint8)\n    return processed\n\ndef remove_small_objects(mask, min_size, threshold):\n    mask = ((mask > threshold) * 255).astype(np.uint8)\n    num_label, label, stats, centroid = cv2.connectedComponentsWithStats(mask, connectivity=8)\n    processed = np.zeros_like(mask)\n    for l in range(1, num_label):\n        if stats[l, cv2.CC_STAT_AREA] >= min_size:\n            processed[label == l] = 1\n    return processed\n\ndef rle_encode(mask):\n    pixel = mask.flatten()\n    pixel = np.concatenate([[0], pixel, [0]])\n    run = np.where(pixel[1:] != pixel[:-1])[0] + 1\n    run[1::2] -= run[::2]\n    rle = ' '.join(str(r) for r in run)\n    if rle == '':\n        rle = '1 0'\n    return rle\n\n#-------------------------------\n\nnet = Net()\nnet = net.eval().cuda()\n\ndef dice_coefficient(pred, target):\n    intersection = np.sum(pred * target)\n    return (2. * intersection) / (np.sum(pred) + np.sum(target))\n\ndef iou(pred, target):\n    intersection = np.sum(pred * target)\n    union = np.sum(pred) + np.sum(target) - intersection\n    return intersection / union\n\ndef evaluate(predictions, ground_truths):\n    dice_scores = []\n    iou_scores = []\n    for pred, gt in zip(predictions, ground_truths):\n        dice_scores.append(dice_coefficient(pred, gt))\n        iou_scores.append(iou(pred, gt))\n    return {\n        'Dice Coefficient': np.mean(dice_scores),\n        'IoU': np.mean(iou_scores)\n    }\n\ndef make_predictions():\n    predictions = []\n    net.eval()\n    for d in test_images:\n        volume = [cv2.imread(f, cv2.IMREAD_GRAYSCALE) for f in d['file']]\n        volume = np.stack(volume)\n        D, H, W = volume.shape\n        \n        predict = np.zeros(d['shape'], dtype=np.float16)\n        axes = [0, 1, 2]  # Axes to consider for TTA\n        for axis in axes:\n            loader = np.array_split(np.arange((D, H, W)[axis]), max(1, int((D, H, W)[axis] // cfg['batch_size'])))\n            num_valid = len(loader)\n            \n            B = 0 \n            start_timer = timer()\n            for t in range(num_valid):\n                print(f'\\r validation: {t}/{num_valid}', end='', flush=True)\n                \n                if axis == 0:\n                    image = volume[loader[t].tolist()]\n                if axis == 1:\n                    image = volume[:, loader[t].tolist()]\n                    image = image.transpose(1, 0, 2)\n                if axis == 2:\n                    image = volume[:, :, loader[t].tolist()]\n                    image = image.transpose(2, 0, 1)\n\n                batch_size, bh, bw = image.shape\n                m = image.reshape(batch_size, -1)\n                m = (m - m.min(keepdims=True)) / (m.max(keepdims=True) - m.min(keepdims=True) + 0.001)\n                m = m.reshape(batch_size, bh, bw)\n                m = np.ascontiguousarray(m)\n                image = torch.from_numpy(m).float().cuda().unsqueeze(1)\n\n                counter = 0\n                vessel, kidney = 0, 0\n                image = image.cuda()\n                with torch.cuda.amp.autocast(enabled=True):\n                    with torch.no_grad():\n                        v, k = net(image)\n                        vessel += v\n                        kidney += k\n                        counter += 1\n\n                        v, k = net(torch.flip(image, dims=[2,]))\n                        vessel += torch.flip(v, dims=[2,])\n                        kidney += torch.flip(k, dims=[2,])\n                        counter += 1\n\n                        v, k = net(torch.flip(image, dims=[3,]))\n                        vessel += torch.flip(v, dims=[3,])\n                        kidney += torch.flip(k, dims=[3,])\n                        counter += 1\n\n                        v, k = net(torch.rot90(image, k=1, dims=[2, 3]))\n                        vessel += torch.rot90(v, k=-1, dims=[2, 3])\n                        kidney += torch.rot90(k, k=-1, dims=[2, 3])\n                        counter += 1\n\n                        v, k = net(torch.rot90(image, k=2, dims=[2, 3]))\n                        vessel += torch.rot90(v, k=-2, dims=[2, 3])\n                        kidney += torch.rot90(k, k=-2, dims=[2, 3])\n                        counter += 1\n\n                        v, k = net(torch.rot90(image, k=3, dims=[2, 3]))\n                        vessel += torch.rot90(v, k=-3, dims=[2, 3])\n                        kidney += torch.rot90(k, k=-3, dims=[2, 3])\n                        counter += 1\n\n                vessel = vessel / counter\n                kidney = kidney / counter\n\n                vessel = vessel.float().data.cpu().numpy()\n                kidney = kidney.float().data.cpu().numpy()\n\n                batch_size = len(vessel)\n                for b in range(batch_size):\n                    mk = kidney[b, 0]\n                    mk = choose_biggest_object(mk, threshold=0.5)\n                    mv = vessel[b, 0]\n                    p = (mv * mk)\n                    if axis == 0:\n                        predict[B + b] += p\n                    if axis == 1:\n                        predict[:, B + b] += p\n                    if axis == 2:\n                        predict[:, :, B + b] += p\n\n                B += batch_size\n\n        predict = predict / len(axes)\n        predict = (predict > cfg['p_threshold']).astype(np.uint8)\n\n        if cfg['cc_threshold'] > 0:\n            predict = cc3d.dust(\n                predict,\n                connectivity=26,\n                threshold=cfg['cc_threshold'],\n                in_place=False\n            )\n\n        predictions.append(predict)\n    print(predictions)\n    return predictions\n\n# Define cfg with necessary parameters\ncfg = {\n    'batch_size': 16,\n    'p_threshold': 0.5,\n    'cc_threshold': 0.5\n}\n\n# Get all image file paths from the test set\nimage_dir = '/kaggle/input/blood-vessel-segmentation/train/kidney_1_voi/images'\nimage_paths = glob.glob(os.path.join(image_dir, '*.tif'))\nimage_paths = image_paths[:5]\n\n# Create valid_meta with file paths and image shapes\ntest_images = []\nfor img_path in image_paths:\n    img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n    volume_shape = (len(image_paths), img.shape[0], img.shape[1])  # Assuming the 3D volume is a stack of 2D images\n    test_images.append({\n        'file': [img_path],  # Modify this if you need to include more files per volume\n        'shape': volume_shape\n    })\n\n# Make predictions\npredictions = make_predictions()\n\nground_truth_masks = []\n\nmask_dir = '/kaggle/input/blood-vessel-segmentation/train/kidney_1_voi/labels'\nmask_paths = glob.glob(os.path.join(mask_dir, '*.tif'))\nmask_paths = mask_paths[:5]\n\nfor mask_path in mask_paths:\n    mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    ground_truth_masks.append(mask)\n\n# Evaluate the predictions\nmetrics = evaluate(predictions, ground_truth_masks)\nprint(\"Evaluation Metrics:\")\nfor metric, score in metrics.items():\n    print(f\"{metric}: {score}\")\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(ground_truth_masks)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"class AugmentedDataset(Dataset):\n    def __init__(self, original_dataset):\n        self.original_dataset = original_dataset\n\n    def __len__(self):\n        return len(self.original_dataset)\n\n    def __getitem__(self, idx):\n        images, masks = self.original_dataset[idx]\n        augmented_images, augmented_masks = self.apply_augmentation(images, masks)\n        return augmented_images, augmented_masks\n\n    @staticmethod\n    def apply_augmentation(images, masks):\n        k = random.randint(0, 3)\n        augmented_images = torch.rot90(images, k=k, dims=(1, 2))\n        augmented_masks = torch.rot90(masks, k=k, dims=(1, 2))\n        if random.random() > 0.5:\n            augmented_images = torch.flip(augmented_images, dims=(1,))\n            augmented_masks = torch.flip(augmented_masks, dims=(1,))\n        if random.random() > 0.5:\n            augmented_images = torch.flip(augmented_images, dims=(2,))\n            augmented_masks = torch.flip(augmented_masks, dims=(2,))\n        return augmented_images, augmented_masks\n\n            \n            \n            \naugmented_train_dataset = MyAugDataset(train_files)  \naugmented_train_loader = DataLoader(augmented_train_dataset, batch_size=16, shuffle=True)\n\n\ncriterion = nn.BCEWithLogitsLoss()\ntotal_epochs = 10\ncheckpoint_dir = './checkpoints'\nstart_epoch = 0\ncheckpoint_path = os.path.join(checkpoint_dir, 'latest_checkpoint.pth')\nif os.path.isfile(checkpoint_path):\n    start_epoch, net, optimizer, _ = load_checkpoint(checkpoint_path, net, optimizer)\n    print(f\"Resuming training from epoch {start_epoch}\")\n\n# Continue training for additional epochs\ntrain(net, augmented_train_loader, val_loader, criterion, optimizer, start_epoch=start_epoch, total_epochs=total_epochs, checkpoint_dir=checkpoint_dir)\n\n","metadata":{}}]}