""" Feature extraction — step 1 of the LLM integration ================================================== Runs Vbai-2.6AD FROZEN and writes everything the LLM will need into a single file. Because the encoder is frozen this is a one-off: every later projector or LoRA experiment reads this cache and never touches the 3D CNN again. Stored per visit: fused_features (512) the representation behind the 0.895 accuracy; the projector's main input mri_features (512) imaging branch tab_features (256) biomarker branch class_probs (3) CN/MCI/AD — the LLM's ANCHOR, also written as text will_progress (1) MCI -> AD risk score label, ptid, split, nifti_path bio_values (13) + bio_mask (13) raw values for the text template (mask = 0 means "not measured" and must be stated, never silently skipped) NOTE — no spatial tokens. token_probe.py showed the post-ASPP 3x3x3 grid is degenerate (attention entropy at 98.9% of maximum), i.e. the 27 tokens carry nothing beyond the pooled vector. The pooled representation itself is sound (fresh probe macro-F1 0.464, MCI F1 0.474). Run: python extract_features.py --tbm --ckpt Vbai-2.6AD.pt --out features.pt """ from __future__ import annotations import argparse import os import sys # --- The modality must be chosen BEFORE config is imported: config decides at # import time which visit manifest to read. _ap = argparse.ArgumentParser(add_help=False) _ap.add_argument("--tbm", action="store_true") _ap.add_argument("--t1", action="store_true") _known, _ = _ap.parse_known_args() if _known.tbm == _known.t1: sys.exit("ERROR: pass exactly one of --tbm / --t1 " "(match whichever modality the checkpoint was trained on).") os.environ["VBAI_USE_TBM"] = "1" if _known.tbm else "0" MODALITY = "TBM" if _known.tbm else "raw T1" def _bootstrap_model_path() -> str: """ Locate config.py / model.py / dataset.py. YOU MUST SET YOUR OWN PATH if they do not sit next to this script: point VBAI_MODEL_DIR at the directory holding them. """ env = os.environ.get("VBAI_MODEL_DIR") cands = ([env] if env else []) + [ os.path.dirname(os.path.abspath(__file__)), os.path.join(os.path.dirname(os.path.abspath(__file__)), "Vbai-2.6AD"), ] for c in cands: if c and os.path.isfile(os.path.join(c, "config.py")): if c not in sys.path: sys.path.insert(0, c) return c raise ImportError( f"Could not locate the model modules (config.py). Tried: {cands}\n" "YOU MUST SET YOUR OWN PATH: point VBAI_MODEL_DIR at the directory " "holding config.py / model.py / dataset.py." ) MODEL_DIR = _bootstrap_model_path() import numpy as np import torch from tqdm import tqdm import config as C from model import Vbai26ADModel from dataset import (PairedVisitDataset, TabularNormalizer, collate_pad, subject_split, load_paired) def remap_paths(df): """ Re-root the volume paths stored in the manifest. A manifest built on one machine carries that machine's absolute paths. Rather than forcing a rebuild, the tail of each path is re-attached to the roots configured here. """ def _fix(p): p0 = str(p) if os.path.exists(p0): return p0 q = p0.replace("\\", "/") i = q.find("/Datasets/") if i >= 0: cand = os.path.join(C.DATASET_ROOT, q[i + len("/Datasets/"):]) if os.path.exists(cand): return cand # Volume root: the manifest tail may or may not include the top folder, # so both spellings are tried. j = q.find("/volumes/") if j >= 0: rest = q[j + len("/volumes/"):] for cand in (os.path.join(C.TBM_ROOT, rest), os.path.join(C.TBM_ROOT, "volumes", rest)): if os.path.exists(cand): return cand return p0 df = df.copy() df["nifti_path"] = df["nifti_path"].map(_fix) ok = int(sum(os.path.exists(str(p)) for p in df["nifti_path"])) print(f"[path] reachable images ({MODALITY}): {ok}/{len(df)}") if ok == 0: raise FileNotFoundError( f"No image is reachable.\n DATASET_ROOT={C.DATASET_ROOT}\n" f" VOLUME_ROOT={C.TBM_ROOT}\n" "YOU MUST SET YOUR OWN PATHS: see VBAI_DATASET_ROOT / " "VBAI_VOLUME_ROOT in config.py." ) return df[df["nifti_path"].map(lambda p: os.path.exists(str(p)))].reset_index(drop=True) def main(): ap = argparse.ArgumentParser() ap.add_argument("--ckpt", default="Vbai-2.6AD.pt", help="Vbai-2.6AD checkpoint") ap.add_argument("--out", required=True, help="output .pt path") ap.add_argument("--tbm", action="store_true") ap.add_argument("--t1", action="store_true") ap.add_argument("--batch-size", type=int, default=4) ap.add_argument("--workers", type=int, default=2) args = ap.parse_args() dev = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"[device] {dev} | [modality] {MODALITY} | " f"[manifest] {os.path.basename(C.PAIRED_PARQUET)}") sd = torch.load(args.ckpt, map_location=dev, weights_only=False) mcfg = C.ModelConfig() for k, v in sd.get("model_cfg", {}).items(): if hasattr(mcfg, k): setattr(mcfg, k, v) model = Vbai26ADModel(mcfg).to(dev) res = model.load_state_dict(sd["model"], strict=False) # Loading with strict=False and ignoring the result is how a silently wrong # checkpoint slips through, so any key mismatch is fatal here. if res.missing_keys or res.unexpected_keys: raise RuntimeError( f"Checkpoint does not match the architecture: " f"{len(res.missing_keys)} missing / {len(res.unexpected_keys)} " f"unexpected keys. Wrong model file?") model.eval() for p in model.parameters(): p.requires_grad_(False) print(f"[ckpt] {len(sd['model'])} keys matched | {sd.get('extra', {}).get('metrics')}") norm = TabularNormalizer() norm.load_state_dict(sd["norm"]) feat_names = sd.get("feature_names", C.FEATURE_NAMES) df = remap_paths(load_paired()) train_ids, val_ids, test_ids = subject_split(df) split_of = {} for s, ids in (("train", train_ids), ("val", val_ids), ("test", test_ids)): for i in ids: split_of[i] = s df["split"] = df["ptid"].map(split_of) ds = PairedVisitDataset(df, norm, mode="multi", augment=False, mcfg=mcfg) dl = torch.utils.data.DataLoader(ds, batch_size=args.batch_size, shuffle=False, collate_fn=collate_pad, num_workers=args.workers) acc = {k: [] for k in ["fused_features", "mri_features", "tab_features", "class_probs", "will_progress", "label"]} n_seen = 0 with torch.no_grad(): for b in tqdm(dl, desc="extracting features"): if "mri" not in b or "tab" not in b: raise RuntimeError("Batch is missing mri or tab — " "modality dropout must be off here.") out = model(mri=b["mri"].to(dev), tab=b["tab"].to(dev)) acc["fused_features"].append(out["fused_features"].cpu()) acc["mri_features"].append(out["mri_features"].cpu()) acc["tab_features"].append(out["tab_features"].cpu()) acc["class_probs"].append(torch.softmax(out["fused_logits"], -1).cpu()) acc["will_progress"].append(out["progression"]["will_progress"].cpu()) acc["label"].append(b["label"]) n_seen += b["label"].size(0) store = {k: torch.cat(v).float() for k, v in acc.items()} store["label"] = store["label"].long() # DataLoader order matches df order (shuffle=False), so metadata lines up. assert n_seen == len(df), f"sample count mismatch: {n_seen} vs {len(df)}" store["ptid"] = df["ptid"].tolist() store["split"] = df["split"].tolist() store["nifti_path"] = df["nifti_path"].tolist() # RAW (un-normalised) biomarker values plus their mask, for the text template. vals = np.stack([np.where(np.isnan(df[f].values.astype(np.float32)), 0.0, df[f].values.astype(np.float32)) for f in feat_names], axis=1) mask = np.stack([df[f"feat_mask_{f}"].values.astype(np.float32) for f in feat_names], axis=1) store["bio_values"] = torch.from_numpy(vals) store["bio_mask"] = torch.from_numpy(mask) store["feature_names"] = list(feat_names) store["class_names"] = sd.get("class_names", C.CLASS_NAMES) store["modality"] = MODALITY store["ckpt"] = os.path.abspath(args.ckpt) os.makedirs(os.path.dirname(os.path.abspath(args.out)) or ".", exist_ok=True) torch.save(store, args.out) from collections import Counter print(f"\n[saved] {args.out}") print(f" visits : {len(df)} | patients: {df['ptid'].nunique()}") print(f" split : {Counter(store['split'])}") print(f" fused_features: {tuple(store['fused_features'].shape)}") print(f" bio_values : {tuple(store['bio_values'].shape)} (with mask)") print(" measured rate : " + ", ".join( f"{n}={store['bio_mask'][:, i].mean():.2f}" for i, n in enumerate(feat_names))) print("\nNext: train the projector on this cache (the LLM stays frozen).") if __name__ == "__main__": main()