{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"},{"sourceId":173621,"sourceType":"modelInstanceVersion","modelInstanceId":145101,"modelId":167660},{"sourceId":174462,"sourceType":"modelInstanceVersion","modelInstanceId":116025,"modelId":139257}],"dockerImageVersionId":30787,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from tensorflow.keras.models import load_model\nfrom torchvision import models, transforms\nfrom torch.utils.data import DataLoader\nfrom torchvision.transforms import v2\nimport tensorflow as tf\nfrom tqdm import tqdm\nfrom PIL import Image\nimport pandas as pd\nimport numpy as np\nimport torch\nimport os","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-11-21T20:50:34.684613Z","iopub.execute_input":"2024-11-21T20:50:34.685564Z","iopub.status.idle":"2024-11-21T20:50:34.690135Z","shell.execute_reply.started":"2024-11-21T20:50:34.685523Z","shell.execute_reply":"2024-11-21T20:50:34.689288Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_data_directory = '/kaggle/input/cassava-leaf-disease-classification/test_images'\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nnum_classes = 5\n\nen_model_path = '/kaggle/input/efficientnetv2-large-test/pytorch/default/4/efficientnet_v2_l_480_8591_ISP_CBP.pth'\nen_image_size = 480\n\nvit_model_path = '/kaggle/input/vit_l_cassava/pytorch/default/6/vit_h_14_518_8867_CBP.pth'\nvit_image_size = 518\n\nmodel_select = \"vit\"\n\nif model_select == \"vit\":\n    model_image_size = vit_image_size\nif model_select == \"en\":\n    model_image_size = en_image_size","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T20:50:34.691750Z","iopub.execute_input":"2024-11-21T20:50:34.692055Z","iopub.status.idle":"2024-11-21T20:50:34.701035Z","shell.execute_reply.started":"2024-11-21T20:50:34.692022Z","shell.execute_reply":"2024-11-21T20:50:34.700183Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def invert_square_pad(img):\n    width, height = img.size\n\n    center_width, center_height = width // 2, height // 2\n    top_left = img.crop((0, 0, center_width, center_height))\n    top_right = img.crop((center_width, 0, width, center_height))\n    bottom_left = img.crop((0, center_height, center_width, height))\n    bottom_right = img.crop((center_width, center_height, width, height))\n\n    top_combined = Image.new('RGB', (width, center_height))\n    top_combined.paste(bottom_right, (0, 0))\n    top_combined.paste(bottom_left, (center_width, 0))\n\n    bottom_combined = Image.new('RGB', (width, center_height))\n    bottom_combined.paste(top_right, (0, 0))\n    bottom_combined.paste(top_left, (center_width, 0))\n\n    flipped_img = Image.new('RGB', (width, height))\n    flipped_img.paste(top_combined, (0, 0))\n    flipped_img.paste(bottom_combined, (0, center_height))\n\n    img = flipped_img.copy()\n    del top_combined, bottom_combined, flipped_img\n\n    max_side = max(width, height)\n    padding = (\n        (max_side - width) // 2,  # left\n        (max_side - height) // 2,  # top\n        (max_side - width) - (max_side - width) // 2,  # right\n        (max_side - height) - (max_side - height) // 2  # bottom\n    )\n\n    padded_img = transforms.functional.pad(img, padding, padding_mode='reflect')\n\n    return padded_img\n\ndef resize_max_side(img, size):\n    height, width = img.size\n    if max(width, height) > size:\n        return img\n\n    if width > height:\n        new_width = size\n        new_height = int(size * height / width)\n    else:\n        new_height = size\n        new_width = int(size * width / height)\n        \n    return transforms.functional.resize(img, (new_width, new_height))\n\ndef pad_to_square(img):\n    width, height = img.size\n    max_side = max(width, height)\n    padding = (\n        (max_side - width) // 2,\n        (max_side - height) // 2,\n        (max_side - width) - (max_side - width) // 2,\n        (max_side - height) - (max_side - height) // 2\n    )  # left, top, right, bottom\n\n    padded_img = transforms.functional.pad(img, padding, padding_mode='reflect')\n\n    return padded_img","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T20:50:34.702267Z","iopub.execute_input":"2024-11-21T20:50:34.702589Z","iopub.status.idle":"2024-11-21T20:50:34.715856Z","shell.execute_reply.started":"2024-11-21T20:50:34.702561Z","shell.execute_reply":"2024-11-21T20:50:34.715065Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_transforms = transforms.Compose([\n    v2.Lambda(lambda img: resize_max_side(img, 800)),\n    v2.Lambda(pad_to_square),\n    # v2.Lambda(invert_square_pad),\n    v2.ToImage(),\n    v2.ToDtype(torch.float32, scale=True),\n    v2.Resize((model_image_size, model_image_size)),\n    v2.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5])\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T20:50:34.717001Z","iopub.execute_input":"2024-11-21T20:50:34.717416Z","iopub.status.idle":"2024-11-21T20:50:34.731616Z","shell.execute_reply.started":"2024-11-21T20:50:34.717364Z","shell.execute_reply":"2024-11-21T20:50:34.730866Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if model_select == \"vit\":\n    vit_model = models.vit_h_14(weights=None, image_size=518)\n    vit_model.heads.head = torch.nn.Linear(vit_model.heads.head.in_features, num_classes)\n    vit_model.load_state_dict(torch.load(vit_model_path, map_location=device, weights_only=True))\n    vit_model.to(device)\n    vit_model.eval()\nif model_select == \"en\":\n    en_model = models.efficientnet_v2_l(weights=None)\n    en_model.classifier[1] = torch.nn.Linear(en_model.classifier[1].in_features, num_classes)\n    en_model.load_state_dict(torch.load(en_model_path, map_location=device, weights_only=True))\n    en_model.to(device)\n    en_model.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T20:50:34.733245Z","iopub.execute_input":"2024-11-21T20:50:34.733488Z","iopub.status.idle":"2024-11-21T20:50:57.798808Z","shell.execute_reply.started":"2024-11-21T20:50:34.733464Z","shell.execute_reply":"2024-11-21T20:50:57.798109Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions = []\nimage_ids = []\n\nfor image_name in tqdm(os.listdir(test_data_directory), desc='Test'):\n    image_path = os.path.join(test_data_directory, image_name)\n\n    image = Image.open(image_path).convert('RGB')\n    transformed_image = val_transforms(image).unsqueeze(0)\n    transformed_image = transformed_image.to(device)\n\n    with torch.no_grad():\n        if model_select == \"vit\":\n            vit_output = vit_model(transformed_image)\n            _, predicted_class = torch.max(vit_output, 1)\n        if model_select == \"en\":\n            en_output = en_model(transformed_image)\n            _, predicted_class = torch.max(en_output, 1)\n\n    predictions.append(predicted_class.item())\n    image_ids.append(image_name)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T20:50:57.799853Z","iopub.execute_input":"2024-11-21T20:50:57.800169Z","iopub.status.idle":"2024-11-21T20:50:58.752553Z","shell.execute_reply.started":"2024-11-21T20:50:57.800141Z","shell.execute_reply":"2024-11-21T20:50:58.751723Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission_df = pd.DataFrame({\n    'image_id': image_ids,\n    'label': predictions\n})\nsubmission_df.to_csv('submission.csv', index=False)\nprint(\"Submission file created: submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T20:50:58.753735Z","iopub.execute_input":"2024-11-21T20:50:58.754116Z","iopub.status.idle":"2024-11-21T20:50:58.765522Z","shell.execute_reply.started":"2024-11-21T20:50:58.754074Z","shell.execute_reply":"2024-11-21T20:50:58.764757Z"}},"outputs":[],"execution_count":null}]}