{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.6.6"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":18647,"databundleVersionId":1126921,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":863262,"sourceType":"datasetVersion","datasetId":458222},{"sourceId":1115430,"sourceType":"datasetVersion","datasetId":625794},{"sourceId":11730823,"sourceType":"datasetVersion","datasetId":7363937},{"sourceId":11745411,"sourceType":"datasetVersion","datasetId":7373259},{"sourceId":22581004,"sourceType":"kernelVersion"},{"sourceId":366178,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":303691,"modelId":324173}],"dockerImageVersionId":29869,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Description\nThis kernel performs inference for [PANDA concat tile pooling starter](https://www.kaggle.com/iafoss/panda-concat-fast-ai-starter) kernel with use of multiple models and 8 fold TTA. Check it for more training details. The image preprocessing pipline is provided [here](https://www.kaggle.com/iafoss/panda-16x128x128-tiles).","metadata":{}},{"cell_type":"code","source":"import cv2\nfrom tqdm import tqdm_notebook as tqdm\nimport fastai\nfrom fastai.vision import *\nimport os\nfrom mish_activation import *\nimport warnings\nwarnings.filterwarnings(\"ignore\")\nimport skimage.io\nimport numpy as np\nimport pandas as pd\nsys.path.insert(0, '../input/semisupervised-imagenet-models/semi-supervised-ImageNet1K-models-master/')\nfrom hubconf import *\nimport cv2\nimport os\n\n# Install torchstain with error handling\n!pip install /kaggle/input/torchstain/torchstain-1.4.1-py3-none-any.whl\nimport torchstain\n\n# Load the same target image used in training\ntarget_path = '/kaggle/input/target/target.png'  # Add this image to your dataset\ntarget = cv2.cvtColor(cv2.imread(target_path), cv2.COLOR_BGR2RGB)\n\nnormalizer = torchstain.normalizers.MacenkoNormalizer(backend='numpy')\nnormalizer.fit(target)\n","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"_kg_hide-output":true,"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2025-05-09T10:02:13.592032Z","iopub.execute_input":"2025-05-09T10:02:13.592285Z","iopub.status.idle":"2025-05-09T10:02:18.655139Z","shell.execute_reply.started":"2025-05-09T10:02:13.59226Z","shell.execute_reply":"2025-05-09T10:02:18.653836Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA = '../input/prostate-cancer-grade-assessment/test_images'\nTEST = '../input/prostate-cancer-grade-assessment/test.csv'\nSAMPLE = '../input/prostate-cancer-grade-assessment/sample_submission.csv'\nMODELS = [f'../input/rnxt50/pytorch/default/1/RNXT50_{i}.pth' for i in range(4)]\n\nsz = 128\nbs = 2\nN = 12\nnworkers = 2","metadata":{"execution":{"iopub.status.busy":"2025-05-09T10:02:18.659835Z","iopub.execute_input":"2025-05-09T10:02:18.663909Z","iopub.status.idle":"2025-05-09T10:02:18.676833Z","shell.execute_reply.started":"2025-05-09T10:02:18.66385Z","shell.execute_reply":"2025-05-09T10:02:18.675799Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Load and fit stain normalizer with error fallback\ntry:\n    target_path = '/kaggle/input/target/target.png'  # Make sure to add this dataset in Kaggle\n    target_img = cv2.cvtColor(cv2.imread(target_path), cv2.COLOR_BGR2RGB)\n    normalizer = torchstain.normalizers.MacenkoNormalizer(backend='numpy')\n    normalizer.fit(target_img)\n    use_stain_norm = True\nexcept Exception as e:\n    print(\"Stain normalization setup failed:\", e)\n    use_stain_norm = False\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-09T10:02:18.678239Z","iopub.execute_input":"2025-05-09T10:02:18.678566Z","iopub.status.idle":"2025-05-09T10:02:18.705463Z","shell.execute_reply.started":"2025-05-09T10:02:18.678527Z","shell.execute_reply":"2025-05-09T10:02:18.704612Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Modify tile function to include stain normalization safely\ndef tile(img, mode=0):\n    result = []\n    img = 255 - img\n    img = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n    thresh = cv2.threshold(img, 0, 255, cv2.THRESH_OTSU)[1]\n    cnts = cv2.findContours(thresh, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    cnts = cnts[0] if len(cnts) == 2 else cnts[1]\n    coords = [cv2.boundingRect(c) for c in cnts]\n    coords = sorted(coords, key=lambda x: x[2]*x[3], reverse=True)[:36]\n    \n    for x,y,w,h in coords:\n        tile_img = img[y:y+h, x:x+w]\n        tile_img = cv2.resize(tile_img, (128, 128))\n        tile_img = cv2.cvtColor(tile_img, cv2.COLOR_GRAY2RGB)\n\n        if use_stain_norm:\n            try:\n                empty_ratio = (tile_img > 220).sum() / tile_img.size\n                if empty_ratio < 0.8:\n                    tile_img, *_ = normalizer.normalize(tile_img)\n                    tile_img = np.clip(tile_img, 0, 255).astype(np.uint8)\n            except Exception as e:\n                print(\"Stain normalization failed on a tile:\", e)\n\n        result.append(tile_img)\n\n    while len(result) < 36:\n        result.append(np.zeros((128, 128, 3), dtype=np.uint8))\n    \n    return np.stack(result)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-09T10:02:18.706836Z","iopub.execute_input":"2025-05-09T10:02:18.707168Z","iopub.status.idle":"2025-05-09T10:02:18.738162Z","shell.execute_reply.started":"2025-05-09T10:02:18.707127Z","shell.execute_reply":"2025-05-09T10:02:18.737131Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"def _resnext(url, block, layers, pretrained, progress, **kwargs):\n    model = ResNet(block, layers, **kwargs)\n    #state_dict = load_state_dict_from_url(url, progress=progress)\n    #model.load_state_dict(state_dict)\n    return model\n\nclass Model(nn.Module):\n    def __init__(self, arch='resnext50_32x4d', n=6, pre=True):\n        super().__init__()\n        #m = torch.hub.load('facebookresearch/semi-supervised-ImageNet1K-models', arch)\n        m = _resnext(semi_supervised_model_urls[arch], Bottleneck, [3, 4, 6, 3], False, \n                progress=False,groups=32,width_per_group=4)\n        self.enc = nn.Sequential(*list(m.children())[:-2])       \n        nc = list(m.children())[-1].in_features\n        self.head = nn.Sequential(AdaptiveConcatPool2d(),Flatten(),nn.Linear(2*nc,512),\n                Mish(),nn.BatchNorm1d(512),nn.Dropout(0.5),nn.Linear(512,n))\n        \n    def forward(self, x):\n        shape = x.shape\n        n = shape[1]\n        x = x.view(-1,shape[2],shape[3],shape[4])\n        x = self.enc(x)\n        shape = x.shape\n        x = x.view(-1,n,shape[1],shape[2],shape[3]).permute(0,2,1,3,4).contiguous()\\\n          .view(-1,shape[1],shape[2]*n,shape[3])\n        x = self.head(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2025-05-09T10:02:18.74098Z","iopub.execute_input":"2025-05-09T10:02:18.741564Z","iopub.status.idle":"2025-05-09T10:02:18.762052Z","shell.execute_reply.started":"2025-05-09T10:02:18.741517Z","shell.execute_reply":"2025-05-09T10:02:18.76111Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"models = []\nfor path in MODELS:\n    state_dict = torch.load(path,map_location=torch.device('cpu'))\n    model = Model()\n    model.load_state_dict(state_dict)\n    model.float()\n    model.eval()\n    model.cuda()\n    models.append(model)\n\ndel state_dict","metadata":{"execution":{"iopub.status.busy":"2025-05-09T10:02:18.763639Z","iopub.execute_input":"2025-05-09T10:02:18.763912Z","iopub.status.idle":"2025-05-09T10:02:20.5515Z","shell.execute_reply.started":"2025-05-09T10:02:18.763882Z","shell.execute_reply":"2025-05-09T10:02:20.550803Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data","metadata":{}},{"cell_type":"code","source":"def tile(img):\n    shape = img.shape\n    pad0, pad1 = (sz - shape[0] % sz) % sz, (sz - shape[1] % sz) % sz\n    img = np.pad(img, [[pad0 // 2, pad0 - pad0 // 2],\n                       [pad1 // 2, pad1 - pad1 // 2], [0, 0]], constant_values=255)\n    img = img.reshape(img.shape[0] // sz, sz, img.shape[1] // sz, sz, 3)\n    img = img.transpose(0, 2, 1, 3, 4).reshape(-1, sz, sz, 3)\n\n    img_out = []\n    for x in img:\n        try:\n            empty_ratio = (x > 220).sum() / (x.shape[0] * x.shape[1] * x.shape[2])\n            if empty_ratio < 0.8:\n                x, *_ = normalizer.normalize(x)\n                x = np.clip(x, 0, 255).astype(np.uint8)\n        except Exception:\n            print(\"no stainnorm\")\n            pass  # skip normalization on failure\n        img_out.append(x)\n\n    img = np.stack(img_out)\n    if len(img) < N:\n        img = np.pad(img, [[0, N - len(img)], [0, 0], [0, 0], [0, 0]], constant_values=255)\n    idxs = np.argsort(img.reshape(img.shape[0], -1).sum(-1))[:N]\n    return img[idxs]\n\n\n\nmean = torch.tensor([1.0-0.90949707, 1.0-0.8188697, 1.0-0.87795304])\nstd = torch.tensor([0.36357649, 0.49984502, 0.40477625])\n\nclass PandaDataset(Dataset):\n    def __init__(self, img_dir, submission_csv):\n        self.img_dir = img_dir\n        # read full list of test IDs\n        self.names = pd.read_csv(submission_csv).image_id.tolist()\n\n    def __len__(self):\n        return len(self.names)\n\n    def __getitem__(self, idx):\n        name = self.names[idx]\n        img_path = os.path.join(self.img_dir, name + '.tiff')\n\n        # try to load the tissue image, else use blank\n        try:\n            frames = skimage.io.MultiImage(img_path)\n            if len(frames) == 0:\n                raise ValueError(\"no frames\")\n            img = frames[-1]\n        except Exception:\n            print(f\"Warning: failed to load {name}, using blank fallback.\")\n            img = np.ones((sz*6, sz*6, 3), dtype=np.uint8) * 255\n\n        # tile + optional stain norm\n        tiles = tile(img)  # your tile() already does stain‐norm if configured\n        # to tensor, invert, normalize\n        t = torch.Tensor(1.0 - tiles/255.0)\n        t = (t - mean) / std\n        return t.permute(0,3,1,2), name\n","metadata":{"execution":{"iopub.status.busy":"2025-05-09T10:31:25.482041Z","iopub.execute_input":"2025-05-09T10:31:25.482295Z","iopub.status.idle":"2025-05-09T10:31:25.497328Z","shell.execute_reply.started":"2025-05-09T10:31:25.482272Z","shell.execute_reply":"2025-05-09T10:31:25.496611Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prediction","metadata":{}},{"cell_type":"code","source":"sub_df = pd.read_csv(SAMPLE)\nif os.path.exists(DATA):\n    ds = PandaDataset(DATA, TEST)\n    dl = DataLoader(ds, batch_size=bs, num_workers=nworkers, shuffle=False)\n    names, preds = [], []\n\n    with torch.no_grad():\n        for x, y in tqdm(dl):\n            try:\n                x = x.cuda()\n                # TTA\n                x = torch.stack([\n                    x, x.flip(-1), x.flip(-2), x.flip(-1, -2),\n                    x.transpose(-1, -2), x.transpose(-1, -2).flip(-1),\n                    x.transpose(-1, -2).flip(-2), x.transpose(-1, -2).flip(-1, -2)\n                ], 1)\n                x = x.view(-1, N, 3, sz, sz)\n                p = [model(x) for model in models]\n                p = torch.stack(p, 1)\n                p = p.view(bs, 8 * len(models), -1).mean(1).argmax(-1).cpu()\n                names.append(y)\n                preds.append(p)\n            except Exception as e:\n                print(f\"Inference failed for {y}: {e}\")\n    \n    names = np.concatenate(names)\n    preds = torch.cat(preds).numpy()\n    sub_df = pd.DataFrame({'image_id': names, 'isup_grade': preds})\n    sub_df.to_csv('submission.csv', index=False)\n    sub_df.head()","metadata":{"execution":{"iopub.status.busy":"2025-05-09T10:31:28.810924Z","iopub.execute_input":"2025-05-09T10:31:28.811197Z","iopub.status.idle":"2025-05-09T10:31:28.835443Z","shell.execute_reply.started":"2025-05-09T10:31:28.81117Z","shell.execute_reply":"2025-05-09T10:31:28.834795Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub_df.to_csv(\"submission.csv\", index=False)\nsub_df.head()","metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","execution":{"iopub.status.busy":"2025-05-09T10:31:32.472161Z","iopub.execute_input":"2025-05-09T10:31:32.472404Z","iopub.status.idle":"2025-05-09T10:31:32.481781Z","shell.execute_reply.started":"2025-05-09T10:31:32.472382Z","shell.execute_reply":"2025-05-09T10:31:32.480937Z"},"trusted":true},"outputs":[],"execution_count":null}]}