{"cells":[{"cell_type":"markdown","metadata":{},"source":"# Raptor CoAtNet family standalone submission\n"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"import os, glob, time, gc, hashlib\nos.environ.setdefault('HF_HUB_OFFLINE', '1')\nos.environ.setdefault('TRANSFORMERS_OFFLINE', '1')\nos.environ.setdefault('HF_HUB_DISABLE_TELEMETRY', '1')\nimport numpy as np\nimport torch, torch.nn as nn, torch.nn.functional as F\nimport timm\ntorch.backends.cudnn.benchmark = True\ntorch.backends.cuda.matmul.allow_tf32 = True\nIMG = 336\nCROP_MM = 140.0\nSPAN_LO, SPAN_HI = 0.02, 0.98\nSLOTS = [(\"Sagittal\", 1, 18), (\"Sagittal\", 0, 14),\n         (\"Coronal\", 1, 12), (\"Coronal\", 0, 8), (\"Axial\", -1, 12)]\nMAXS = sum(slot[2] for slot in SLOTS)\nK_EVAL = 62\nNORM = \"imagenet\"\nLAB = [\"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\", \"Medial OA\",\n       \"Lateral OA\", \"PF OA\", \"Effusion\", \"Synovitis\", \"Baker's\",\n       \"Contusion\", \"Fracture\"]\n_MEAN = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)\n_STD = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)\n_SLOTS64 = [(\"Sagittal\", 1, 18), (\"Sagittal\", 0, 14),\n            (\"Coronal\", 1, 12), (\"Coronal\", 0, 8), (\"Axial\", -1, 12)]\n_SLOTS44 = [(\"Sagittal\", 1, 12), (\"Sagittal\", 0, 10),\n            (\"Coronal\", 1, 8), (\"Coronal\", 0, 6), (\"Axial\", -1, 8)]\nARMS = [\n    {\"name\": \"maxspan-v5\", \"file\": \"raptor_ft_coatnet_v5_full_swa.pt\",\n     \"arch\": \"coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k\", \"res\": 384,\n     \"img\": 336, \"slots\": _SLOTS64, \"span\": (0.02, 0.98), \"k_eval\": 62,\n     \"reverse\": False, \"w\": 0.60},\n    {\"name\": \"native384dense-v10\", \"file\": \"raptor_ft_coatnet_v10_full.pt\",\n     \"arch\": \"coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k\", \"res\": 384,\n     \"img\": 384, \"slots\": _SLOTS64, \"span\": (0.02, 0.98), \"k_eval\": 62,\n     \"reverse\": False, \"w\": 0.10},\n    {\"name\": \"maxspan-v5-reverse\", \"file\": \"raptor_ft_coatnet_v5_full_swa.pt\",\n     \"arch\": \"coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k\", \"res\": 384,\n     \"img\": 336, \"slots\": _SLOTS64, \"span\": (0.02, 0.98), \"k_eval\": 62,\n     \"reverse\": True, \"w\": 0.10},\n    {\"name\": \"native384-v8\", \"file\": \"raptor_ft_coatnet_v8_full_swa.pt\",\n     \"arch\": \"coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k\", \"res\": 384,\n     \"img\": 384, \"slots\": _SLOTS44, \"span\": (0.06, 0.94), \"k_eval\": 42,\n     \"reverse\": False, \"w\": 0.20},\n]\n\ndef build_backbone(arch, pretrained=False):\n    hybrid = arch.startswith(('maxvit', 'maxxvit', 'coatnet', 'coat_', 'convnext'))\n    is_vit = not hybrid and any((k in arch for k in ('vit', 'deit', 'dinov2', 'eva', 'beit')))\n    kw = dict(pretrained=pretrained, num_classes=0, in_chans=3)\n    if is_vit:\n        kw.update(global_pool='token', dynamic_img_size=True)\n    else:\n        kw.update(global_pool='avg')\n    return timm.create_model(arch, **kw)\n\nclass RaptorClassifier(nn.Module):\n\n    def __init__(self, backbone, F_dim=768, n=12, drop=0.2):\n        super().__init__()\n        self.backbone = backbone\n        self.norm = nn.LayerNorm(F_dim)\n        self.att = nn.Sequential(nn.Linear(F_dim, 256), nn.Tanh(), nn.Dropout(drop), nn.Linear(256, n))\n        self.clsW = nn.Parameter(torch.zeros(n, F_dim))\n        self.clsb = nn.Parameter(torch.zeros(n))\n        nn.init.trunc_normal_(self.clsW, std=0.02)\n        self.n = n\n\n    def encode(self, x):\n        B, K = x.shape[:2]\n        f = self.backbone(x.flatten(0, 1))\n        return f.view(B, K, -1)\n\n    def head(self, feats):\n        h = self.norm(feats)\n        a = self.att(h)\n        a = torch.softmax(a, dim=1)\n        pooled = torch.einsum('bkn,bkf->bnf', a, h)\n        logits = (pooled * self.clsW).sum(-1) + self.clsb\n        return logits\n\n    def forward(self, x):\n        return self.head(self.encode(x))\n\ndef load_model(pt_path, arch_default, res_default, device, ngpu=1):\n    ck = torch.load(pt_path, map_location='cpu', weights_only=False)\n    arch = ck.get('arch', arch_default)\n    ck_res = int(ck.get('res', res_default))\n    bb = build_backbone(arch, pretrained=False)\n    model = RaptorClassifier(bb, F_dim=bb.num_features)\n    model.load_state_dict(ck['model'], strict=True)\n    model.eval().to(device)\n    del ck\n    gc.collect()\n    return (model, ck_res)\n\ndef load_refit_head(pt_path, feature_dim, device):\n    ck = torch.load(pt_path, map_location='cpu', weights_only=False)\n    state = ck.get('model', ck)\n    head = RaptorClassifier(nn.Identity(), F_dim=int(feature_dim))\n    head_state = {name: tensor for name, tensor in state.items()\n                  if not name.startswith('backbone.')}\n    head.load_state_dict(head_state, strict=True)\n    head.eval().to(device)\n    del ck, state, head_state\n    gc.collect()\n    return head\n\ndef _eval_centers(mask, D, k):\n    valid = np.where(mask > 0)[0]\n    if len(valid) < 3:\n        valid = np.arange(min(3, D))\n    lo, hi = (int(valid.min()), int(valid.max()))\n    cs = [c for c in range(lo + 1, hi) if c - 1 >= lo and c + 1 <= hi]\n    if not cs:\n        cs = [max(1, min((lo + hi) // 2, D - 2))]\n    idx = np.linspace(0, len(cs) - 1, k).round().astype(int)\n    return [cs[i] for i in idx]\n\ndef eval_windows(vol, mask, k, res, norm=NORM):\n    D = vol.shape[0]\n    cs = _eval_centers(mask, D, k)\n    wins = np.empty((len(cs), 3, res, res), np.float32)\n    for j, c in enumerate(cs):\n        c = max(1, min(c, D - 2))\n        tri = np.stack([vol[c - 1], vol[c], vol[c + 1]], 0).astype(np.float32) / 255.0\n        t = torch.from_numpy(tri)\n        if t.shape[-1] != res:\n            t = F.interpolate(t[None], size=(res, res), mode='bilinear', align_corners=False)[0]\n        wins[j] = t.numpy()\n    x = torch.from_numpy(wins)\n    if norm == 'imagenet':\n        x = (x - _MEAN) / _STD\n    return x\n\n@torch.no_grad()\ndef infer_probs(model, xwins, device):\n    x = xwins.unsqueeze(0).to(device)\n    use_cuda = device != 'cpu' and str(device).startswith('cuda')\n    if use_cuda:\n        try:\n            with torch.autocast('cuda', dtype=torch.float16):\n                o = torch.sigmoid(model(x).float())\n            return o[0].cpu().numpy()\n        except RuntimeError:\n            torch.cuda.empty_cache()\n            o = torch.sigmoid(model(x).float())\n            return o[0].cpu().numpy()\n    o = torch.sigmoid(model(x).float())\n    return o[0].cpu().numpy()\n\n@torch.no_grad()\ndef infer_probs_two_heads(model, refit_head, xwins, device):\n    x = xwins.unsqueeze(0).to(device)\n    use_cuda = device != 'cpu' and str(device).startswith('cuda')\n    def forward_heads():\n        features = model.encode(x)\n        original = torch.sigmoid(model.head(features).float())\n        refitted = torch.sigmoid(refit_head.head(features).float())\n        return (original[0].cpu().numpy(), refitted[0].cpu().numpy())\n    if use_cuda:\n        try:\n            with torch.autocast('cuda', dtype=torch.float16):\n                return forward_heads()\n        except RuntimeError:\n            torch.cuda.empty_cache()\n            return forward_heads()\n    return forward_heads()\n\ndef rankpct(x):\n    order = x.argsort(0).argsort(0).astype(np.float64)\n    return order / max(1, x.shape[0] - 1)\n\ndef _make_reader():\n    import pydicom, cv2\n    from pydicom.pixel_data_handlers.util import apply_modality_lut\n\n    def order_and_meta(sdir):\n        fs = glob.glob(sdir + '/*.dcm')\n        recs = []\n        ps_list = []\n        for f in fs:\n            try:\n                h = pydicom.dcmread(f, stop_before_pixels=True)\n                iop = getattr(h, 'ImageOrientationPatient', None)\n                ipp = getattr(h, 'ImagePositionPatient', None)\n                if iop is not None and ipp is not None and (len(iop) == 6):\n                    r = np.array(iop[:3], float)\n                    c = np.array(iop[3:], float)\n                    n = np.cross(r, c)\n                    pos = float(np.dot(np.array(ipp, float), n))\n                else:\n                    pos = float(getattr(h, 'InstanceNumber', 0) or 0)\n                ps = getattr(h, 'PixelSpacing', None)\n                ps = float(ps[0]) if ps is not None else 0.5\n                ps_list.append(ps)\n                recs.append((pos, f, ps))\n            except Exception:\n                recs.append((0.0, f, 0.5))\n        recs.sort(key=lambda x: x[0])\n        med_ps = float(np.median(ps_list)) if ps_list else 0.5\n        return ([(f, ps) for _, f, ps in recs], med_ps)\n\n    def read_px(f):\n        d = pydicom.dcmread(f)\n        a = apply_modality_lut(d.pixel_array, d).astype(np.float32)\n        if str(getattr(d, 'PhotometricInterpretation', '')) == 'MONOCHROME1':\n            a = a.max() - a\n        return a\n\n    def mm_crop_resize(a, ps):\n        h, w = a.shape\n        cpx = int(round(CROP_MM / max(ps, 0.001)))\n        cpx = min(cpx, min(h, w))\n        y0 = (h - cpx) // 2\n        x0 = (w - cpx) // 2\n        a = a[y0:y0 + cpx, x0:x0 + cpx]\n        return cv2.resize(a, (IMG, IMG), interpolation=cv2.INTER_AREA)\n    return (order_and_meta, read_px, mm_crop_resize)\n\ndef _pick_series_for_slot(rows, plane, fluid, used):\n    cands = [r for r in rows if r['Anatomical_Plane'] == plane and r['SeriesInstanceUID'] not in used]\n    if fluid in (0, 1):\n        pref = [r for r in cands if int(r.get('Fluid_Sensitive', 0) or 0) == fluid]\n        if pref:\n            return pref[0]\n    return cands[0] if cands else None\n\ndef build_study(sid, ser_records, tsdir, reader):\n    order_and_meta, read_px, mm_crop_resize = reader\n    rows = ser_records.get(sid, [])\n    vol = np.zeros((MAXS, IMG, IMG), np.uint8)\n    idx = 0\n    used = set()\n    for plane, fluid, k in SLOTS:\n        r = _pick_series_for_slot(rows, plane, fluid, used)\n        if r is None:\n            idx += k\n            continue\n        used.add(r['SeriesInstanceUID'])\n        files, med_ps = order_and_meta(f\"{tsdir}/{sid}/{r['SeriesInstanceUID']}\")\n        if not files:\n            idx += k\n            continue\n        n = len(files)\n        lo, hi = (int(n * SPAN_LO), int(n * SPAN_HI) - 1)\n        hi = max(hi, lo)\n        picks = np.linspace(lo, hi, k).round().astype(int) if n > 1 else [0] * k\n        arrs = []\n        pss = []\n        for p in picks:\n            fp, ps = files[min(p, n - 1)]\n            try:\n                arrs.append(read_px(fp))\n                pss.append(ps)\n            except Exception:\n                arrs.append(None)\n                pss.append(med_ps)\n        valid = [a for a in arrs if a is not None]\n        if valid:\n            allpx = np.concatenate([a.ravel() for a in valid])\n            loq, hiq = np.percentile(allpx, [2.0, 98.0])\n        else:\n            loq, hiq = (0.0, 1.0)\n        for a, ps in zip(arrs, pss):\n            if idx >= MAXS:\n                break\n            if a is None:\n                idx += 1\n                continue\n            aw = np.clip((a - loq) / (hiq - loq + 1e-06), 0, 1)\n            aw = mm_crop_resize(aw, ps if ps > 0 else med_ps)\n            vol[idx] = (aw * 255).astype(np.uint8)\n            idx += 1\n        if idx >= MAXS:\n            break\n    mask = (vol.reshape(MAXS, -1).sum(1) > 0).astype(np.uint8)\n    return (vol, mask)\n\ndef find_test_root():\n    cands = ['/kaggle/input/competitions/rsna-knee-abnormality-detection', '/kaggle/input/rsna-knee-abnormality-detection']\n    for b in cands:\n        if os.path.exists(b + '/test.csv'):\n            return b\n    for d, _, f in os.walk('/kaggle/input'):\n        if 'test.csv' in f and (os.path.isdir(d + '/test_series') or os.path.isdir(d + '/test_images')):\n            return d\n    for d, _, f in os.walk('/kaggle/input'):\n        if 'test.csv' in f:\n            return d\n    raise RuntimeError('no test root under /kaggle/input')\n\ndef find_weight_file(fname):\n    direct = [f'/kaggle/input/raptor-knee-maxspan/{fname}', f'/kaggle/input/raptor-knee-native384dense/{fname}', f'/kaggle/input/raptor-knee-native384/{fname}', f'/kaggle/input/raptor-knee-arms/{fname}', f'/kaggle/input/raptor-knee-arms/1/{fname}', f'/kaggle/input/raptor-cnn336/{fname}']\n    for p in direct:\n        if os.path.exists(p):\n            return p\n    for d in sorted(glob.glob('/kaggle/input/*/')):\n        if 'competition' in d.lower():\n            continue\n        hits = glob.glob(os.path.join(d, '**', fname), recursive=True)\n        if hits:\n            return hits[0]\n    raise RuntimeError(f'{fname} not found under /kaggle/input')\n\ndef find_optional_verified_weight(fname, expected_sha256, root='/kaggle/input'):\n    hits = []\n    for directory in sorted(glob.glob(os.path.join(root, '*/'))):\n        if 'competition' in directory.lower():\n            continue\n        hits.extend(glob.glob(os.path.join(directory, '**', fname), recursive=True))\n    for path in sorted(set(hits)):\n        digest = hashlib.sha256()\n        with open(path, 'rb') as stream:\n            for chunk in iter(lambda: stream.read(8 * 1024 * 1024), b''):\n                digest.update(chunk)\n        if digest.hexdigest() == expected_sha256:\n            return path\n        print(f'[head-refit] ignored hash-mismatched optional checkpoint: {path}', flush=True)\n    return None\n\ndef main():\n    import pandas as pd\n    t0 = time.time()\n    dev = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    print(f\"device {dev} | gpus {torch.cuda.device_count()} | torch {torch.__version__}\", flush=True)\n    root = find_test_root()\n    tsdir = root + \"/test_series\"\n    if not os.path.isdir(tsdir):\n        tsdir = root + \"/test_images\"\n    print(\"test root:\", root, \"| series dir:\", tsdir, flush=True)\n    test = pd.read_csv(root + \"/test.csv\")\n    test[\"StudyInstanceUID\"] = test[\"StudyInstanceUID\"].astype(str)\n    test_ids = test[\"StudyInstanceUID\"].tolist()\n    tser = pd.read_csv(root + \"/test_series.csv\")\n    tser[\"StudyInstanceUID\"] = tser[\"StudyInstanceUID\"].astype(str)\n    tser[\"SeriesInstanceUID\"] = tser[\"SeriesInstanceUID\"].astype(str)\n    series = {key: frame.to_dict(\"records\") for key, frame in tser.groupby(\"StudyInstanceUID\")}\n    print(f\"test studies {len(test_ids)} | test series {len(tser)}\", flush=True)\n    sub_cols = [\"StudyInstanceUID\"] + LAB\n    sample = os.path.join(root, \"sample_submission.csv\")\n    if os.path.exists(sample):\n        sub_cols = list(pd.read_csv(sample, nrows=1).columns)\n    reader = _make_reader()\n    n_study, n_arm = len(test_ids), len(ARMS)\n    arm_probs = [np.full((n_study, len(LAB)), 0.5, np.float32) for _ in range(n_arm)]\n    for arm_index, arm in enumerate(ARMS):\n        globals()[\"IMG\"] = int(arm[\"img\"])\n        globals()[\"SLOTS\"] = list(arm[\"slots\"])\n        globals()[\"MAXS\"] = sum(slot[2] for slot in SLOTS)\n        globals()[\"SPAN_LO\"], globals()[\"SPAN_HI\"] = map(float, arm[\"span\"])\n        globals()[\"K_EVAL\"] = int(arm[\"k_eval\"])\n        weight_path = find_weight_file(arm[\"file\"])\n        model, resolution = load_model(weight_path, arm[\"arch\"], arm[\"res\"], dev)\n        print(f\"[arm {arm_index}] {arm['name']} | img {IMG} | slices {MAXS} | \"\n              f\"span {SPAN_LO:.2f}-{SPAN_HI:.2f} | windows {K_EVAL} | \"\n              f\"res {resolution} | {time.time() - t0:.0f}s\", flush=True)\n        for study_index, study_uid in enumerate(test_ids):\n            try:\n                volume, mask = build_study(study_uid, series, tsdir, reader)\n                windows = eval_windows(volume, mask, k=K_EVAL, res=resolution, norm=NORM)\n                if bool(arm.get(\"reverse\", False)):\n                    windows = windows.flip(1).contiguous()\n                arm_probs[arm_index][study_index] = infer_probs(model, windows, dev)\n                del volume, mask, windows\n            except Exception as error:\n                print(f\"  [arm {arm_index}] study {study_index} {study_uid[:16]} FALLBACK \"\n                      f\"({type(error).__name__}: {error})\", flush=True)\n            if (study_index + 1) % 100 == 0 or study_index + 1 == n_study:\n                print(f\"  [arm {arm_index}] {study_index + 1}/{n_study} | \"\n                      f\"{time.time() - t0:.0f}s\", flush=True)\n        del model\n        gc.collect()\n        if str(dev).startswith(\"cuda\"):\n            torch.cuda.empty_cache()\n        print(f\"[arm {arm_index}] done + freed | {time.time() - t0:.0f}s\", flush=True)\n    weights = np.array([float(arm.get(\"w\", 1.0)) for arm in ARMS], dtype=np.float64)\n    weights /= weights.sum()\n    print(f\"[blend] global probability mean w=\"\n          f\"{dict(zip([arm['name'] for arm in ARMS], weights.round(4)))}\", flush=True)\n    probability_blend = np.tensordot(\n        weights, np.stack([np.clip(values, 0, 1) for values in arm_probs]), axes=(0, 0))\n    ranks = rankpct(probability_blend)\n    if not np.isfinite(ranks).all():\n        ranks[~np.isfinite(ranks)] = 0.5\n    submission = pd.DataFrame(ranks.astype(np.float32), columns=LAB)\n    submission.insert(0, \"StudyInstanceUID\", test_ids)\n    submission = submission[sub_cols]\n    assert submission[\"StudyInstanceUID\"].tolist() == test_ids\n    assert np.isfinite(submission[LAB].values).all()\n    out = \"/kaggle/working/submission.csv\"\n    submission.to_csv(out, index=False)\n    print(\"wrote\", out, \"|\", len(submission), \"rows x\", len(submission.columns), \"cols\", flush=True)\n    print(submission.head().to_string(index=False), flush=True)\n    print(f\"DONE {time.time() - t0:.0f}s\", flush=True)\n\n\nmain()\n"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11.13"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":154281,"sourceType":"competition"},{"sourceId":18842180,"sourceType":"datasetVersion"},{"sourceId":18839182,"sourceType":"datasetVersion"},{"sourceId":18229736,"sourceType":"datasetVersion"},{"sourceId":18956429,"sourceType":"datasetVersion"},{"sourceId":19075719,"sourceType":"datasetVersion"},{"sourceId":18875869,"sourceType":"datasetVersion"},{"sourceId":18673646,"sourceType":"datasetVersion"},{"sourceId":18673450,"sourceType":"datasetVersion"},{"sourceId":18716507,"sourceType":"datasetVersion"},{"sourceId":18879001,"sourceType":"datasetVersion"},{"sourceId":18757740,"sourceType":"datasetVersion"},{"sourceId":342671664,"sourceType":"kernelVersion"},{"sourceId":342849430,"sourceType":"kernelVersion"},{"sourceId":4533,"sourceType":"modelInstanceVersion"}],"dockerImageVersionId":31430,"isGpuEnabled":true,"isInternetEnabled":false,"language":"python","sourceType":"notebook"},"rsna_optimization":{"official_source_score":0.891,"revision":"v66-v65-parent-legacy-dino-002","source":"pilkwang/rsna-knee-baseline-v1"}},"nbformat":4,"nbformat_minor":4}