| """ |
| Dataset F Viewer — FastAPI backend. |
| Serves images from the HF dataset with the EXACT same preprocessing the model sees: |
| 1. PIL convert("RGB") |
| 2. Letterbox to 384×384 (BICUBIC resize + zero-pad) |
| 3. ImageNet normalize: (x - mean) / std |
| |
| Three image modes: |
| - raw: Original 350×350 image as-is |
| - letterbox: After letterbox to 384×384 (before normalization) — what the model literally receives |
| - normalized: Letterbox + ImageNet normalize (de-normalized for display — what the model "sees") |
| |
| Also joins CystoDS metadata for source C images to show the 8 NML subclasses. |
| """ |
| import io, json, base64, pathlib, sys |
| from collections import Counter, defaultdict |
| import numpy as np |
| import pandas as pd |
| from PIL import Image |
| from fastapi import FastAPI, Query, HTTPException |
| from fastapi.responses import HTMLResponse, JSONResponse, Response |
| from datasets import load_from_disk |
|
|
| |
| F_HF = pathlib.Path("/home/kccitadmin/jupyterlab/sandbox/datasets/F_hf_dataset") |
| CYSTO_META = pathlib.Path("/home/kccitadmin/jupyterlab/sandbox/data/C/data/metadata.parquet") |
| VIEWER_DIR = pathlib.Path("/home/kccitadmin/jupyterlab/sandbox/tools/viewer") |
|
|
| |
| SIZE = 384 |
| IMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32) |
| IMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32) |
|
|
| |
| TARGET3_NAMES = {0: "Malignant Tumor", 1: "Non-Malignant Lesion", 2: "Non-ROI"} |
| CANCER_NAMES = {0: "Non-cancer", 1: "Cancer"} |
| GRADE_NAMES = {0: "High Grade", 1: "Low Grade", 2: "Non-cancer"} |
| SUBCLASS_NAMES = {0: "MT", 1: "NML", 2: "Normal", 3: "Foreign Body", 4: "Landmark"} |
| IMAGING_NAMES = {0: "WLI", 1: "BLC", 2: "NBI"} |
| SPLIT_NAMES = {"train": "Train", "validation": "Validation", "test": "Test"} |
|
|
| |
| CYSTO_NML_SUBCLASSES = { |
| "BenignNOS": "Benign NOS (reactive changes, atypia, dysplasia)", |
| "InflammationNOS": "Inflammation NOS (non-specific cystitis)", |
| "CCG": "Cystitis cystica & glandularis", |
| "Denuded": "Denuded urothelium", |
| "UrothelialPapilloma": "Urothelial papilloma", |
| "SquamousMetaplasia": "Squamous metaplasia", |
| "NephrogenicAdenoma": "Nephrogenic adenoma", |
| "BenignRare": "Benign rare (malakoplakia, melanosis)", |
| } |
|
|
| |
| print("Loading F dataset...", flush=True) |
| f_ds = load_from_disk(str(F_HF)) |
|
|
| |
| print("Building index...", flush=True) |
| INDEX = {} |
| for split in ["train", "validation", "test"]: |
| ds = f_ds[split] |
| n = len(ds) |
| cols = {col: ds[col] for col in ds.column_names if col != 'image'} |
| records = [] |
| for i in range(n): |
| rec = { |
| "idx": i, |
| "split": split, |
| "filename": cols["original_filename"][i], |
| "target3": cols["target3"][i], |
| "cancer_label": cols["cancer_label"][i], |
| "grade_label": cols["grade_label"][i], |
| "subclass_label": cols["subclass_label"][i], |
| "source_dataset": cols["source_dataset"][i], |
| "patient_id": cols["patient_id"][i], |
| "imaging_type": cols["imaging_type"][i], |
| "track_id": cols["track_id"][i], |
| "cv_fold": cols["cv_fold"][i], |
| } |
| records.append(rec) |
| INDEX[split] = records |
|
|
| |
| print("Loading CystoDS metadata...", flush=True) |
| cysto_meta = pd.read_parquet(str(CYSTO_META)) |
| cysto_meta.set_index('filename', inplace=True) |
| cysto_lookup = cysto_meta[['class', 'subclass', 'subclass2', 'stage', 'morphology', 'modality', 'lesion', 'visit', 'pid']].to_dict('index') |
|
|
| |
| def _clean(v): |
| if v is None: |
| return None |
| try: |
| if pd.isna(v): |
| return None |
| except (TypeError, ValueError): |
| pass |
| return v |
|
|
| for split in ["train", "validation", "test"]: |
| for rec in INDEX[split]: |
| fn = rec["filename"] |
| if rec["source_dataset"] == "C" and fn in cysto_lookup: |
| c = cysto_lookup[fn] |
| rec["cysto_class"] = _clean(c.get("class")) |
| rec["cysto_subclass"] = _clean(c.get("subclass")) |
| rec["cysto_subclass2"] = _clean(c.get("subclass2")) |
| rec["cysto_stage"] = _clean(c.get("stage")) |
| rec["cysto_morphology"] = _clean(c.get("morphology")) |
| rec["cysto_modality"] = _clean(c.get("modality")) |
| rec["cysto_lesion"] = _clean(c.get("lesion")) |
| rec["cysto_visit"] = _clean(c.get("visit")) |
| rec["cysto_pid"] = _clean(c.get("pid")) |
| else: |
| for k in ["cysto_class","cysto_subclass","cysto_subclass2","cysto_stage", |
| "cysto_morphology","cysto_modality","cysto_lesion","cysto_visit","cysto_pid"]: |
| rec[k] = None |
|
|
| |
| ALL_RECORDS = [] |
| for split in ["train", "validation", "test"]: |
| ALL_RECORDS.extend(INDEX[split]) |
|
|
| TOTAL = len(ALL_RECORDS) |
| print(f"Loaded {TOTAL} records across 3 splits.", flush=True) |
|
|
| |
| def letterbox(img, size=SIZE): |
| """Aspect-ratio-preserving resize with zero-padding (same as train_convnext_mlp.py).""" |
| w, h = img.size |
| scale = size / max(w, h) |
| new_w = max(int(round(w * scale)), 1) |
| new_h = max(int(round(h * scale)), 1) |
| img = img.resize((new_w, new_h), Image.BICUBIC) |
| padded = Image.new("RGB", (size, size), (0, 0, 0)) |
| padded.paste(img, ((size - new_w) // 2, (size - new_h) // 2)) |
| return padded |
|
|
| def get_image(split, idx, mode="letterbox"): |
| """Return PNG bytes for the requested image mode.""" |
| ds = f_ds[split] |
| img = ds[int(idx)]["image"].convert("RGB") |
| |
| if mode == "raw": |
| |
| buf = io.BytesIO() |
| img.save(buf, format="PNG") |
| return buf.getvalue() |
| |
| elif mode == "letterbox": |
| |
| img = letterbox(img, SIZE) |
| buf = io.BytesIO() |
| img.save(buf, format="PNG") |
| return buf.getvalue() |
| |
| elif mode == "normalized": |
| |
| img = letterbox(img, SIZE) |
| arr = np.array(img, dtype=np.float32) / 255.0 |
| arr = (arr - IMAGENET_MEAN) / IMAGENET_STD |
| |
| arr = arr * IMAGENET_STD + IMAGENET_MEAN |
| arr = np.clip(arr, 0, 1) |
| arr = (arr * 255).astype(np.uint8) |
| img = Image.fromarray(arr) |
| buf = io.BytesIO() |
| img.save(buf, format="PNG") |
| return buf.getvalue() |
| |
| else: |
| raise ValueError(f"Unknown mode: {mode}") |
|
|
| |
| app = FastAPI(title="Dataset F Viewer") |
|
|
| @app.get("/", response_class=HTMLResponse) |
| async def home(): |
| html = (VIEWER_DIR / "templates" / "index.html").read_text() |
| return HTMLResponse(html) |
|
|
| @app.get("/api/stats") |
| async def stats(): |
| """Return dataset statistics for dashboard.""" |
| |
| t3_dist = {} |
| for split in ["train", "validation", "test"]: |
| t3_dist[split] = dict(Counter(r["target3"] for r in INDEX[split])) |
| |
| |
| img_dist = {} |
| for split in ["train", "validation", "test"]: |
| img_dist[split] = dict(Counter(r["imaging_type"] for r in INDEX[split])) |
| |
| |
| src_dist = {} |
| for split in ["train", "validation", "test"]: |
| src_dist[split] = dict(Counter(r["source_dataset"] for r in INDEX[split])) |
| |
| |
| cross = defaultdict(lambda: defaultdict(int)) |
| for r in ALL_RECORDS: |
| cross[r["target3"]][r["imaging_type"]] += 1 |
| cross_tab = {str(k): dict(v) for k, v in cross.items()} |
| |
| |
| cysto_nml_sub = Counter() |
| for r in ALL_RECORDS: |
| if r["source_dataset"] == "C" and r["target3"] == 1 and r.get("cysto_subclass"): |
| cysto_nml_sub[r["cysto_subclass"]] += 1 |
| |
| |
| all_patients = set(r["patient_id"] for r in ALL_RECORDS) |
| |
| |
| src_img = defaultdict(lambda: defaultdict(int)) |
| for r in ALL_RECORDS: |
| src_img[r["source_dataset"]][r["imaging_type"]] += 1 |
| src_img_tab = {k: dict(v) for k, v in src_img.items()} |
| |
| return JSONResponse({ |
| "total": TOTAL, |
| "splits": {s: len(INDEX[s]) for s in ["train", "validation", "test"]}, |
| "target3": t3_dist, |
| "imaging_type": img_dist, |
| "source": src_dist, |
| "target3_x_imaging": cross_tab, |
| "source_x_imaging": src_img_tab, |
| "cysto_nml_subclass": dict(cysto_nml_sub), |
| "num_patients": len(all_patients), |
| "label_maps": { |
| "target3": TARGET3_NAMES, |
| "imaging_type": IMAGING_NAMES, |
| "subclass_label": SUBCLASS_NAMES, |
| "cancer_label": CANCER_NAMES, |
| "grade_label": GRADE_NAMES, |
| }, |
| "cysto_nml_subclasses": CYSTO_NML_SUBCLASSES, |
| }) |
|
|
| @app.get("/api/browse") |
| async def browse( |
| split: str = Query("", description="Filter by split"), |
| target3: str = Query("", description="Filter by target3 class"), |
| imaging_type: str = Query("", description="Filter by imaging type"), |
| source: str = Query("", description="Filter by source dataset"), |
| patient_id: str = Query("", description="Filter by patient ID"), |
| cv_fold: str = Query("", description="Filter by CV fold"), |
| track_id: str = Query("", description="Filter by track ID"), |
| subclass: str = Query("", description="Filter by CystoDS subclass (source C)"), |
| search: str = Query("", description="Search filename"), |
| sort: str = Query("idx", description="Sort field"), |
| page: int = Query(1, ge=1), |
| per_page: int = Query(50, ge=1, le=200), |
| ): |
| """Browse images with filters, return paginated results.""" |
| records = list(ALL_RECORDS) |
| |
| |
| if split: |
| records = [r for r in records if r["split"] == split] |
| if target3 != "": |
| records = [r for r in records if r["target3"] == int(target3)] |
| if imaging_type != "": |
| records = [r for r in records if r["imaging_type"] == int(imaging_type)] |
| if source: |
| records = [r for r in records if r["source_dataset"] == source] |
| if patient_id != "": |
| records = [r for r in records if r["patient_id"] == int(patient_id)] |
| if cv_fold != "": |
| records = [r for r in records if r["cv_fold"] == int(cv_fold)] |
| if track_id: |
| records = [r for r in records if r["track_id"] == track_id] |
| if subclass: |
| records = [r for r in records if r.get("cysto_subclass") == subclass] |
| if search: |
| records = [r for r in records if search.lower() in r["filename"].lower()] |
| |
| |
| sort_map = { |
| "idx": lambda r: (r["split"], r["idx"]), |
| "patient": lambda r: (r["patient_id"], r["idx"]), |
| "filename": lambda r: (r["filename"]), |
| "target3": lambda r: (r["target3"], r["idx"]), |
| "imaging": lambda r: (r["imaging_type"], r["idx"]), |
| "source": lambda r: (r["source_dataset"], r["idx"]), |
| } |
| records.sort(key=sort_map.get(sort, sort_map["idx"])) |
| |
| total = len(records) |
| start = (page - 1) * per_page |
| end = start + per_page |
| page_records = records[start:end] |
| |
| return JSONResponse({ |
| "total": total, |
| "page": page, |
| "per_page": per_page, |
| "pages": (total + per_page - 1) // per_page, |
| "records": page_records, |
| }) |
|
|
| @app.get("/api/detail/{split}/{idx}") |
| async def detail(split: str, idx: int): |
| """Get full metadata for a single image.""" |
| if split not in INDEX or idx < 0 or idx >= len(INDEX[split]): |
| raise HTTPException(404, "Image not found") |
| return JSONResponse(INDEX[split][idx]) |
|
|
| @app.get("/api/patient/{pid}") |
| async def patient_view(pid: int): |
| """Get all images for a patient, grouped by track_id.""" |
| patient_records = [r for r in ALL_RECORDS if r["patient_id"] == pid] |
| |
| |
| tracks = defaultdict(list) |
| for r in patient_records: |
| tracks[r["track_id"]].append(r) |
| |
| |
| for tid in tracks: |
| tracks[tid].sort(key=lambda r: r["idx"]) |
| |
| |
| t3_dist = Counter(r["target3"] for r in patient_records) |
| src_dist = Counter(r["source_dataset"] for r in patient_records) |
| img_dist = Counter(r["imaging_type"] for r in patient_records) |
| |
| return JSONResponse({ |
| "patient_id": pid, |
| "total_images": len(patient_records), |
| "splits": dict(Counter(r["split"] for r in patient_records)), |
| "target3_dist": dict(t3_dist), |
| "source_dist": dict(src_dist), |
| "imaging_dist": dict(img_dist), |
| "tracks": {tid: recs for tid, recs in sorted(tracks.items())}, |
| "cysto_subclasses": dict(Counter(r["cysto_subclass"] for r in patient_records |
| if r.get("cysto_subclass"))), |
| }) |
|
|
| @app.get("/api/filters") |
| async def get_filter_options(): |
| """Get available filter options (unique values).""" |
| patients = sorted(set(r["patient_id"] for r in ALL_RECORDS)) |
| tracks = sorted(set(r["track_id"] for r in ALL_RECORDS if r["track_id"] != "NA")) |
| cysto_subs = sorted(set(r["cysto_subclass"] for r in ALL_RECORDS if r.get("cysto_subclass"))) |
| |
| return JSONResponse({ |
| "patients": patients, |
| "tracks": tracks, |
| "cysto_subclasses": cysto_subs, |
| "splits": list(INDEX.keys()), |
| "target3": [{"value": k, "name": v} for k, v in TARGET3_NAMES.items()], |
| "imaging_types": [{"value": k, "name": v} for k, v in IMAGING_NAMES.items()], |
| "sources": ["B", "C", "D"], |
| "cv_folds": [0, 1, 2, 3, 4], |
| }) |
|
|
| @app.get("/image/{split}/{idx}") |
| async def serve_image(split: str, idx: int, mode: str = Query("letterbox")): |
| """Serve an image in the requested mode.""" |
| if split not in f_ds or idx < 0 or idx >= len(f_ds[split]): |
| raise HTTPException(404, "Image not found") |
| img_bytes = get_image(split, idx, mode) |
| return Response(content=img_bytes, media_type="image/png") |
|
|
| @app.get("/thumb/{split}/{idx}") |
| async def serve_thumb(split: str, idx: int, mode: str = Query("letterbox")): |
| """Serve a thumbnail (100×100) of an image.""" |
| if split not in f_ds or idx < 0 or idx >= len(f_ds[split]): |
| raise HTTPException(404, "Image not found") |
| img_bytes = get_image(split, idx, mode) |
| img = Image.open(io.BytesIO(img_bytes)) |
| img = img.resize((100, 100), Image.BICUBIC) |
| buf = io.BytesIO() |
| img.save(buf, format="PNG") |
| return Response(content=buf.getvalue(), media_type="image/png") |
|
|
| if __name__ == "__main__": |
| import uvicorn |
| uvicorn.run(app, host="0.0.0.0", port=8910) |
|
|