{"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":"# Inference\n\nSimple pipeline to make inference by only taking tiles from the centre of an image, making prediction on each tile and then averaging the probabilities.","metadata":{}},{"cell_type":"code","source":"#imports\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport os\n\nimport joblib\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms as transforms\nfrom torchvision import models\n\nImage.MAX_IMAGE_PIXELS = None\n\nencoder = joblib.load('/kaggle/input/mymodels/UBCO/label_encoder.joblib')\ncheckpoint = torch.load('/kaggle/input/mymodels/UBCO/EfficientNet_V2_small_v2.pth')","metadata":{"execution":{"iopub.status.busy":"2023-10-21T00:11:11.892483Z","iopub.execute_input":"2023-10-21T00:11:11.892771Z","iopub.status.idle":"2023-10-21T00:11:23.73522Z","shell.execute_reply.started":"2023-10-21T00:11:11.892734Z","shell.execute_reply":"2023-10-21T00:11:23.734283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = {\n    'BATCH_SIZE' : 8,\n    'DEVICE':torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n}","metadata":{"execution":{"iopub.status.busy":"2023-10-21T00:11:23.736852Z","iopub.execute_input":"2023-10-21T00:11:23.737402Z","iopub.status.idle":"2023-10-21T00:11:23.742504Z","shell.execute_reply.started":"2023-10-21T00:11:23.737371Z","shell.execute_reply":"2023-10-21T00:11:23.741551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test = pd.read_csv('/kaggle/input/UBC-OCEAN/test.csv')\ntest['path'] = test['image_id'].apply(lambda x: f'/kaggle/input/UBC-OCEAN/test_images/{x}.png')","metadata":{"execution":{"iopub.status.busy":"2023-10-21T00:11:23.743801Z","iopub.execute_input":"2023-10-21T00:11:23.744103Z","iopub.status.idle":"2023-10-21T00:11:23.776514Z","shell.execute_reply.started":"2023-10-21T00:11:23.744078Z","shell.execute_reply":"2023-10-21T00:11:23.775635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = models.efficientnet_v2_s()\nmodel_in_features = model.classifier[-1].in_features\n\nfor param in model.parameters():\n    param.requires_grad = False   \n\nmodel.classifier = nn.Sequential(*list(model.classifier.children())[:-1], nn.Linear(model_in_features, 5))\n\nmodel = model.to(config['DEVICE'])\n\nmodel.load_state_dict(checkpoint['model_state_dict'])\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-10-21T00:11:23.778829Z","iopub.execute_input":"2023-10-21T00:11:23.779163Z","iopub.status.idle":"2023-10-21T00:11:24.301372Z","shell.execute_reply.started":"2023-10-21T00:11:23.779136Z","shell.execute_reply":"2023-10-21T00:11:24.300424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_tile(img, x, y, width, height):\n    return img.crop((x, y, x+width, y+height))","metadata":{"execution":{"iopub.status.busy":"2023-10-21T00:11:24.302633Z","iopub.execute_input":"2023-10-21T00:11:24.303006Z","iopub.status.idle":"2023-10-21T00:11:24.30792Z","shell.execute_reply.started":"2023-10-21T00:11:24.302973Z","shell.execute_reply":"2023-10-21T00:11:24.306992Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Display Sample Images","metadata":{}},{"cell_type":"code","source":"# image = Image.open('/kaggle/input/UBC-OCEAN/train_images/10896.png')\n# tile_size = 256\n\n# # Number of tiles in both dimensions\n# num_rows = 8\n# num_cols = 8\n\n# # Determine the center of the image\n# center_x = image.width // 2\n# center_y = image.height // 2\n\n# # Calculate the starting coordinates to get the centered region\n# start_x = center_x - (num_cols * tile_size) // 2\n# start_y = center_y - (num_rows * tile_size) // 2\n\n# # Get coordinates for the center tiles\n# tiles_20 = [(start_x + i * tile_size, start_y + j * tile_size) for i in range(num_cols) for j in range(num_rows)]\n\n# fig, ax = plt.subplots(num_rows, num_cols, figsize=(10,8))\n\n# for i, (x, y) in enumerate(tiles_20):\n#     tile = get_tile(image, x, y, tile_size, tile_size)\n#     tile_np = np.array(tile)  \n\n#     row_index = i // num_cols\n#     col_index = i % num_cols\n    \n#     ax[row_index, col_index].imshow(tile_np)\n#     ax[row_index, col_index].axis('off')\n\n# plt.suptitle('Center 64 Tiles')\n# plt.tight_layout()\n# plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-10-21T00:11:24.309353Z","iopub.execute_input":"2023-10-21T00:11:24.309714Z","iopub.status.idle":"2023-10-21T00:12:47.991899Z","shell.execute_reply.started":"2023-10-21T00:11:24.309681Z","shell.execute_reply":"2023-10-21T00:12:47.990682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"transform = transforms.Compose([\n        transforms.Resize((256, 256), antialias=True),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n    ])\n\n\ndef preprocess_tiles_batch_torch(tiles_batch):\n    tiles_np = [np.array(tile) for tile in tiles_batch]\n    tiles_tensor = torch.stack([transform(tensor) for tensor in tiles_tensor]).to(config['DEVICE']).half()\n\n    return tiles_tensor\n","metadata":{"execution":{"iopub.status.busy":"2023-10-21T00:12:47.993381Z","iopub.execute_input":"2023-10-21T00:12:47.993736Z","iopub.status.idle":"2023-10-21T00:12:48.001779Z","shell.execute_reply.started":"2023-10-21T00:12:47.993705Z","shell.execute_reply":"2023-10-21T00:12:48.00075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_probability(image_path, num_tiles_side=8, tile_size=256):\n    tiles = []\n\n    with Image.open(image_path) as img:\n        # Determine the center of the image\n        center_x = img.width // 2\n        center_y = img.height // 2\n\n        # Calculate the starting coordinates to get the centered region\n        start_x = center_x - (num_tiles_side * tile_size) // 2\n        start_y = center_y - (num_tiles_side * tile_size) // 2\n\n\n        # Extract tiles from the centered region\n        for x_offset in range(0, num_tiles_side * tile_size, tile_size):\n            for y_offset in range(0, num_tiles_side * tile_size, tile_size):\n                tile = get_tile(img, start_x + x_offset, start_y + y_offset, tile_size, tile_size)\n                tiles.append(tile)\n\n        total_prob = torch.zeros(1,5).to(config['DEVICE'])\n        n = 0\n\n\n        if len(tiles) == config['BATCH_SIZE']:\n            tiles_tensor = preprocess_tiles_batch_torch(tiles)\n            tiles_tensor = tiles_tensor.to(config['DEVICE']).half()\n\n            with torch.no_grad():\n                output_logits = model(tiles_tensor)\n                output_prob = torch.softmax(output_logits, dim=1)\n                total_prob += output_prob.sum(dim=0)\n            n += len(tile_batch)\n\n            del tiles_tensor\n            del tiles \n            tiles = []\n            torch.cuda.empty_cache()\n\n\n\n    mean_prob = total_prob / n\n    return mean_prob","metadata":{"execution":{"iopub.status.busy":"2023-10-21T00:12:48.002866Z","iopub.execute_input":"2023-10-21T00:12:48.003181Z","iopub.status.idle":"2023-10-21T00:12:48.016491Z","shell.execute_reply.started":"2023-10-21T00:12:48.003157Z","shell.execute_reply":"2023-10-21T00:12:48.015414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.half()\nmodel.to(config['DEVICE'])\n\nresults = []\nfor i, row in test.iterrows():\n    img_path = row['path']\n    output = get_probability(img_path)\n    results.append(output)\n    del output","metadata":{"execution":{"iopub.status.busy":"2023-10-21T00:12:48.017743Z","iopub.execute_input":"2023-10-21T00:12:48.018032Z","iopub.status.idle":"2023-10-21T00:13:07.719215Z","shell.execute_reply.started":"2023-10-21T00:12:48.018006Z","shell.execute_reply":"2023-10-21T00:13:07.718419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"results_label = [result.argmax().item() for result in results]\ntest['label'] = results_label\ntest = test[['image_id', 'label']]\ntest['label'] = test['label'].map(lambda x: encoder.inverse_transform([x]).item())\n\ntest.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-10-21T00:13:07.721552Z","iopub.execute_input":"2023-10-21T00:13:07.721837Z","iopub.status.idle":"2023-10-21T00:13:07.752026Z","shell.execute_reply.started":"2023-10-21T00:13:07.721811Z","shell.execute_reply":"2023-10-21T00:13:07.751263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}