{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","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":24800,"datasetId":1042002,"databundleVersionId":1831594},{"sourceType":"datasetVersion","sourceId":12425246,"datasetId":7837104,"databundleVersionId":12993792},{"sourceType":"datasetVersion","sourceId":12121793,"datasetId":7632744,"databundleVersionId":12653139}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q segmentation-models-pytorch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-10T08:19:17.333427Z","iopub.execute_input":"2025-07-10T08:19:17.334184Z","iopub.status.idle":"2025-07-10T08:19:20.777218Z","shell.execute_reply.started":"2025-07-10T08:19:17.334145Z","shell.execute_reply":"2025-07-10T08:19:20.776361Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport numpy as np\nimport pydicom\nfrom PIL import Image\nfrom torchvision import transforms\nfrom torchvision.models import resnet101, densenet121\nimport segmentation_models_pytorch as smp\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# ========= 🔹 Step 1: Load All Models Once\ndef load_ensemble_models():\n    # Transforms\n    transform_chexnet = transforms.Compose([\n        transforms.Resize((256, 256)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.5482]*3, std=[0.2667]*3)\n    ])\n    transform_seg = transforms.Compose([\n        transforms.Resize((512, 512)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.5482]*3, std=[0.2667]*3)\n    ])\n\n    # Segmentation model\n    seg_model = smp.Unet(encoder_name='resnet34', in_channels=3, classes=1)\n    seg_model.load_state_dict(torch.load(\"/kaggle/input/x-ray-segmention-model/xray_Segmention_model.pth\", map_location=DEVICE))\n    seg_model.eval().to(DEVICE)\n\n    # Classification models\n    def get_resnet_model():\n        model = resnet101(weights=None)\n        model.fc = nn.Sequential(nn.Dropout(0.5), nn.Linear(2048, 2))\n        return model.to(DEVICE)\n\n    def get_chexnet_model():\n        model = densenet121(weights=None)\n        model.classifier = nn.Linear(1024, 2)\n        return model.to(DEVICE)\n\n    model_seg = get_resnet_model()\n    model_seg.load_state_dict(torch.load(\"/kaggle/input/models-normal-vs-abnormal/best_fc_only_resnet101_epoch10_Segmentation.pth\", map_location=DEVICE))\n    model_seg.eval()\n\n    model_chex_all = get_chexnet_model()\n    model_chex_all.load_state_dict(torch.load(\"/kaggle/input/models-normal-vs-abnormal/chexnet_fc_only_epoch7.pth\", map_location=DEVICE))\n    model_chex_all.eval()\n\n    model_chex_2 = get_chexnet_model()\n    model_chex_2.load_state_dict(torch.load(\"/kaggle/input/models-normal-vs-abnormal/chexnet_fc_only_epoch9_2classes.pth\", map_location=DEVICE))\n    model_chex_2.eval()\n\n    return {\n        \"transform_chexnet\": transform_chexnet,\n        \"transform_seg\": transform_seg,\n        \"seg_model\": seg_model,\n        \"model_seg\": model_seg,\n        \"model_chex_all\": model_chex_all,\n        \"model_chex_2\": model_chex_2\n    }\n\n# ========= 🔹 Step 2: Predict Single DICOM Image\ndef predict_image_class(dicom_path, models_dict, threshold=0.71):\n    # Load image\n    dicom = pydicom.dcmread(dicom_path)\n    image_array = dicom.pixel_array.astype(np.float32)\n    image_array -= image_array.min()\n    image_array /= image_array.max()\n    image_array *= 255\n    image_array = image_array.astype(np.uint8)\n\n    if len(image_array.shape) == 2:\n        image_pil = Image.fromarray(image_array).convert(\"RGB\")\n    else:\n        image_pil = Image.fromarray(image_array)\n\n    # Transforms\n    transform_chexnet = models_dict[\"transform_chexnet\"]\n    transform_seg = models_dict[\"transform_seg\"]\n\n    img_chexnet = transform_chexnet(image_pil).unsqueeze(0).to(DEVICE)\n    img_seg_input = transform_seg(image_pil).unsqueeze(0).to(DEVICE)\n\n    # Apply segmentation mask\n    with torch.no_grad():\n        mask = torch.sigmoid(models_dict[\"seg_model\"](img_seg_input)).squeeze().cpu().numpy()\n        mask = (mask > 0.5).astype(np.float32)\n\n    image_resized = image_pil.resize((512, 512))\n    image_np = np.array(image_resized).astype(np.float32) / 255.0\n    masked_image = image_np * np.expand_dims(mask, axis=-1)\n    masked_pil = Image.fromarray((masked_image * 255).astype(np.uint8))\n    img_segmented = transform_chexnet(masked_pil).unsqueeze(0).to(DEVICE)\n\n    # Get probabilities from each model\n    def get_prob(model, img_tensor):\n        with torch.no_grad():\n            out = model(img_tensor)\n            return torch.softmax(out, dim=1)[0, 1].item()\n\n    p1 = get_prob(models_dict[\"model_seg\"], img_segmented)\n    p2 = get_prob(models_dict[\"model_chex_all\"], img_chexnet)\n    p3 = get_prob(models_dict[\"model_chex_2\"], img_chexnet)\n\n    probs = [p1, p2, p3]\n    final_prob = np.mean(probs)\n    pred_class = int(final_prob >= threshold)\n    label = \"Abnormal\" if pred_class else \"Normal\"\n\n    # print(f\"\\n📁 DICOM File: {dicom_path}\")\n    # print(f\"📊 Ensemble Probability: {final_prob:.4f}\")\n    # print(f\"🧠 Prediction: {label}\")\n\n    return label\n#once\nmodels_dict = load_ensemble_models()\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-07-10T08:22:32.77853Z","iopub.execute_input":"2025-07-10T08:22:32.77888Z","iopub.status.idle":"2025-07-10T08:22:34.854204Z","shell.execute_reply.started":"2025-07-10T08:22:32.778854Z","shell.execute_reply":"2025-07-10T08:22:34.853593Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predict_image_class(\"/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train/000434271f63a053c4128a0ba6352c7f.dicom\", models_dict)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-10T08:22:35.153958Z","iopub.execute_input":"2025-07-10T08:22:35.154224Z","iopub.status.idle":"2025-07-10T08:22:35.416641Z","shell.execute_reply.started":"2025-07-10T08:22:35.154204Z","shell.execute_reply":"2025-07-10T08:22:35.416037Z"}},"outputs":[],"execution_count":null}]}