{"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":9988,"databundleVersionId":868324,"sourceType":"competition"},{"sourceId":9619,"sourceType":"modelInstanceVersion","modelInstanceId":7816},{"sourceId":9637,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":7832},{"sourceId":9653,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":7848}],"dockerImageVersionId":30648,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\n# from torchvision.transforms import v2\nfrom torchvision.io import read_image\nfrom torch.utils.data import  Dataset\nfrom PIL import Image\nimport torchvision.transforms as transforms\nfrom torch.utils.data import DataLoader\n# from skimage.morphology import label\nimport cv2","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:41:37.958582Z","iopub.execute_input":"2024-02-07T21:41:37.959318Z","iopub.status.idle":"2024-02-07T21:41:37.964665Z","shell.execute_reply.started":"2024-02-07T21:41:37.959283Z","shell.execute_reply":"2024-02-07T21:41:37.963577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = os.listdir('/kaggle/input/airbus-ship-detection/train_v2')\ntest_data = os.listdir('/kaggle/input/airbus-ship-detection/test_v2')\n\n\nimage_path_train = '/kaggle/input/airbus-ship-detection/train_v2'\nimage_path_test = '/kaggle/input/airbus-ship-detection/test_v2'\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-02-07T21:41:38.089827Z","iopub.execute_input":"2024-02-07T21:41:38.090320Z","iopub.status.idle":"2024-02-07T21:41:38.202825Z","shell.execute_reply.started":"2024-02-07T21:41:38.090294Z","shell.execute_reply":"2024-02-07T21:41:38.201914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.read_csv('/kaggle/input/airbus-ship-detection/sample_submission_v2.csv')\nsubmission.head()","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:41:38.269907Z","iopub.execute_input":"2024-02-07T21:41:38.270455Z","iopub.status.idle":"2024-02-07T21:41:38.288830Z","shell.execute_reply.started":"2024-02-07T21:41:38.270427Z","shell.execute_reply":"2024-02-07T21:41:38.288008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask = pd.read_csv('/kaggle/input/airbus-ship-detection/train_ship_segmentations_v2.csv')\ndisplay(mask.head(5))\nprint((mask.shape))","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:41:38.425451Z","iopub.execute_input":"2024-02-07T21:41:38.425715Z","iopub.status.idle":"2024-02-07T21:41:39.005426Z","shell.execute_reply.started":"2024-02-07T21:41:38.425693Z","shell.execute_reply":"2024-02-07T21:41:39.004475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"empty_ship = mask['EncodedPixels'].isna().sum()\nprint('Empty image without ships: ',empty_ship)\n\nContaining_ship = mask['EncodedPixels'].notna().sum()\nprint('Image containing ships: ',Containing_ship)","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:41:39.006929Z","iopub.execute_input":"2024-02-07T21:41:39.007220Z","iopub.status.idle":"2024-02-07T21:41:39.041874Z","shell.execute_reply.started":"2024-02-07T21:41:39.007194Z","shell.execute_reply":"2024-02-07T21:41:39.040856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = mask.groupby('ImageId').agg({'EncodedPixels': 'count'})\ndf = df.rename(columns={'EncodedPixels' : 'ships'})\ndf['has_ship'] = df['ships'].map(lambda x: 1 if x > 0 else 0)\ndf","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:41:39.043144Z","iopub.execute_input":"2024-02-07T21:41:39.043504Z","iopub.status.idle":"2024-02-07T21:41:39.407059Z","shell.execute_reply.started":"2024-02-07T21:41:39.043471Z","shell.execute_reply":"2024-02-07T21:41:39.406084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask['ships'] = mask['EncodedPixels'].map(lambda x: 1 if not pd.isna(x) else 0)\nmask","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:41:39.409547Z","iopub.execute_input":"2024-02-07T21:41:39.409913Z","iopub.status.idle":"2024-02-07T21:41:39.695825Z","shell.execute_reply.started":"2024-02-07T21:41:39.409881Z","shell.execute_reply":"2024-02-07T21:41:39.694956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"filtered_image_ids = set(mask['ImageId'].unique()) \n\ntrain_data_2 = [x for x in train_data if x in filtered_image_ids]\nset(train_data_2)\nlen(train_data_2)","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:41:39.696830Z","iopub.execute_input":"2024-02-07T21:41:39.697078Z","iopub.status.idle":"2024-02-07T21:41:39.819948Z","shell.execute_reply.started":"2024-02-07T21:41:39.697055Z","shell.execute_reply":"2024-02-07T21:41:39.819022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision import transforms\nfrom PIL import Image\nimport numpy as np\n\nclass CustomDataset(Dataset):\n    def __init__(self, df, image_dir, transform=None):\n        self.df = df\n        self.image_dir = image_dir\n        self.transform = transform\n        \n        self.image_names = df.loc[df['ships'] == 1, 'ImageId'].tolist()\n    \n    def __len__(self):\n        return len(self.image_names)\n    \n    def __getitem__(self, idx):\n        image_name = self.image_names[idx]\n        image_path = os.path.join(self.image_dir, image_name)\n        image = Image.open(image_path).convert('RGB')\n        \n        # Получаем маску из датасета\n#         mask = self.masks_as_image(self.df.query('ImageId == \"image_name\"')['EncodedPixels'])\n        mask = self.masks_as_image(self.df.query(f'ImageId == \"{image_name}\"')['EncodedPixels'])\n\n        \n        # Применяем трансформации\n        if self.transform is not None:\n            image = self.transform(image)\n            mask = self.transform(mask)\n            mask = (mask > 0).float()  # Бинаризация маски\n            \n#             image = image/255.0\n                 \n        return image, mask\n    \n    def rle_decode(self, mask_rle, shape=(768, 768)):\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        return img.reshape(shape).T \n    \n    def masks_as_image(self, mask_list):\n        # Take the individual ship masks and create a single mask array for all ships\n        mask = np.zeros((768, 768), dtype=np.uint8)\n\n        # Check if mask_list is not a float (i.e., it's not a single value)\n        if not isinstance(mask_list, float):\n            if isinstance(mask_list, str):\n                mask |= self.rle_decode(mask_list)\n            elif pd.notna(mask_list).any():\n                for massk in mask_list:\n                    if isinstance(massk, str):\n                        mask |= self.rle_decode(massk)\n\n        return Image.fromarray(mask)  # Convert NumPy array to PIL Image object\n    \n# Define transformations for images\ntransform = transforms.Compose([\n    transforms.Resize((256, 256)),\n    transforms.ToTensor(),\n#     transforms.Normalize(mean=[0.485, 0.456, 0.406],std=[0.229, 0.224, 0.225])\n])\n\nmask_transform = transforms.Compose([\n    transforms.Resize((256, 256)),\n    transforms.ToTensor()\n])\n\nbatch_size = 28\n\n# Create dataset instance\ndataset = CustomDataset(df=mask, image_dir=image_path_train, transform=transform)\n\nlen(dataset)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:41:39.821195Z","iopub.execute_input":"2024-02-07T21:41:39.821504Z","iopub.status.idle":"2024-02-07T21:41:39.863818Z","shell.execute_reply.started":"2024-02-07T21:41:39.821478Z","shell.execute_reply":"2024-02-07T21:41:39.862965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision import transforms\nfrom PIL import Image\nimport numpy as np\n\nclass CustomDataset_for_test_data(Dataset):\n    def __init__(self, image_dir, transform=None):\n        self.image_dir = image_dir\n        self.transform = transform\n        \n        self.image_names = os.listdir(self.image_dir)\n    \n    def __len__(self):\n        return len(self.image_names)\n    \n    def __getitem__(self, idx):\n        image_name = self.image_names[idx]\n        image_path = os.path.join(self.image_dir, image_name)\n        image = Image.open(image_path).convert('RGB')\n\n        if self.transform is not None:\n            image = self.transform(image)\n                 \n        return image\n    \ntransform = transforms.Compose([\n    transforms.Resize((256, 256)),\n    transforms.ToTensor(),\n])\n\nbatch_size = 28\n\ndataset_test = CustomDataset_for_test_data(image_dir=image_path_test, transform = transform)\nlen(dataset_test)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:41:39.865972Z","iopub.execute_input":"2024-02-07T21:41:39.866673Z","iopub.status.idle":"2024-02-07T21:41:39.883315Z","shell.execute_reply.started":"2024-02-07T21:41:39.866636Z","shell.execute_reply":"2024-02-07T21:41:39.882451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imagee, maskk = dataset[150]\n\nimport matplotlib.pyplot as plt\nplt.imshow(imagee.permute(1, 2, 0))  # Изменение порядка размерностей для отображения в matplotlib\nplt.title('Image')\nplt.axis('off')\nplt.show()\n\n# Отображение маски\nplt.imshow(maskk[0], cmap='gray')  # Предполагаем, что маска имеет размерность (1, H, W)\nplt.title('Mask')\nplt.axis('off')\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:41:40.457701Z","iopub.execute_input":"2024-02-07T21:41:40.458291Z","iopub.status.idle":"2024-02-07T21:41:40.754421Z","shell.execute_reply.started":"2024-02-07T21:41:40.458263Z","shell.execute_reply":"2024-02-07T21:41:40.753554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"maskk.unique()","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:41:47.093540Z","iopub.execute_input":"2024-02-07T21:41:47.094389Z","iopub.status.idle":"2024-02-07T21:41:47.101541Z","shell.execute_reply.started":"2024-02-07T21:41:47.094351Z","shell.execute_reply":"2024-02-07T21:41:47.100581Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"size = maskk.size()\nsize\nimage_size = imagee.size()\nimage_size","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:41:48.100330Z","iopub.execute_input":"2024-02-07T21:41:48.101110Z","iopub.status.idle":"2024-02-07T21:41:48.107175Z","shell.execute_reply.started":"2024-02-07T21:41:48.101077Z","shell.execute_reply":"2024-02-07T21:41:48.106265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import random_split\n\ntrain_size = int(0.8 * len(dataset)) \ntest_size = len(dataset) - train_size\n\n\ntrain_dataset, test_dataset = random_split(dataset, [train_size, test_size])\n\ntrain_data_loader = DataLoader(train_dataset, batch_size = batch_size, shuffle = True, num_workers= 4)\ntest_data_loader = DataLoader(test_dataset, batch_size = batch_size, shuffle = True, num_workers= 4)\n\nprint(len(train_dataset))\nprint(len(test_dataset))","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:41:50.451715Z","iopub.execute_input":"2024-02-07T21:41:50.452731Z","iopub.status.idle":"2024-02-07T21:41:50.468972Z","shell.execute_reply.started":"2024-02-07T21:41:50.452681Z","shell.execute_reply":"2024-02-07T21:41:50.467972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.transforms.functional as TF\n\nclass DoubleConv(nn.Module):\n    def __init__(self, input_channels, output_channels):\n        super(DoubleConv, self).__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(input_channels,\n                      output_channels,\n                      kernel_size=3,\n                      stride=1,\n                      padding=1,\n                      bias=False),\n            nn.BatchNorm2d(output_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(output_channels, \n                      output_channels,\n                      kernel_size=3,\n                      stride=1,\n                      padding=1,\n                      bias=False),\n            nn.BatchNorm2d(output_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.conv(x)\n\nclass Unet(nn.Module):\n    def __init__(self, input_channels=3, output_channels=1, features=[64, 128, 256, 512]):\n        super(Unet, self).__init__()\n        self.ups = nn.ModuleList()\n        self.downs = nn.ModuleList()\n        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)  \n\n        # Down part of Unet\n        for feature in features:\n            self.downs.append(DoubleConv(input_channels, feature))  \n            input_channels = feature\n\n        # Up part of Unet\n        for feature in reversed(features):\n            self.ups.append(\n                nn.ConvTranspose2d(feature * 2,\n                                   feature,\n                                   kernel_size=2,\n                                   stride=2))\n            self.ups.append(DoubleConv(feature * 2, feature))\n\n        self.bottleneck = DoubleConv(features[-1], features[-1] * 2)\n        self.final_conv = nn.Conv2d(features[0], output_channels, kernel_size=1) \n\n    def forward(self, x):\n        skip_connections = []\n        for down in self.downs:\n            x = down(x)\n            skip_connections.append(x)\n            x = self.pool(x)\n        x = self.bottleneck(x)\n\n        skip_connections = skip_connections[::-1]\n        for idx in range(0, len(self.ups), 2):\n            x = self.ups[idx](x)\n            skip_connection = skip_connections[idx // 2]\n            if skip_connection.shape != x.shape:\n                x = TF.resize(x, size=skip_connection.shape[2:])  \n            concat_skip_connection = torch.cat((skip_connection, x), dim=1)\n            x = self.ups[idx + 1](concat_skip_connection)\n\n#         return torch.sigmoid(self.final_conv(x))\n        return self.final_conv(x)\n\n\ndef test():\n    x = torch.rand(size=(32, 1, 256,\n                         256)) \n    model = Unet(input_channels = 1, output_channels=1)\n    prediction = model(x)\n    print(x.shape)\n    print(prediction.shape)\n    assert prediction.shape == x.shape\n\ntest()\n","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:41:51.433717Z","iopub.execute_input":"2024-02-07T21:41:51.434044Z","iopub.status.idle":"2024-02-07T21:42:16.345273Z","shell.execute_reply.started":"2024-02-07T21:41:51.434017Z","shell.execute_reply":"2024-02-07T21:42:16.344447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_checkpoint(state, filename=\"my_checkpoint.pth.tar\"):\n    print(\"Saving checkpoint\")\n    torch.save(state, filename)\n\ndef load_checkpoint(checkpoint, model):\n    print('Loading checkpoint')\n    model.load_state_dict(checkpoint['state_dict'])","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:42:16.346859Z","iopub.execute_input":"2024-02-07T21:42:16.347157Z","iopub.status.idle":"2024-02-07T21:42:16.352450Z","shell.execute_reply.started":"2024-02-07T21:42:16.347132Z","shell.execute_reply":"2024-02-07T21:42:16.351508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def check_accuracy(loader, model, device = 'cuda'):\n    num_correct = 0\n    num_pixels = 0\n    dice_score = 0\n\n    model.eval()\n\n    with torch.no_grad():\n        for x, y in loader:\n            x = x.to(device)\n            y = y.to(device)\n\n            preds = torch.sigmoid(model(x))\n            preds = (preds > 0.5).float()\n            num_correct += (preds == y).sum()\n            num_pixels += torch.numel(preds)\n            \n            dice_score += (2 * (preds * y).sum()) / ((preds + y).sum() + 1e-8)\n\n    print(f'Got {num_correct}/{num_pixels} with accuracy {num_correct/num_pixels*100:.2f}')\n    print(f'Dice score: {dice_score/len(loader)}')\n    model.train()","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:42:16.353760Z","iopub.execute_input":"2024-02-07T21:42:16.354197Z","iopub.status.idle":"2024-02-07T21:42:16.366831Z","shell.execute_reply.started":"2024-02-07T21:42:16.354164Z","shell.execute_reply":"2024-02-07T21:42:16.366023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def check_accuracy_1(loader, model, device='cuda'):\n    num_correct = 0\n    num_pixels = 0\n    dice_score = 0\n    predicted_masks = []\n\n    model.eval()\n\n    with torch.no_grad():\n        for x, y in loader:\n            x = x.to(device)\n            y = y.to(device)\n\n            preds = torch.sigmoid(model(x))\n            preds = (preds > 0.5).float()\n            num_correct += (preds == y).sum()\n            num_pixels += torch.numel(preds)\n            \n            dice_score += (2 * (preds * y).sum()) / ((preds + y).sum() + 1e-8)\n\n            # Сохраняем предсказанные маски\n            predicted_masks.append(preds.cpu().numpy())\n\n    print(f'Got {num_correct}/{num_pixels} with accuracy {num_correct/num_pixels*100:.2f}')\n    print(f'Dice score: {dice_score/len(loader)}')\n    model.train()\n\n    return predicted_masks","metadata":{"execution":{"iopub.status.busy":"2024-02-07T20:53:43.046262Z","iopub.execute_input":"2024-02-07T20:53:43.046685Z","iopub.status.idle":"2024-02-07T20:53:43.055407Z","shell.execute_reply.started":"2024-02-07T20:53:43.046648Z","shell.execute_reply":"2024-02-07T20:53:43.054431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\nmodel = Unet(input_channels=3, output_channels=1).to(device) ","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:42:16.368396Z","iopub.execute_input":"2024-02-07T21:42:16.368686Z","iopub.status.idle":"2024-02-07T21:42:16.654100Z","shell.execute_reply.started":"2024-02-07T21:42:16.368663Z","shell.execute_reply":"2024-02-07T21:42:16.653228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(torch.load('/kaggle/input/123321/pytorch/123321/1/actually5_17 (1).pth', map_location=torch.device('cpu')))","metadata":{"execution":{"iopub.status.busy":"2024-02-07T18:31:37.325387Z","iopub.execute_input":"2024-02-07T18:31:37.326121Z","iopub.status.idle":"2024-02-07T18:31:37.447026Z","shell.execute_reply.started":"2024-02-07T18:31:37.326085Z","shell.execute_reply":"2024-02-07T18:31:37.445962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom tqdm import tqdm\nimport torch.nn as nn \nimport torch.optim as optim \n\nlearning_rate = 1e-6\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\nnum_epochs = 5\nnum_workers = 10\npin_memory = True\n\ndef train_fn(loader, model, optimizer, loss_fn):\n    loop = tqdm(loader)\n    \n    for batch_idx, (data, target) in enumerate(loop):\n        data = data.to(device)\n        target = target.float().to(device)\n        \n        # Forward\n        predictions = model(data)\n        loss = loss_fn(predictions, target)\n            \n        # Backward\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        # Update tqdm loop\n        loop.set_postfix(loss=loss.item())\n\n# loss_fn = nn.BCELoss()\nloss_fn = nn.BCEWithLogitsLoss()\n# loss_fn = criterion\noptimizer = optim.Adam(model.parameters(), lr=learning_rate)        \n\nfor epoch in range(num_epochs):\n    train_fn(train_data_loader, model, optimizer, loss_fn)\n    \n    #save_model\n    checkpoint = {\n        'state_dict': model.state_dict(),\n        'optimizer': optimizer.state_dict()\n    }\n    save_checkpoint(checkpoint)\n    torch.save(model.state_dict(), 'actually5_18.pth')\n    check_accuracy(test_data_loader, model, device = device)","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:45:16.513706Z","iopub.execute_input":"2024-02-07T21:45:16.514076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from IPython.display import FileLink\n\n# # Создайте объект FileLink для скачивания файла\n# FileLink(r'./actually5_17.pth')","metadata":{"execution":{"iopub.status.busy":"2024-02-07T19:34:23.263129Z","iopub.execute_input":"2024-02-07T19:34:23.263788Z","iopub.status.idle":"2024-02-07T19:34:23.270979Z","shell.execute_reply.started":"2024-02-07T19:34:23.263747Z","shell.execute_reply":"2024-02-07T19:34:23.270000Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predicted_masks = check_accuracy_1(test_data_loader, model, device=device)","metadata":{"execution":{"iopub.status.busy":"2024-02-07T20:56:36.347460Z","iopub.execute_input":"2024-02-07T20:56:36.348311Z","iopub.status.idle":"2024-02-07T20:59:52.377483Z","shell.execute_reply.started":"2024-02-07T20:56:36.348275Z","shell.execute_reply":"2024-02-07T20:59:52.376027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_segmentation(image, mask):\n    plt.figure(figsize=(10, 5))\n    \n    # Преобразование изображения к правильному формату\n    image = np.moveaxis(image, 0, -1)\n    \n    # Изображение\n    plt.subplot(1, 2, 1)\n    plt.imshow(image)\n    plt.title('Изображение')\n    plt.axis('off')\n    \n    # Маска\n    plt.subplot(1, 2, 2)\n    plt.imshow(mask.squeeze(), cmap='gray')\n    plt.title('Маска сегментации')\n    plt.axis('off')\n    \n    plt.show()\n\n# Проходим по первым 50 элементам в тестовом DataLoader\nfor idx, (image, mask) in enumerate(test_data_loader):\n    if idx >= 50:\n        break  # Прерываем цикл после 50 элементов\n    \n    # Берем только одно изображение из батча\n    image = image[0]\n    \n    # Берем соответствующую маску из предсказанных масок\n    mask = predicted_masks[idx][0]\n    \n    visualize_segmentation(image.numpy(), mask)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:15:57.798874Z","iopub.execute_input":"2024-02-07T21:15:57.799249Z","iopub.status.idle":"2024-02-07T21:16:19.895206Z","shell.execute_reply.started":"2024-02-07T21:15:57.799212Z","shell.execute_reply":"2024-02-07T21:16:19.894219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_segmentation(image, mask, test_mask):\n    plt.figure(figsize=(15, 5))\n    \n    # Преобразование изображения к правильному формату\n    image = np.moveaxis(image, 0, -1)\n    \n    # Изображение\n    plt.subplot(1, 3, 1)\n    plt.imshow(image)\n    plt.title('Изображение')\n    plt.axis('off')\n    \n    # Предсказанная маска\n    plt.subplot(1, 3, 2)\n    plt.imshow(mask.squeeze(), cmap='gray')\n    plt.title('Предсказанная маска')\n    plt.axis('off')\n    \n    # Третье изображение маски\n    plt.subplot(1, 3, 3)\n    plt.imshow(test_mask.squeeze(), cmap='gray')\n    plt.title('Настоящая маска')\n    plt.axis('off')\n    \n    plt.show()\n\n# Проходим по первым 50 элементам в тестовом DataLoader\nfor idx, (image, mask) in enumerate(test_data_loader):\n    if idx >= 50:\n        break  # Прерываем цикл после 50 элементов\n    \n    # Берем только одно изображение из батча\n    image = image[0]\n    \n    # Берем соответствующую маску из предсказанных масок\n    mask_pred = predicted_masks[idx][0]\n    mask_true = mask[0]\n    \n    visualize_segmentation(image.numpy(), mask_pred, mask_true)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:22:40.500291Z","iopub.execute_input":"2024-02-07T21:22:40.500695Z","iopub.status.idle":"2024-02-07T21:23:06.577435Z","shell.execute_reply.started":"2024-02-07T21:22:40.500660Z","shell.execute_reply":"2024-02-07T21:23:06.576047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# image.squeeze(0).shape\npredicted_masks[idx].shape","metadata":{"execution":{"iopub.status.busy":"2024-02-07T21:12:40.837824Z","iopub.execute_input":"2024-02-07T21:12:40.838192Z","iopub.status.idle":"2024-02-07T21:12:40.844441Z","shell.execute_reply.started":"2024-02-07T21:12:40.838160Z","shell.execute_reply":"2024-02-07T21:12:40.843562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torchvision.transforms as transforms\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport numpy as np\n\n# Функция для загрузки и предобработки изображения\ndef preprocess_image(image_path):\n    transform = transforms.Compose([\n        transforms.Resize((256, 256)),\n        transforms.ToTensor(),\n#         transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    ])\n    image = Image.open(image_path).convert('RGB')\n    image = transform(image)\n    return image\n\n# Функция для визуализации изображения и его маски\ndef visualize_segmentation(image, mask):\n    plt.figure(figsize=(10, 5))\n    \n    # Изображение\n    plt.subplot(1, 2, 1)\n    plt.imshow(image.permute(1, 2, 0))\n    plt.title('Image')\n    plt.axis('off')\n    \n    # Маска\n    plt.subplot(1, 2, 2)\n    plt.imshow(mask.squeeze(), cmap='gray')\n    plt.title('Segmentation Mask')\n    plt.axis('off')\n    \n    plt.show()\n\n# Путь к вашей модели\nmodel_path = \"/kaggle/working/actually5_17.pth\"\n\n# Путь к тестовому изображению\ntest_image_path = \"/kaggle/input/airbus-ship-detection/train_v2/000fd9827.jpg\"\n\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\nmodel = Unet(input_channels=3, output_channels=1).to(device) \nmodel.load_state_dict(torch.load('/kaggle/input/123321/pytorch/123321/1/actually5_17 (1).pth', map_location=torch.device('cpu')))\nmodel.eval()\n\n\n\n# Предобработка и вывод тестового изображения\ntest_image = preprocess_image(test_image_path).to(device)\nwith torch.no_grad():\n    output = model(test_image.unsqueeze(0))  # Добавляем размерность пакета\npredicted_mask = torch.argmax(output, dim=1).cpu()  # Возвращаем маску на CPU для визуализации\n\n# Визуализация изображения и маски\n# Визуализация изображения и маски\nvisualize_segmentation(test_image.cpu(), predicted_mask.cpu())\n","metadata":{"execution":{"iopub.status.busy":"2024-02-07T20:51:47.314502Z","iopub.execute_input":"2024-02-07T20:51:47.315410Z","iopub.status.idle":"2024-02-07T20:51:47.980010Z","shell.execute_reply.started":"2024-02-07T20:51:47.315371Z","shell.execute_reply":"2024-02-07T20:51:47.979113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_image.unsqueeze(0).shape","metadata":{"execution":{"iopub.status.busy":"2024-02-07T20:49:22.106679Z","iopub.execute_input":"2024-02-07T20:49:22.107626Z","iopub.status.idle":"2024-02-07T20:49:22.113766Z","shell.execute_reply.started":"2024-02-07T20:49:22.107587Z","shell.execute_reply":"2024-02-07T20:49:22.112822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}