{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":14774,"databundleVersionId":875431,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":15907135,"datasetId":10200657,"databundleVersionId":16862695}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport cv2\nimport torch\nimport timm\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport albumentations as A\n\n# =========================\n# CONFIG\n# =========================\n\nIMG_SIZE = 384\nNUM_CLASSES = 5\nMODEL_NAME = \"tf_efficientnetv2_s.in21k_ft_in1k\"\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nCLASS_NAMES = {\n    0: \"No DR\",\n    1: \"Mild DR\",\n    2: \"Moderate DR\",\n    3: \"Severe DR\",\n    4: \"Proliferative DR\"\n}\n\n# =========================\n# MODEL ARCHITECTURE\n# =========================\n\nclass DRModel(nn.Module):\n    def __init__(self, model_name=MODEL_NAME, num_classes=NUM_CLASSES):\n        super().__init__()\n\n        self.backbone = timm.create_model(\n            model_name,\n            pretrained=False,\n            num_classes=0,\n            global_pool=\"avg\"\n        )\n\n        in_features = self.backbone.num_features\n\n        self.head = nn.Sequential(\n            nn.Dropout(0.35),\n            nn.Linear(in_features, 512),\n            nn.SiLU(),\n            nn.Dropout(0.30),\n            nn.Linear(512, num_classes)\n        )\n\n    def forward(self, x):\n        x = self.backbone(x)\n        x = self.head(x)\n        return x\n\n# =========================\n# LOAD MODEL\n# =========================\n\nmodel = DRModel().to(DEVICE)\n\ncheckpoint = torch.load(\n    \"/kaggle/input/datasets/abderrahmanegamga/dataffo/best_pytorch_h100_aptos.pth\",\n    map_location=DEVICE\n)\n\nstate_dict = checkpoint[\"model_state_dict\"]\n\n# Fix if model was saved after torch.compile()\nnew_state_dict = {}\nfor k, v in state_dict.items():\n    new_key = k.replace(\"_orig_mod.\", \"\")\n    new_state_dict[new_key] = v\n\nmodel.load_state_dict(new_state_dict)\nmodel.eval()\n\n# =========================\n# LOAD THRESHOLDS\n# =========================\n\nif os.path.exists(\"/kaggle/input/datasets/abderrahmanegamga/dataffo/best_thresholds.csv\"):\n    thresholds = pd.read_csv(\"/kaggle/input/datasets/abderrahmanegamga/dataffo/best_thresholds.csv\")[\"threshold\"].values\nelse:\n    thresholds = np.array([0.5, 1.5, 2.5, 3.5])\n\nprint(\"Thresholds:\", thresholds)\n\n# =========================\n# PREPROCESS IMAGE\n# =========================\n\nvalid_transforms = A.Compose([\n    A.Resize(IMG_SIZE, IMG_SIZE),\n    A.Normalize(\n        mean=(0.485, 0.456, 0.406),\n        std=(0.229, 0.224, 0.225)\n    )\n])\n\ndef crop_image_from_gray(img, tol=7):\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n    mask = gray > tol\n\n    if mask.sum() == 0:\n        return img\n\n    return img[np.ix_(mask.any(1), mask.any(0))]\n\ndef read_image(path):\n    img = cv2.imread(path)\n\n    if img is None:\n        raise ValueError(f\"Image not found: {path}\")\n\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    img = crop_image_from_gray(img)\n\n    return img\n\n# =========================\n# PREDICT FUNCTION\n# =========================\n\n@torch.no_grad()\ndef predict_retinopathy(image_path):\n    img = read_image(image_path)\n    img = valid_transforms(image=img)[\"image\"]\n\n    img = torch.tensor(img, dtype=torch.float32).permute(2, 0, 1)\n    img = img.unsqueeze(0).to(DEVICE)\n\n    with torch.autocast(device_type=\"cuda\", dtype=torch.bfloat16):\n        output = model(img)\n\n    probs = F.softmax(output.float(), dim=1).cpu().numpy()[0]\n\n    continuous_score = float(np.sum(probs * np.arange(NUM_CLASSES)))\n    class_id = int(np.digitize(continuous_score, thresholds))\n\n    result = {\n        \"class_id\": class_id,\n        \"diagnosis\": CLASS_NAMES[class_id],\n        \"confidence\": float(np.max(probs)),\n        \"probabilities\": {\n            CLASS_NAMES[i]: float(probs[i]) for i in range(NUM_CLASSES)\n        }\n    }\n\n    return result\n\n# =========================\n# TEST EXAMPLE\n# =========================\n\nimage_path = \"/kaggle/input/competitions/aptos2019-blindness-detection/train_images/00cb6555d108.png\"\n\nresult = predict_retinopathy(image_path)\nresult","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-24T03:33:39.910948Z","iopub.execute_input":"2026-04-24T03:33:39.911841Z","iopub.status.idle":"2026-04-24T03:33:40.827496Z","shell.execute_reply.started":"2026-04-24T03:33:39.911802Z","shell.execute_reply":"2026-04-24T03:33:40.826803Z"}},"outputs":[],"execution_count":null}]}