{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"cada716a-e3d5-4470-9b47-e05eedb7f4fb","cell_type":"markdown","source":"# RSNA Knee — Submission (baseline)\n\nPipeline: DICOM → canonicalización → 32 cortes equiespaciados → DINOv2 ViT-B/14 → 5 cabezas promediadas → `submission.csv`.\n\n\n| Dataset | Contenido |\n|---|---|\n| `rsna-wheels` | wheels offline |\n| `dinov2-vitb14-reg` | el `model.safetensors`|\n| `rsna-heads` | `head_f0.pt` … `head_f4.pt` |\n\n**Antes del submission real:** poner `FEAT` según lo que gane (`cls` vs `both`) y portar la función de crop exacta de tu script de preprocesado (marcada con TODO).","metadata":{}},{"id":"8aa420a7-1afd-4e17-82bb-fc6e97b385c5","cell_type":"code","source":"import os, glob, time, collections\nimport numpy as np, pandas as pd\n\nT0 = time.time()\nBUDGET_S = 9 * 3600\n\nCOMP       = '/kaggle/input/competitions/rsna-knee-abnormality-detection'\nD_WHEELS   = '/kaggle/input/datasets/sadamtorres/rsna-wheels'\nD_BACKBONE = '/kaggle/input/datasets/sadamtorres/dinov2-vitb14-reg'\nD_HEADS    = '/kaggle/input/datasets/sadamtorres/rsna-heads'\n\nFEAT = 'both'          # <-- 'cls' o 'both', DEBE coincidir con el entrenamiento de las cabezas\nK    = 32\nIMG  = 224\n\nCLASSES = ['ACL','MCL','Medial Meniscus','Lateral Meniscus','Medial OA','Lateral OA',\n           'PF OA','Effusion','Synovitis',\"Baker's\",'Contusion','Fracture']\n# prevalencias de train (media de labels blandos) — fallback si un estudio falla\nPREV = dict(zip(CLASSES, [0.196,0.103,0.429,0.210,0.228,0.157,0.307,0.410,0.407,0.242,0.216,0.223]))\n\nTEST_DIR = f'{COMP}/test_series'\ntest_studies = sorted(os.listdir(TEST_DIR))\nprint(len(test_studies), 'estudios de test')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T21:38:02.52428Z","iopub.execute_input":"2026-08-08T21:38:02.525301Z","iopub.status.idle":"2026-08-08T21:38:02.534047Z","shell.execute_reply.started":"2026-08-08T21:38:02.525267Z","shell.execute_reply":"2026-08-08T21:38:02.53311Z"}},"outputs":[],"execution_count":null},{"id":"49623aa7-07ff-464b-be13-b1c17632ebc1","cell_type":"code","source":"# --- instalación offline (defensa contra JPEG 2000 en test) -----------------\n# train era 100% Explicit VR LE sin comprimir, pero el test puede no serlo.\n!pip install --no-index --find-links {D_WHEELS} \\\n    pylibjpeg pylibjpeg-libjpeg pylibjpeg-openjpeg python-gdcm -q\n\nimport pydicom, cv2, torch\nprint('pydicom', pydicom.__version__, '| cuda', torch.cuda.is_available())\n\n# censo rápido de transfer syntaxes en el test (1 archivo por estudio, ~1 min)\nts = collections.Counter()\nfor s in test_studies[:200]:\n    f = glob.glob(f'{TEST_DIR}/{s}/*/*.dcm')\n    if f:\n        ts[str(pydicom.dcmread(f[0], stop_before_pixels=True).file_meta.TransferSyntaxUID.name)] += 1\nprint(ts)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T21:38:05.330102Z","iopub.execute_input":"2026-08-08T21:38:05.330493Z","iopub.status.idle":"2026-08-08T21:38:08.597078Z","shell.execute_reply.started":"2026-08-08T21:38:05.330465Z","shell.execute_reply":"2026-08-08T21:38:08.596117Z"}},"outputs":[],"execution_count":null},{"id":"866055dd-62ac-436d-9d42-f487fc4380d7","cell_type":"code","source":"# --- backbone (offline, con interpolación de pos-embeddings 518→224) ---------\nimport timm\n\nmodel = timm.create_model(\n    'vit_base_patch14_reg4_dinov2', pretrained=True, num_classes=0, img_size=IMG,\n    pretrained_cfg_overlay=dict(file=f'{D_BACKBONE}/model.safetensors'))\nmodel = model.eval().cuda()\nNPREFIX = model.num_prefix_tokens          # 5 = 1 CLS + 4 registers\nD_PROBE = model.num_features * (2 if FEAT == 'both' else 1)\nIM_MEAN = np.array([0.485,0.456,0.406], np.float32)\nIM_STD  = np.array([0.229,0.224,0.225], np.float32)\nprint('backbone listo, d_in de la cabeza =', D_PROBE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T21:38:10.067973Z","iopub.execute_input":"2026-08-08T21:38:10.068624Z","iopub.status.idle":"2026-08-08T21:38:11.374279Z","shell.execute_reply.started":"2026-08-08T21:38:10.068589Z","shell.execute_reply":"2026-08-08T21:38:11.373422Z"}},"outputs":[],"execution_count":null},{"id":"c3be498a-c093-4dce-928b-3b0fb7a597d9","cell_type":"markdown","source":"## Preprocesado DICOM\n\nDebe replicar el preprocesado de train **exactamente**: misma canonicalización, misma normalización, mismo crop. El crop está marcado como TODO — portar la función real de tu script de prep y verificar contra 2-3 estudios de train (misma salida píxel a píxel tras el resize).","metadata":{}},{"id":"42699998-507e-4adf-b2ad-817a9c6cfbad","cell_type":"code","source":"PLANES = {'sagittal': 0, 'coronal': 1, 'axial': 2}\nDEAD_ZONE_MM = 25.0\n\ndef plane_from_iop(iop):\n    row, col = np.array(iop[:3], float), np.array(iop[3:], float)\n    n = np.cross(row, col)\n    return ['sagittal','coronal','axial'][int(np.argmax(np.abs(n)))], row, col\n\ndef center_x(ds):\n    \"\"\"x del CENTRO de la imagen (IPP es la esquina; sin la corrección la\n    concordancia con el tag cae de 96.9% a 58.8%).\"\"\"\n    ipp = np.array(ds.ImagePositionPatient, float)\n    row = np.array(ds.ImageOrientationPatient[:3], float)\n    col = np.array(ds.ImageOrientationPatient[3:], float)\n    py, px = map(float, ds.PixelSpacing)          # [row_spacing, col_spacing]\n    c = ipp + row * px * (int(ds.Columns) / 2) + col * py * (int(ds.Rows) / 2)\n    return float(c[0])\n\ndef study_laterality(headers):\n    \"\"\"Geometría manda fuera de la zona muerta de ±25 mm; dentro, manda el tag.\"\"\"\n    xs = [center_x(ds) for ds in headers if hasattr(ds, 'ImagePositionPatient')]\n    xm = float(np.mean(xs)) if xs else 0.0\n    tags = {str(getattr(ds, k, '') or '').upper()[:1]\n            for ds in headers for k in ('ImageLaterality','Laterality')} - {''}\n    tag = tags.pop() if len(tags) == 1 else None\n    if abs(xm) < DEAD_ZONE_MM and tag in ('L','R'):\n        return tag\n    return 'L' if xm > 0 else 'R'\n\ndef crop_background(vol_u8):\n    \"\"\"TODO: SUSTITUIR por la función exacta del prep de train.\n    Placeholder: bbox del tejido por umbral, cuadrado, margen 4 px.\"\"\"\n    m = vol_u8.max(axis=0) > 8\n    ys, xs = np.where(m)\n    if len(ys) == 0:\n        return vol_u8\n    y0, y1, x0, x1 = ys.min(), ys.max()+1, xs.min(), xs.max()+1\n    y0, x0 = max(0, y0-4), max(0, x0-4)\n    y1, x1 = min(vol_u8.shape[1], y1+4), min(vol_u8.shape[2], x1+4)\n    return vol_u8[:, y0:y1, x0:x1]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T21:38:14.230415Z","iopub.execute_input":"2026-08-08T21:38:14.230871Z","iopub.status.idle":"2026-08-08T21:38:14.2418Z","shell.execute_reply.started":"2026-08-08T21:38:14.230797Z","shell.execute_reply":"2026-08-08T21:38:14.240875Z"}},"outputs":[],"execution_count":null},{"id":"968ef622-a880-4335-af9d-bbaf3cea5abe","cell_type":"code","source":"def load_series(series_dir, laterality):\n    \"\"\"→ (x, plane_id): x = (K, 3, IMG, IMG) float32 normalizado ImageNet.\n    Canonicalización a 'rodilla derecha': hflip en coronal/axial izquierdas,\n    orden sagital con corte 0 = más medial.\"\"\"\n    files = sorted(glob.glob(f'{series_dir}/*.dcm'))\n    if not files:\n        return None, None\n    heads = [pydicom.dcmread(f, stop_before_pixels=True) for f in files]\n    iop = next(h.ImageOrientationPatient for h in heads if hasattr(h, 'ImageOrientationPatient'))\n    plane, row, col = plane_from_iop(iop)\n\n    # orden espacial: proyección de IPP sobre la normal (no fiarse de InstanceNumber)\n    nrm = np.cross(row, col)\n    pos = np.array([np.dot(np.array(h.ImagePositionPatient, float), nrm) for h in heads])\n    order = np.argsort(pos)\n    if plane == 'sagittal':\n        # normal sagital ≈ eje x (+x = izquierda del paciente). Corte 0 = más medial:\n        # rodilla derecha → x máximo primero; izquierda → x mínimo primero.\n        # TODO: verificar signo contra el prep de train (n=142 verificado allí).\n        xs = np.array([center_x(h) for h in heads])\n        order = np.argsort(xs)[::-1] if laterality == 'R' else np.argsort(xs)\n\n    sel = np.linspace(0, len(files) - 1, K).round().astype(int)   # K equiespaciados\n    idx = [order[i] for i in sel]\n    vol = np.stack([pydicom.dcmread(files[i]).pixel_array for i in idx]).astype(np.float32)\n\n    # normalización por percentiles sobre la serie (aquí: sobre los K muestreados,\n    # aproximación suficiente de la serie completa)\n    lo, hi = np.percentile(vol, [0.5, 99.5])\n    vol = np.clip((vol - lo) / max(hi - lo, 1e-6), 0, 1)\n    vol = (vol * 255).astype(np.uint8)\n    vol = crop_background(vol)\n\n    if plane in ('coronal', 'axial') and laterality == 'L':\n        vol = vol[:, :, ::-1]                                     # espejo a 'derecha'\n\n    out = np.empty((K, IMG, IMG), np.float32)\n    for j in range(K):\n        out[j] = cv2.resize(vol[j], (IMG, IMG), interpolation=cv2.INTER_AREA)\n    out /= 255.0\n    out = out[:, None].repeat(3, axis=1)\n    out = (out - IM_MEAN[None,:,None,None]) / IM_STD[None,:,None,None]\n    return torch.from_numpy(np.ascontiguousarray(out)), PLANES[plane]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T21:38:17.279042Z","iopub.execute_input":"2026-08-08T21:38:17.279434Z","iopub.status.idle":"2026-08-08T21:38:17.289622Z","shell.execute_reply.started":"2026-08-08T21:38:17.279393Z","shell.execute_reply":"2026-08-08T21:38:17.288595Z"}},"outputs":[],"execution_count":null},{"id":"97c77ae1-b912-4308-9ca3-929e9956dca3","cell_type":"code","source":"# --- cabeza: copia EXACTA de train_head.py + carga de los 5 folds ------------\nimport torch.nn as nn\n\nclass Head(nn.Module):\n    def __init__(self, d_in, d=256, K=32, n_cls=12, p_drop=0.3):\n        super().__init__()\n        self.proj = nn.Sequential(nn.LayerNorm(d_in), nn.Linear(d_in, d))\n        self.plane_emb = nn.Embedding(3, d)\n        self.register_buffer('pos', torch.linspace(0, 1, K).view(1, 1, K, 1))\n        self.pos_proj = nn.Linear(1, d)\n        self.att = nn.Sequential(nn.Linear(d, d // 2), nn.Tanh(), nn.Linear(d // 2, 1))\n        self.mlp = nn.Sequential(nn.LayerNorm(3 * d), nn.Dropout(p_drop),\n                                 nn.Linear(3 * d, d), nn.GELU(),\n                                 nn.Dropout(p_drop), nn.Linear(d, n_cls))\n    def forward(self, x, pl, mask):\n        h = self.proj(x) + self.pos_proj(self.pos) + self.plane_emb(pl)[:, :, None]\n        a = self.att(h).softmax(dim=2)\n        sv = (a * h).sum(dim=2)\n        outs = []\n        for p in range(3):\n            sel = (pl == p) & mask\n            n = sel.sum(1, keepdim=True).clamp(min=1)\n            outs.append((sv * sel[..., None]).sum(1) / n)\n        return self.mlp(torch.cat(outs, dim=1))\n\nheads = []\nfor f in range(5):\n    h = Head(D_PROBE, K=K).cuda().eval()\n    h.load_state_dict(torch.load(f'{D_HEADS}/head_f{f}.pt', map_location='cuda'))\n    heads.append(h)\nprint(len(heads), 'cabezas cargadas')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T21:38:19.798939Z","iopub.execute_input":"2026-08-08T21:38:19.79934Z","iopub.status.idle":"2026-08-08T21:38:19.991433Z","shell.execute_reply.started":"2026-08-08T21:38:19.799312Z","shell.execute_reply":"2026-08-08T21:38:19.990792Z"}},"outputs":[],"execution_count":null},{"id":"68fa8cf0-9bf3-449a-8a63-afb6e1b1def9","cell_type":"code","source":"# --- inferencia --------------------------------------------------------------\n@torch.no_grad()\ndef predict_study(study_dir):\n    series_dirs = sorted(glob.glob(f'{study_dir}/*'))\n    # pasada de lateralidad con el primer header de cada serie\n    first_heads = []\n    for sd in series_dirs:\n        f = glob.glob(f'{sd}/*.dcm')\n        if f:\n            first_heads.append(pydicom.dcmread(f[0], stop_before_pixels=True))\n    lat = study_laterality(first_heads)\n\n    embs, planes = [], []\n    for sd in series_dirs:\n        x, pl = load_series(sd, lat)\n        if x is None:\n            continue\n        with torch.autocast('cuda', dtype=torch.float16):\n            feats = model.forward_features(x.cuda())            # (K, 5+256, 768)\n        cls, pool = feats[:, 0], feats[:, NPREFIX:].mean(1)\n        e = cls if FEAT == 'cls' else (pool if FEAT == 'pool' else torch.cat([cls, pool], 1))\n        embs.append(e.float()); planes.append(pl)\n    if not embs:\n        raise RuntimeError('estudio sin series legibles')\n\n    x = torch.stack(embs)[None]                                  # (1, S, K, D)\n    pl = torch.tensor(planes, device='cuda')[None]\n    m = torch.ones_like(pl, dtype=torch.bool)\n    probs = torch.stack([torch.sigmoid(h(x, pl, m)) for h in heads]).mean(0)\n    return probs[0].cpu().numpy()\n\nrows, fails = [], []\nfor i, st in enumerate(test_studies):\n    try:\n        p = predict_study(f'{TEST_DIR}/{st}')\n    except Exception as e:\n        fails.append((st, repr(e)))\n        p = np.array([PREV[c] for c in CLASSES])                 # fallback: prevalencias\n    rows.append([st, *p])\n    if i % 100 == 0:\n        el = time.time() - T0\n        eta = el / max(i, 1) * len(test_studies)\n        print(f'{i:>5}/{len(test_studies)}  {el/60:6.1f} min  ETA total {eta/60:6.1f} min')\n\nprint(f'\\nfallos: {len(fails)}')\nfor st, e in fails[:10]:\n    print(' ', st, e)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T21:38:22.897653Z","iopub.execute_input":"2026-08-08T21:38:22.898601Z","iopub.status.idle":"2026-08-08T21:38:36.800118Z","shell.execute_reply.started":"2026-08-08T21:38:22.898565Z","shell.execute_reply":"2026-08-08T21:38:36.799315Z"}},"outputs":[],"execution_count":null},{"id":"925ed0b3-29ea-4326-958c-d051679302ed","cell_type":"code","source":"# --- submission --------------------------------------------------------------\nsub = pd.DataFrame(rows, columns=['StudyInstanceUID', *CLASSES])\nsample = pd.read_csv(f'{COMP}/sample_submission.csv')\n# contrato de formato: mismas columnas, mismo orden de filas que el sample\nsub = sample[['StudyInstanceUID']].merge(sub, on='StudyInstanceUID', how='left')\nassert list(sub.columns) == list(sample.columns), 'columnas no coinciden con sample_submission'\nassert sub[CLASSES].notna().all().all(), 'hay estudios sin predicción'\nassert ((sub[CLASSES] >= 0) & (sub[CLASSES] <= 1)).all().all()\nsub.to_csv('submission.csv', index=False)\nprint(sub.shape, f'| runtime total {(time.time()-T0)/60:.1f} min de {BUDGET_S/60:.0f}')\nsub.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T21:39:00.475207Z","iopub.execute_input":"2026-08-08T21:39:00.475617Z","iopub.status.idle":"2026-08-08T21:39:00.546037Z","shell.execute_reply.started":"2026-08-08T21:39:00.475588Z","shell.execute_reply":"2026-08-08T21:39:00.545261Z"}},"outputs":[],"execution_count":null}]}