{"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":"code","source":"# Install required packages\n!pip install segmentation_models_pytorch\n!pip install -q cairosvg==2.5.2\n!pip install -q reportlab==3.5.65\n!pip install -q cssutils==2.2.0\n!pip install segmentation_models_pytorch\n!pip install -U segmentation-models-pytorch albumentations --user \n\n\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport glob\nimport os\nfrom PIL import Image\nfrom skimage.io import imread\nfrom tqdm import tqdm_notebook as tqdm\nimport matplotlib.pyplot as plt\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport segmentation_models_pytorch as smp\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import confusion_matrix\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-05-28T16:01:11.951071Z","iopub.execute_input":"2023-05-28T16:01:11.951692Z","iopub.status.idle":"2023-05-28T16:03:29.782254Z","shell.execute_reply.started":"2023-05-28T16:01:11.951659Z","shell.execute_reply":"2023-05-28T16:03:29.781101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Set device\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n\n# Validate that there are images in the extracted directory\nIMAGE_PATH = '../input/airbus-ship-detection/train_v2'\nimages = os.listdir(IMAGE_PATH)\nprint(len(images))\n\n# Read masks\nmasks = pd.read_csv('../input/airbus-ship-detection/train_ship_segmentations_v2.csv')\nmasks.head()\n\n# Define image size\nIMG_SIZE = 96\n\n# Define image transformation\nmy_transform = transforms.Compose([\n    transforms.Resize(size=(IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor()\n])","metadata":{"execution":{"iopub.status.busy":"2023-05-28T16:04:26.038972Z","iopub.execute_input":"2023-05-28T16:04:26.039426Z","iopub.status.idle":"2023-05-28T16:04:28.833896Z","shell.execute_reply.started":"2023-05-28T16:04:26.039388Z","shell.execute_reply":"2023-05-28T16:04:28.832755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define RLE decode function\ndef rle_decode(mask_rle, shape=(768, 768)):\n    try:\n        s = mask_rle.split()\n        starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n        starts -= 1\n        ends = starts + lengths\n        img = np.zeros(shape[0] * shape[1], dtype=np.uint8)\n        for lo, hi in zip(starts, ends):\n            img[lo:hi] = 1\n    except:\n        img = np.zeros(shape[0] * shape[1], dtype=np.uint8)  # no mask found / no encoding\n    return img.reshape(shape).T\n","metadata":{"execution":{"iopub.status.busy":"2023-05-28T16:04:44.971987Z","iopub.execute_input":"2023-05-28T16:04:44.972376Z","iopub.status.idle":"2023-05-28T16:04:44.981161Z","shell.execute_reply.started":"2023-05-28T16:04:44.972348Z","shell.execute_reply":"2023-05-28T16:04:44.979862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define Airbus Dataset class\nclass AirbusDataset(Dataset):\n    def __init__(self, images, masks, transforms=None):\n        self.images = images\n        self.masks = masks\n        self.transform = transforms\n        \n    def __getitem__(self, idx):\n        img = Image.open(IMAGE_PATH + \"/\" + self.images[idx])\n        img_masks = self.masks.loc[self.masks['ImageId'] == self.images[idx], 'EncodedPixels'].tolist()\n        \n        if self.transform:\n            img = self.transform(img)\n\n        # Take the individual ship masks and create a single mask array for all ships\n        all_masks = np.zeros((IMG_SIZE, IMG_SIZE))\n        for mask in img_masks:\n            all_masks += transforms.Resize(size=(IMG_SIZE, IMG_SIZE))(Image.fromarray(rle_decode(mask)))\n        \n        # The resize function distorts the mask values to be less than 1, so we switch them back to binary values.\n        all_masks[all_masks > 0] = 1\n        \n        return img, all_masks\n        \n    def __len__(self):\n        return len(self.images)","metadata":{"execution":{"iopub.status.busy":"2023-05-28T16:05:59.902771Z","iopub.execute_input":"2023-05-28T16:05:59.903212Z","iopub.status.idle":"2023-05-28T16:05:59.913435Z","shell.execute_reply.started":"2023-05-28T16:05:59.903179Z","shell.execute_reply":"2023-05-28T16:05:59.912521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define helper function for data visualization\ndef visualize(**images):\n    \"\"\"Plot images in a grid\"\"\"\n    n = len(images)\n    plt.figure(figsize=(15, 5))\n    for i, (name, image) in enumerate(images.items()):\n        plt.subplot(1, n, i+1)\n        plt.xticks([])\n        plt.yticks([])\n        plt.title(name)\n        if image.shape[0] == 3:  # RGB image\n            image = np.transpose(image, (1, 2, 0))  # Convert from (C, H, W) to (H, W, C)\n            plt.imshow(image)\n        else:  # Grayscale image\n            plt.imshow(image, cmap='gray')\n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-05-28T16:06:13.058837Z","iopub.execute_input":"2023-05-28T16:06:13.059273Z","iopub.status.idle":"2023-05-28T16:06:13.068087Z","shell.execute_reply.started":"2023-05-28T16:06:13.059243Z","shell.execute_reply":"2023-05-28T16:06:13.06677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create an instance of the Airbus dataset\ndataset = AirbusDataset(images, masks, transforms=my_transform)\n\n# Visualize a sample image and mask\nsample_idx = 0\nsample_image, sample_mask = dataset[sample_idx]\nvisualize(Image=sample_image, Mask=sample_mask)\n\n# Split the dataset into train and validation sets\ntrain_images, val_images = train_test_split(images, test_size=0.2, random_state=42)\ntrain_dataset = AirbusDataset(train_images, masks, transforms=my_transform)\nval_dataset = AirbusDataset(val_images, masks, transforms=my_transform)\n\n# Create data loaders for training and validation\nbatch_size = 16\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=batch_size)\n\n# Define the model architecture\nmodel = smp.Unet('resnet34', encoder_weights='imagenet', classes=1).to(DEVICE)\n\n# Define the loss function and optimizer\nloss_fn = torch.nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)\n\n# Training loop\nnum_epochs = 10\n\nfor epoch in range(num_epochs):\n    train_loss = 0.0\n    val_loss = 0.0\n    \n    # Training\n    model.train()\n    for images, masks in train_loader:\n        images = images.to(DEVICE)\n        masks = masks.to(DEVICE)\n        \n        optimizer.zero_grad()\n        \n        outputs = model(images)\n        loss = loss_fn(outputs, masks.unsqueeze(1))\n        \n        loss.backward()\n        optimizer.step()\n        \n        train_loss += loss.item() * images.size(0)","metadata":{"execution":{"iopub.status.busy":"2023-05-28T16:06:42.578887Z","iopub.execute_input":"2023-05-28T16:06:42.579938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Validation\n    model.eval()\n    with torch.no_grad():\n        for images, masks in val_loader:\n            images = images.to(DEVICE)\n            masks = masks.to(DEVICE)\n            \n            outputs = model(images)\n            loss = loss_fn(outputs, masks.unsqueeze(1))\n            \n            val_loss += loss.item() * images.size(0)\n    \n    # Calculate average loss\n    train_loss /= len(train_dataset)\n    val_loss /= len(val_dataset)\n    \n    # Print loss for the epoch\n    print(f'Epoch {epoch+1}/{num_epochs}, Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}')","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}