{"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":[{"cell_type":"markdown","source":"## 1. Setup & Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport re\nimport json\nimport warnings\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nfrom sklearn.metrics import roc_auc_score\nfrom tqdm.auto import tqdm\n\nwarnings.filterwarnings('ignore')\n\n# ── Paths ─────────────────────────────────────────────────────────────────────\nDATA_DIR   = Path('/kaggle/input/competitions/rsna-knee-abnormality-detection')\nOUTPUT_DIR = Path('/kaggle/working')\nOUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n\n# ── Auto-detect Gemma model directory (handles Kaggle's nested layout) ────────\n_MODEL_ROOT = Path('/kaggle/input/models/google/gemma-3/transformers/gemma-3-4b-it/1')\n\n   # root added by Kaggle hub\n_config_files = sorted(_MODEL_ROOT.rglob('config.json'))\nif not _config_files:\n    print(\"ERROR: No config.json found. Listing tree:\")\n    for p in sorted(_MODEL_ROOT.rglob('*'))[:30]:\n        print(' ', p)\n    raise FileNotFoundError(\"Attach the Gemma-3 model: Add Data → Models → gemma-3-4b-it\")\n\nMODEL_PATH = _config_files[0].parent        # ← the real weights directory\nprint(f\"✅ Model resolved: {MODEL_PATH}\")\n\n# ── Label columns ──────────────────────────────────────────────────────────────\nLABEL_COLS = [\n    'ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus',\n    'Medial OA', 'Lateral OA', 'PF OA',\n    'Effusion', \"Synovitis\", \"Baker's\", 'Contusion', 'Fracture',\n]\n\nprint('Data dir:', DATA_DIR)\nprint('Files:', [f.name for f in DATA_DIR.iterdir() if f.is_file()])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T08:07:01.631794Z","iopub.execute_input":"2026-08-25T08:07:01.632177Z","iopub.status.idle":"2026-08-25T08:07:01.649341Z","shell.execute_reply.started":"2026-08-25T08:07:01.632151Z","shell.execute_reply":"2026-08-25T08:07:01.648341Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Load Training Data","metadata":{}},{"cell_type":"code","source":"# Read train.csv — handles embedded newlines in Report field\ntrain_df = pd.read_csv(\n    DATA_DIR / 'train.csv',\n    quoting=1,   # QUOTE_ALL\n    engine='python',\n    on_bad_lines='skip',\n)\nprint(f'Train studies: {len(train_df)}')\nprint(train_df.head(2))\n\n# Separate labeled vs. report-only studies\nhas_labels = train_df[LABEL_COLS].notna().any(axis=1)\nlabeled_df   = train_df[has_labels].copy()\nunlabeled_df = train_df[~has_labels].copy()\nprint(f'\\nStudies with gold labels: {len(labeled_df)}')\nprint(f'Studies with report only: {len(unlabeled_df)}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T08:07:01.651174Z","iopub.execute_input":"2026-08-25T08:07:01.651447Z","iopub.status.idle":"2026-08-25T08:07:01.852962Z","shell.execute_reply.started":"2026-08-25T08:07:01.651424Z","shell.execute_reply":"2026-08-25T08:07:01.852158Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Fast Keyword Extraction (Pass 1)","metadata":{}},{"cell_type":"code","source":"# ── Keyword lists per label (English terms)\n# The LLM will handle other languages; keywords give a quick first pass.\nKEYWORDS_POS = {\n    'ACL':              ['acl tear', 'acl rupture', 'anterior cruciate', 'acl injury', 'acl lesion', 'anterior cruciate ligament tear'],\n    'MCL':              ['mcl tear', 'mcl injury', 'medial collateral', 'mcl rupture', 'medial collateral ligament tear'],\n    'Medial Meniscus':  ['medial meniscus tear', 'medial meniscal tear', 'medial meniscus lesion', 'medial meniscus injury'],\n    'Lateral Meniscus': ['lateral meniscus tear', 'lateral meniscal tear', 'lateral meniscus lesion'],\n    'Medial OA':        ['medial osteoarthritis', 'medial compartment oa', 'medial tibiofemoral oa', 'medial joint space narrowing', 'medial chondral'],\n    'Lateral OA':       ['lateral osteoarthritis', 'lateral compartment oa', 'lateral tibiofemoral oa', 'lateral joint space narrowing'],\n    'PF OA':            ['patellofemoral osteoarthritis', 'pf oa', 'patellofemoral oa', 'patellofemoral chondromalacia'],\n    'Effusion':         ['joint effusion', 'effusion', 'fluid', 'intra-articular fluid'],\n    'Synovitis':        ['synovitis', 'synovial', 'synovial thickening'],\n    \"Baker's\":          [\"baker's cyst\", 'popliteal cyst', 'baker cyst'],\n    'Contusion':        ['bone contusion', 'bone bruise', 'contusion', 'trabecular edema', 'bone marrow edema'],\n    'Fracture':         ['fracture', 'fract', 'avulsion'],\n}\n\nKEYWORDS_NEG = {\n    'ACL':              ['acl intact', 'acl normal', 'no acl tear', 'acl unremarkable'],\n    'MCL':              ['mcl intact', 'no mcl tear', 'mcl normal'],\n    'Medial Meniscus':  ['medial meniscus intact', 'no medial meniscus tear'],\n    'Lateral Meniscus': ['lateral meniscus intact', 'no lateral meniscus tear'],\n    'Fracture':         ['no fracture', 'no fx', 'fracture excluded'],\n    'Effusion':         ['no effusion', 'no joint fluid'],\n    \"Baker's\":          [\"no baker's cyst\", 'no popliteal cyst'],\n}\n\ndef keyword_extract(report: str) -> dict:\n    \"\"\"Return {label: 0/1/NaN} via keyword matching.\"\"\"\n    if not isinstance(report, str) or len(report.strip()) == 0:\n        return {col: np.nan for col in LABEL_COLS}\n\n    text = report.lower()\n    result = {}\n\n    for col in LABEL_COLS:\n        # Check negative first (higher precision)\n        neg_hits = [kw for kw in KEYWORDS_NEG.get(col, []) if kw in text]\n        pos_hits = [kw for kw in KEYWORDS_POS.get(col, []) if kw in text]\n\n        if neg_hits and not pos_hits:\n            result[col] = 0\n        elif pos_hits:\n            result[col] = 1\n        else:\n            result[col] = np.nan  # uncertain → send to LLM\n\n    return result\n\n\n# Apply keyword extraction\ntqdm.pandas(desc='Keyword extraction')\nkw_rows = train_df['Report'].progress_apply(keyword_extract).tolist()\nkw_df = pd.DataFrame(kw_rows, index=train_df['StudyInstanceUID'])\n\n# Fraction confident\nconfident = kw_df.notna().mean()\nprint('\\nFraction of studies with confident keyword labels per condition:')\nprint(confident.round(3).to_string())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T08:07:01.854017Z","iopub.execute_input":"2026-08-25T08:07:01.854303Z","iopub.status.idle":"2026-08-25T08:07:02.390967Z","shell.execute_reply.started":"2026-08-25T08:07:01.854273Z","shell.execute_reply":"2026-08-25T08:07:02.389869Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. LLM Structured Extraction (Pass 2)","metadata":{}},{"cell_type":"code","source":"import torch\nfrom transformers import AutoTokenizer, AutoModelForCausalLM\n\nmodel_path_str = str(MODEL_PATH.resolve())   # ← key fix\n\ntokenizer = AutoTokenizer.from_pretrained(\n    model_path_str,\n    local_files_only=True,    # ← prevents HF hub lookup\n)\nllm = AutoModelForCausalLM.from_pretrained(\n    model_path_str,\n    torch_dtype=torch.bfloat16,\n    device_map='auto',\n    local_files_only=True,    # ← prevents HF hub lookup\n)\nllm.eval()\nprint('LLM loaded successfully.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T08:07:02.393088Z","iopub.execute_input":"2026-08-25T08:07:02.393382Z","iopub.status.idle":"2026-08-25T08:07:31.377026Z","shell.execute_reply.started":"2026-08-25T08:07:02.393358Z","shell.execute_reply":"2026-08-25T08:07:31.373067Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SYSTEM_PROMPT = \"\"\"You are a radiologist assistant that reads knee MRI reports.\nFor each report, extract binary presence (1) or absence (0) of exactly 12 findings.\nIf a finding is not mentioned, output 0.\nRespond ONLY with a valid JSON object with exactly these 12 keys:\nACL, MCL, \"Medial Meniscus\", \"Lateral Meniscus\",\n\"Medial OA\", \"Lateral OA\", \"PF OA\",\nEffusion, Synovitis, \"Baker's\", Contusion, Fracture\nEach value must be 0 or 1. No explanations, no text outside the JSON.\"\"\"\n\nFEW_SHOT = [\n    {\n        \"role\": \"user\",\n        \"content\": \"Report: Complete tear of the ACL. No meniscal tear identified. Small joint effusion present. No fracture.\"\n    },\n    {\n        \"role\": \"assistant\",\n        \"content\": '{\"ACL\": 1, \"MCL\": 0, \"Medial Meniscus\": 0, \"Lateral Meniscus\": 0, \"Medial OA\": 0, \"Lateral OA\": 0, \"PF OA\": 0, \"Effusion\": 1, \"Synovitis\": 0, \"Baker\\'s\": 0, \"Contusion\": 0, \"Fracture\": 0}'\n    },\n]\n\nMAX_REPORT_TOKENS = 800\n\ndef llm_extract(report: str, max_new_tokens: int = 150) -> dict:\n    \"\"\"Run LLM inference on a single report, return parsed JSON dict.\"\"\"\n    # Truncate long reports\n    tokens = tokenizer.encode(report, add_special_tokens=False)\n    if len(tokens) > MAX_REPORT_TOKENS:\n        tokens = tokens[:MAX_REPORT_TOKENS]\n        report = tokenizer.decode(tokens, skip_special_tokens=True)\n\n    messages = [\n        {\"role\": \"system\", \"content\": SYSTEM_PROMPT},\n        *FEW_SHOT,\n        {\"role\": \"user\", \"content\": f\"Report: {report}\"},\n    ]\n\n    prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)\n    inputs = tokenizer(prompt, return_tensors='pt', truncation=True, max_length=2048).to(device)\n\n    with torch.no_grad():\n        outputs = llm.generate(\n            **inputs,\n            max_new_tokens=max_new_tokens,\n            do_sample=False,\n            temperature=None,\n            top_p=None,\n            pad_token_id=tokenizer.eos_token_id,\n        )\n\n    generated = outputs[0][inputs['input_ids'].shape[1]:]\n    response = tokenizer.decode(generated, skip_special_tokens=True).strip()\n\n    # Parse JSON\n    try:\n        # Find JSON in response\n        json_match = re.search(r'\\{[^}]+\\}', response, re.DOTALL)\n        if json_match:\n            data = json.loads(json_match.group())\n            # Validate and sanitize\n            result = {}\n            for col in LABEL_COLS:\n                val = data.get(col, 0)\n                result[col] = int(val) if val in (0, 1, '0', '1') else 0\n            return result\n    except Exception:\n        pass\n\n    # Fallback: return zeros\n    return {col: 0 for col in LABEL_COLS}\n\n\nprint('LLM extraction function defined.')\nprint('\\nTest on a sample report:')\nsample_report = train_df['Report'].dropna().iloc[0]\nprint('Report snippet:', sample_report[:200])\nsample_result = llm_extract(sample_report)\nprint('Extracted labels:', sample_result)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T08:07:31.383084Z","iopub.execute_input":"2026-08-25T08:07:31.383694Z","iopub.status.idle":"2026-08-25T08:07:31.559253Z","shell.execute_reply.started":"2026-08-25T08:07:31.383657Z","shell.execute_reply":"2026-08-25T08:07:31.557873Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Determine which studies need LLM (those with any NaN from keyword pass)\n# Run LLM on ALL non-gold studies (keyword pass used as initialization)\n\n# Studies without gold labels\nunlabeled_uids = unlabeled_df['StudyInstanceUID'].tolist()\nprint(f'Running LLM on {len(unlabeled_uids)} non-gold studies...')\n\nllm_results = []\nfor uid in tqdm(unlabeled_uids, desc='LLM extraction'):\n    row = train_df[train_df['StudyInstanceUID'] == uid]\n    report = row['Report'].values[0] if len(row) > 0 else ''\n\n    if not isinstance(report, str) or len(report.strip()) < 10:\n        llm_results.append({col: 0 for col in LABEL_COLS})\n        continue\n\n    extracted = llm_extract(report)\n    llm_results.append(extracted)\n\nllm_df = pd.DataFrame(llm_results, index=unlabeled_uids)\nprint('LLM extraction complete.')\nprint('Label prevalence from LLM:')\nprint(llm_df.mean().round(3).to_string())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T08:07:31.56027Z","iopub.status.idle":"2026-08-25T08:07:31.560749Z","shell.execute_reply.started":"2026-08-25T08:07:31.560526Z","shell.execute_reply":"2026-08-25T08:07:31.560555Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Validate Against Gold Labels","metadata":{}},{"cell_type":"code","source":"# Validate LLM quality against gold-standard labeled studies\ngold_uids = labeled_df['StudyInstanceUID'].tolist()\nprint(f'Validating on {len(gold_uids)} gold-labeled studies...')\n\ngold_llm_results = []\nfor uid in tqdm(gold_uids, desc='LLM validation pass'):\n    row = train_df[train_df['StudyInstanceUID'] == uid]\n    report = row['Report'].values[0] if len(row) > 0 else ''\n    extracted = llm_extract(report) if isinstance(report, str) else {col: 0 for col in LABEL_COLS}\n    gold_llm_results.append(extracted)\n\ngold_llm_df = pd.DataFrame(gold_llm_results, index=gold_uids)\n\n# Compare to actual gold labels\ngold_true = labeled_df.set_index('StudyInstanceUID')[LABEL_COLS]\n\nprint('\\nPer-label AUC (LLM vs. Gold Labels):')\naucs = {}\nfor col in LABEL_COLS:\n    t = gold_true[col].fillna(0).values\n    p = gold_llm_df[col].values\n    if len(np.unique(t)) > 1:\n        auc = roc_auc_score(t, p)\n        aucs[col] = auc\n        print(f'  {col:25s}: {auc:.4f}')\n    else:\n        aucs[col] = np.nan\n        print(f'  {col:25s}: N/A (no positives in gold set)')\n\nvalid_aucs = [v for v in aucs.values() if not np.isnan(v)]\nprint(f'\\nMacro-average AUC (LLM pseudo-labeler quality): {np.mean(valid_aucs):.4f}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T08:07:31.562338Z","iopub.status.idle":"2026-08-25T08:07:31.562967Z","shell.execute_reply.started":"2026-08-25T08:07:31.562583Z","shell.execute_reply":"2026-08-25T08:07:31.562611Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Build & Save Final Pseudo-Label CSV","metadata":{}},{"cell_type":"code","source":"# Merge: gold labels take priority, LLM labels for the rest\npseudo_rows = []\n\nfor _, row in train_df.iterrows():\n    uid = row['StudyInstanceUID']\n    entry = {'StudyInstanceUID': uid}\n\n    if uid in gold_true.index:\n        # Gold labels — use directly\n        for col in LABEL_COLS:\n            val = gold_true.loc[uid, col]\n            entry[col] = float(val) if not pd.isna(val) else 0.0\n        entry['label_source'] = 'gold'\n    elif uid in llm_df.index:\n        # LLM-derived pseudo-label\n        for col in LABEL_COLS:\n            entry[col] = float(llm_df.loc[uid, col])\n        entry['label_source'] = 'llm'\n    else:\n        # No report or couldn't process → neutral\n        for col in LABEL_COLS:\n            entry[col] = 0.0\n        entry['label_source'] = 'default'\n\n    pseudo_rows.append(entry)\n\npseudo_df = pd.DataFrame(pseudo_rows)\npseudo_df.to_csv(OUTPUT_DIR / 'train_pseudolabels.csv', index=False)\n\nprint(f'Pseudo-label CSV saved: {OUTPUT_DIR / \"train_pseudolabels.csv\"}')\nprint(f'\\nSource breakdown:')\nprint(pseudo_df['label_source'].value_counts().to_string())\nprint(f'\\nLabel prevalence in pseudo-labels:')\nprint(pseudo_df[LABEL_COLS].mean().round(3).to_string())\nprint(f'\\nShape: {pseudo_df.shape}')\npseudo_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T08:07:31.564727Z","iopub.status.idle":"2026-08-25T08:07:31.565155Z","shell.execute_reply.started":"2026-08-25T08:07:31.564945Z","shell.execute_reply":"2026-08-25T08:07:31.564969Z"}},"outputs":[],"execution_count":null}]}