{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"print(\"Starting notebook\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-18T13:00:16.751548Z","iopub.execute_input":"2026-08-18T13:00:16.752335Z","iopub.status.idle":"2026-08-18T13:00:16.756329Z","shell.execute_reply.started":"2026-08-18T13:00:16.752301Z","shell.execute_reply":"2026-08-18T13:00:16.755664Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Starting RSNA Knee Axial Supervised Baseline\")\n\n\"\"\"\nRSNA Knee Abnormality Detection\n--------------------------------\n\nSupervised axial-only baseline using weak labels extracted from\nthe LLM/radiology-report label dataset.\n\nIMPORTANT:\n- We intentionally use Chunk 1 for training and Chunk 2 for testing.\n- We do NOT care about patient/study leakage for this experiment.\n- Existing tensor caches are reused.\n- Only one image tensor chunk is kept in RAM at a time.\n- GPU memory is protected with AMP + moderate batch size.\n- Labels are study-level and therefore repeated for every axial slice.\n- Study-level prediction is obtained by top-k slice pooling.\n\nTargets:\n    ACL\n    MCL\n    Medial Meniscus\n    Lateral Meniscus\n    Medial OA\n    Lateral OA\n    PF OA\n    Effusion\n    Synovitis\n    Baker's\n    Contusion\n    Fracture\n\nLLM labels:\n    probabilities are converted to binary using > 0.5\n\nOutputs:\n    /kaggle/working/knee_classifier_weights.pt\n    /kaggle/working/knee_classifier_results.json\n\"\"\"\n\nimport os\nimport glob\nimport json\nimport time\nimport gc\nimport re\nimport warnings\n\nfrom typing import Dict, List, Tuple, Optional\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport pydicom\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\n\nfrom sklearn.metrics import (\n    roc_auc_score,\n    average_precision_score,\n    f1_score,\n    accuracy_score,\n    precision_score,\n    recall_score,\n    confusion_matrix,\n)\n\nwarnings.filterwarnings(\"ignore\")\n\n\n# =====================================================================\n# 1. CONFIGURATION\n# =====================================================================\n\nCOMPETITION_PATH = (\n    \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n)\n\nLLM_DIR = (\n    \"/kaggle/input/datasets/stevenleehans/rsna-knee-llm-report-labels\"\n)\n\nAXIAL_PATHS_CACHE = (\n    \"/kaggle/working/axial_paths_cache.json\"\n)\n\nCHUNK_FILE_PATTERN = (\n    \"/kaggle/working/axial_chunk_{}.pt\"\n)\n\nLABEL_CHUNK_FILE_PATTERN = (\n    \"/kaggle/working/axial_labels_chunk_{}.pt\"\n)\n\nMODEL_WEIGHTS_FILE = (\n    \"/kaggle/working/knee_classifier_weights.pt\"\n)\n\nRESULTS_FILE = (\n    \"/kaggle/working/knee_classifier_results.json\"\n)\n\nLABEL_MAPPING_FILE = (\n    \"/kaggle/working/knee_label_mapping.json\"\n)\n\nTARGET_RESOLUTION = (128, 128)\n\nTARGET_COLUMNS = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\",\n]\n\n# ---------------------------------------------------------------------\n# Training configuration\n# ---------------------------------------------------------------------\n\nEPOCHS = 8\n\nBATCH_SIZE = 128\n\nLEARNING_RATE = 3e-4\n\nWEIGHT_DECAY = 1e-4\n\nNUM_WORKERS = 2\n\nPIN_MEMORY = True\n\n# Number of slices used for study-level pooling.\n# top-k=5 means the 5 strongest abnormal-looking slices are averaged.\nTOP_K = 5\n\n# AMP significantly reduces T4 memory usage.\nUSE_AMP = True\n\nSEED = 42\n\n\n# =====================================================================\n# 2. REPRODUCIBILITY\n# =====================================================================\n\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\n\nDEVICE = torch.device(\n    \"cuda\" if torch.cuda.is_available() else \"cpu\"\n)\n\nprint(\"=\" * 70)\nprint(\"CONFIGURATION\")\nprint(\"=\" * 70)\nprint(f\"Device       : {DEVICE}\")\n\nif torch.cuda.is_available():\n    print(f\"GPU          : {torch.cuda.get_device_name(0)}\")\n    print(\n        f\"GPU memory   : \"\n        f\"{torch.cuda.get_device_properties(0).total_memory / 1024**3:.2f} GB\"\n    )\n\nprint(f\"Resolution   : {TARGET_RESOLUTION}\")\nprint(f\"Batch size   : {BATCH_SIZE}\")\nprint(f\"Epochs       : {EPOCHS}\")\nprint(f\"AMP          : {USE_AMP}\")\nprint(f\"Top-k pooling: {TOP_K}\")\nprint(\"=\" * 70)\n\n\n# =====================================================================\n# 3. FIND LLM LABEL FILE\n# =====================================================================\n\ndef find_label_files() -> List[str]:\n    \"\"\"\n    Find CSV/Parquet/JSON files in the supplied LLM label directory.\n    \"\"\"\n\n    patterns = [\n        \"*.csv\",\n        \"*.CSV\",\n        \"*.parquet\",\n        \"*.json\",\n        \"*.jsonl\",\n    ]\n\n    files = []\n\n    for pattern in patterns:\n        files.extend(\n            glob.glob(\n                os.path.join(LLM_DIR, \"**\", pattern),\n                recursive=True,\n            )\n        )\n\n    return sorted(set(files))\n\n\ndef inspect_label_files():\n    \"\"\"\n    Print available files and their columns.\n\n    This is intentionally done before guessing the schema.\n    \"\"\"\n\n    files = find_label_files()\n\n    print(\"\\n\" + \"=\" * 70)\n    print(\"LLM LABEL DATASET INSPECTION\")\n    print(\"=\" * 70)\n\n    if not files:\n        raise FileNotFoundError(\n            f\"No CSV/Parquet/JSON files found in:\\n{LLM_DIR}\"\n        )\n\n    for file_path in files:\n        print(f\"\\nFILE: {file_path}\")\n\n        try:\n            if file_path.endswith(\".csv\"):\n                df = pd.read_csv(file_path, nrows=5)\n\n            elif file_path.endswith(\".parquet\"):\n                df = pd.read_parquet(file_path).head()\n\n            elif file_path.endswith(\".json\"):\n                df = pd.read_json(file_path).head()\n\n            elif file_path.endswith(\".jsonl\"):\n                df = pd.read_json(\n                    file_path,\n                    lines=True,\n                    nrows=5,\n                )\n\n            else:\n                continue\n\n            print(\"Columns:\")\n            print(list(df.columns))\n\n        except Exception as e:\n            print(f\"Could not inspect: {e}\")\n\n    print(\"=\" * 70)\n\n    return files\n\n\n# =====================================================================\n# 4. NORMALIZE COLUMN NAMES\n# =====================================================================\n\ndef normalize_column_name(x: str) -> str:\n    \"\"\"\n    Convert column names into a comparable representation.\n    \"\"\"\n\n    x = str(x).lower().strip()\n\n    x = x.replace(\"_\", \" \")\n    x = x.replace(\"-\", \" \")\n    x = re.sub(r\"\\s+\", \" \", x)\n\n    return x\n\n\ndef find_column(\n    columns,\n    candidates: List[str],\n) -> Optional[str]:\n    \"\"\"\n    Find a column using normalized exact matching first,\n    then substring matching.\n    \"\"\"\n\n    normalized = {\n        normalize_column_name(c): c\n        for c in columns\n    }\n\n    # Exact match\n    for candidate in candidates:\n\n        candidate_norm = normalize_column_name(candidate)\n\n        if candidate_norm in normalized:\n            return normalized[candidate_norm]\n\n    # Substring match\n    for candidate in candidates:\n\n        candidate_norm = normalize_column_name(candidate)\n\n        for norm_col, original_col in normalized.items():\n\n            if candidate_norm in norm_col:\n                return original_col\n\n    return None\n\n\n# =====================================================================\n# 5. LOAD LLM LABELS\n# =====================================================================\n\ndef load_label_dataframe() -> pd.DataFrame:\n    \"\"\"\n    Load the most likely label table.\n\n    We prefer a table containing:\n        - StudyInstanceUID\n        - multiple target columns\n    \"\"\"\n\n    files = find_label_files()\n\n    candidates = []\n\n    for file_path in files:\n\n        try:\n\n            if file_path.endswith(\".csv\"):\n                df = pd.read_csv(file_path)\n\n            elif file_path.endswith(\".parquet\"):\n                df = pd.read_parquet(file_path)\n\n            elif file_path.endswith(\".json\"):\n                df = pd.read_json(file_path)\n\n            elif file_path.endswith(\".jsonl\"):\n                df = pd.read_json(\n                    file_path,\n                    lines=True,\n                )\n\n            else:\n                continue\n\n            columns = list(df.columns)\n\n            uid_col = find_column(\n                columns,\n                [\n                    \"StudyInstanceUID\",\n                    \"Study Instance UID\",\n                    \"study_uid\",\n                    \"studyinstanceuid\",\n                    \"study id\",\n                ],\n            )\n\n            target_matches = 0\n\n            for target in TARGET_COLUMNS:\n\n                aliases = [\n                    target,\n                    target.lower(),\n                    target.replace(\" \", \"_\"),\n                ]\n\n                if target == \"Baker's\":\n                    aliases += [\n                        \"Baker cyst\",\n                        \"Bakers cyst\",\n                        \"Baker's cyst\",\n                        \"Baker\",\n                    ]\n\n                if target == \"Medial Meniscus\":\n                    aliases += [\n                        \"medial meniscus tear\",\n                        \"medial meniscus injury\",\n                    ]\n\n                if target == \"Lateral Meniscus\":\n                    aliases += [\n                        \"lateral meniscus tear\",\n                        \"lateral meniscus injury\",\n                    ]\n\n                if target == \"Medial OA\":\n                    aliases += [\n                        \"medial osteoarthritis\",\n                        \"medial oa\",\n                    ]\n\n                if target == \"Lateral OA\":\n                    aliases += [\n                        \"lateral osteoarthritis\",\n                        \"lateral oa\",\n                    ]\n\n                if target == \"PF OA\":\n                    aliases += [\n                        \"patellofemoral oa\",\n                        \"patellofemoral osteoarthritis\",\n                        \"pf osteoarthritis\",\n                    ]\n\n                if find_column(columns, aliases):\n                    target_matches += 1\n\n            if uid_col is not None and target_matches >= 2:\n\n                candidates.append(\n                    (\n                        target_matches,\n                        len(df),\n                        file_path,\n                        df,\n                        uid_col,\n                    )\n                )\n\n        except Exception as e:\n\n            print(\n                f\"Could not read {file_path}: {e}\"\n            )\n\n    if not candidates:\n\n        raise RuntimeError(\n            \"\\nCould not automatically identify the LLM label table.\\n\"\n            \"Run the notebook once and inspect the printed columns above.\"\n        )\n\n    candidates.sort(\n        key=lambda x: (x[0], x[1]),\n        reverse=True,\n    )\n\n    best = candidates[0]\n\n    print(\"\\n\" + \"=\" * 70)\n    print(\"SELECTED LLM LABEL FILE\")\n    print(\"=\" * 70)\n\n    print(f\"File        : {best[2]}\")\n    print(f\"Rows        : {len(best[3]):,}\")\n    print(f\"UID column  : {best[4]}\")\n    print(f\"Target cols : {best[0]}\")\n    print(\"=\" * 70)\n\n    return best[3]\n\n\ndef build_label_lookup(\n    df: pd.DataFrame,\n) -> Tuple[Dict[str, np.ndarray], Dict[str, str]]:\n    \"\"\"\n    Build:\n\n        StudyInstanceUID -> 12-element binary vector\n\n    Values > 0.5 become 1.\n\n    Missing labels become 0.\n    \"\"\"\n\n    columns = list(df.columns)\n\n    uid_col = find_column(\n        columns,\n        [\n            \"StudyInstanceUID\",\n            \"Study Instance UID\",\n            \"study_uid\",\n            \"studyinstanceuid\",\n            \"study id\",\n        ],\n    )\n\n    if uid_col is None:\n        raise RuntimeError(\n            \"Could not find StudyInstanceUID in label dataframe.\"\n        )\n\n    selected_columns = {}\n\n    for target in TARGET_COLUMNS:\n\n        aliases = [\n            target,\n            target.lower(),\n            target.replace(\" \", \"_\"),\n        ]\n\n        if target == \"Baker's\":\n            aliases += [\n                \"Baker cyst\",\n                \"Bakers cyst\",\n                \"Baker's cyst\",\n                \"Baker\",\n            ]\n\n        if target == \"Medial Meniscus\":\n            aliases += [\n                \"medial meniscus tear\",\n                \"medial meniscus injury\",\n            ]\n\n        if target == \"Lateral Meniscus\":\n            aliases += [\n                \"lateral meniscus tear\",\n                \"lateral meniscus injury\",\n            ]\n\n        if target == \"Medial OA\":\n            aliases += [\n                \"medial osteoarthritis\",\n            ]\n\n        if target == \"Lateral OA\":\n            aliases += [\n                \"lateral osteoarthritis\",\n            ]\n\n        if target == \"PF OA\":\n            aliases += [\n                \"patellofemoral oa\",\n                \"patellofemoral osteoarthritis\",\n            ]\n\n        col = find_column(columns, aliases)\n\n        selected_columns[target] = col\n\n    print(\"\\nLABEL COLUMN MAPPING\")\n    print(\"-\" * 70)\n\n    for target, col in selected_columns.items():\n        print(\n            f\"{target:20s} -> \"\n            f\"{str(col)}\"\n        )\n\n    print(\"-\" * 70)\n\n    # -------------------------------------------------------------\n    # Build lookup\n    # -------------------------------------------------------------\n\n    lookup = {}\n\n    for _, row in df.iterrows():\n\n        uid = str(row[uid_col])\n\n        if uid == \"nan\":\n            continue\n\n        labels = []\n\n        for target in TARGET_COLUMNS:\n\n            col = selected_columns[target]\n\n            if col is None:\n\n                # Missing target column.\n                # We use 0, but report it clearly.\n                value = 0.0\n\n            else:\n\n                value = pd.to_numeric(\n                    row[col],\n                    errors=\"coerce\",\n                )\n\n                if pd.isna(value):\n                    value = 0.0\n\n            # Requested threshold:\n            # probability > .5 -> 1\n            label = 1.0 if float(value) > 0.5 else 0.0\n\n            labels.append(label)\n\n        lookup[uid] = np.asarray(\n            labels,\n            dtype=np.float32,\n        )\n\n    print(\n        f\"Built label lookup for \"\n        f\"{len(lookup):,} studies.\"\n    )\n\n    # Save mapping for reproducibility\n    with open(\n        LABEL_MAPPING_FILE,\n        \"w\",\n    ) as f:\n\n        json.dump(\n            {\n                \"uid_column\": uid_col,\n                \"columns\": selected_columns,\n                \"threshold\": 0.5,\n                \"targets\": TARGET_COLUMNS,\n            },\n            f,\n            indent=2,\n        )\n\n    return lookup, selected_columns\n\n\n# =====================================================================\n# 6. AXIAL PATH EXTRACTION\n# =====================================================================\n\ndef get_axial_paths_from_csv() -> List[str]:\n    \"\"\"\n    Load cached axial paths if available.\n\n    Otherwise construct them from train_series.csv.\n    \"\"\"\n\n    if os.path.exists(AXIAL_PATHS_CACHE):\n\n        print(\n            f\"Loading cached paths:\\n\"\n            f\"{AXIAL_PATHS_CACHE}\"\n        )\n\n        with open(\n            AXIAL_PATHS_CACHE,\n            \"r\",\n        ) as f:\n\n            return json.load(f)\n\n    print(\n        \"Axial path cache does not exist.\"\n    )\n\n    train_series_csv = (\n        f\"{COMPETITION_PATH}/train_series.csv\"\n    )\n\n    df_series = pd.read_csv(\n        train_series_csv\n    )\n\n    axial_df = df_series[\n        df_series[\"Anatomical_Plane\"]\n        .astype(str)\n        .str.lower()\n        .eq(\"axial\")\n    ]\n\n    print(\n        f\"Total series : {len(df_series):,}\"\n    )\n\n    print(\n        f\"Axial series : {len(axial_df):,}\"\n    )\n\n    axial_paths = []\n\n    for _, row in axial_df.iterrows():\n\n        study_uid = row[\"StudyInstanceUID\"]\n        series_uid = row[\"SeriesInstanceUID\"]\n\n        series_dir = (\n            f\"{COMPETITION_PATH}/train_series/\"\n            f\"{study_uid}/{series_uid}\"\n        )\n\n        axial_paths.extend(\n            glob.glob(\n                f\"{series_dir}/*.dcm\"\n            )\n        )\n\n    with open(\n        AXIAL_PATHS_CACHE,\n        \"w\",\n    ) as f:\n\n        json.dump(\n            axial_paths,\n            f,\n        )\n\n    print(\n        f\"Found {len(axial_paths):,} axial DICOMs.\"\n    )\n\n    return axial_paths\n\n\n# =====================================================================\n# 7. STUDY UID FROM PATH\n# =====================================================================\n\ndef study_uid_from_path(path: str) -> str:\n    \"\"\"\n    Expected:\n\n    .../train_series/<StudyUID>/<SeriesUID>/<file>.dcm\n\n    Therefore StudyUID is parent-parent directory.\n    \"\"\"\n\n    return os.path.basename(\n        os.path.dirname(\n            os.path.dirname(path)\n        )\n    )\n\n\n# =====================================================================\n# 8. CREATE LABEL TENSOR WITHOUT READING DICOM\n# =====================================================================\n\ndef build_chunk_labels(\n    paths: List[str],\n    chunk_id: int,\n    label_lookup: Dict[str, np.ndarray],\n) -> torch.Tensor:\n    \"\"\"\n    Generate labels from paths.\n\n    This does NOT read the DICOM files.\n    \"\"\"\n\n    label_cache = (\n        LABEL_CHUNK_FILE_PATTERN.format(\n            chunk_id\n        )\n    )\n\n    if os.path.exists(label_cache):\n\n        print(\n            f\"Loading cached labels:\"\n            f\" {label_cache}\"\n        )\n\n        return torch.load(\n            label_cache,\n            map_location=\"cpu\",\n        )\n\n    print(\n        f\"Building labels for Chunk \"\n        f\"{chunk_id}...\"\n    )\n\n    labels = np.zeros(\n        (\n            len(paths),\n            len(TARGET_COLUMNS),\n        ),\n        dtype=np.float32,\n    )\n\n    missing = 0\n\n    for i, path in enumerate(paths):\n\n        uid = study_uid_from_path(path)\n\n        if uid in label_lookup:\n\n            labels[i] = label_lookup[uid]\n\n        else:\n\n            missing += 1\n\n    tensor = torch.from_numpy(labels)\n\n    torch.save(\n        tensor,\n        label_cache,\n    )\n\n    print(\n        f\"Saved labels: {label_cache}\"\n    )\n\n    print(\n        f\"Missing study labels: \"\n        f\"{missing:,}/{len(paths):,}\"\n    )\n\n    return tensor\n\n\n# =====================================================================\n# 9. LOAD IMAGE CHUNK\n# =====================================================================\n\ndef load_chunk_tensor(\n    paths: List[str],\n    chunk_id: int,\n) -> torch.Tensor:\n    \"\"\"\n    Load an image tensor from disk if cached.\n\n    Otherwise read DICOMs once and cache the resulting tensor.\n\n    WARNING:\n        A 58k x 128 x 128 float32 tensor is approximately 3.8 GB.\n        This is intentional, but only ONE chunk should be resident.\n    \"\"\"\n\n    cache_pt = (\n        CHUNK_FILE_PATTERN.format(\n            chunk_id\n        )\n    )\n\n    if os.path.exists(cache_pt):\n\n        print(\n            f\"\\nLoading existing image tensor:\"\n            f\"\\n{cache_pt}\"\n        )\n\n        tensor = torch.load(\n            cache_pt,\n            map_location=\"cpu\",\n        )\n\n        print(\n            f\"Chunk {chunk_id} shape: \"\n            f\"{tuple(tensor.shape)}\"\n        )\n\n        return tensor\n\n    print(\n        f\"\\nReading DICOMs for Chunk \"\n        f\"{chunk_id}: {len(paths):,}\"\n    )\n\n    start = time.time()\n\n    slices = []\n\n    for idx, file_path in enumerate(paths):\n\n        if idx % 1000 == 0:\n\n            elapsed = time.time() - start\n\n            rate = (\n                idx / elapsed\n                if elapsed > 0\n                else 0\n            )\n\n            remaining = (\n                len(paths) - idx\n            )\n\n            eta = (\n                remaining / rate\n                if rate > 0\n                else 0\n            )\n\n            print(\n                f\"Chunk {chunk_id}: \"\n                f\"{idx:,}/{len(paths):,} | \"\n                f\"ETA {eta/60:.1f} min\"\n            )\n\n        try:\n\n            ds = pydicom.dcmread(\n                file_path,\n                force=True,\n            )\n\n            img = ds.pixel_array.astype(\n                np.float32\n            )\n\n            # Per-slice intensity normalization\n            mn = img.min()\n            mx = img.max()\n\n            if mx > mn:\n\n                img = (\n                    img - mn\n                ) ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}