histopath-pni / app.py
@PrajwalDange
Claude Opus 4.7
Add 10x v4.1 classifier and phone-camera preprocessing module
73b9e48
Raw
History Blame Contribute Delete
16.5 kB
"""
PNI Detection in Histopathology β€” Gradio Web App
AI-powered two-stage detection of perineural invasion in H&E-stained
histopathology images using the Phikon-v2 foundation model.
Stage 1: Nerve detection (identifies nerve structures)
Stage 2: PNI classification (determines if tumour invades the nerve)
"""
import os
import torch
import joblib
import numpy as np
import pandas as pd
import gradio as gr
from pathlib import Path
from datetime import datetime, timezone
from transformers import AutoModel, AutoImageProcessor
from inference import run_inference
from feedback import submit_feedback, SHEETS_ENABLED
from phone_preprocess import preprocess_phone_image
# ── Global model loading (runs once at startup) ────────────────────────
print("Loading Phikon-v2 foundation model...")
DEVICE = "cuda" if torch.cuda.is_available() else (
"mps" if torch.backends.mps.is_available() else "cpu"
)
print(f" Device: {DEVICE}")
dtype = torch.float16 if DEVICE == "cuda" else torch.float32
MODEL = AutoModel.from_pretrained(
"owkin/phikon-v2",
trust_remote_code=True,
torch_dtype=dtype,
).to(DEVICE).eval()
PROCESSOR = AutoImageProcessor.from_pretrained(
"owkin/phikon-v2",
trust_remote_code=True,
use_fast=True,
)
print("Loading pre-trained classifiers...")
BASE_DIR = Path(__file__).parent
NERVE_CLF = joblib.load(BASE_DIR / "classifiers" / "nerve_clf.pkl")
PNI_CLF = joblib.load(BASE_DIR / "classifiers" / "pni_clf.pkl")
# 10x-specific classifiers (optional β€” gracefully absent before training)
_10x_nerve = BASE_DIR / "classifiers" / "nerve_clf_10x.pkl"
_10x_pni = BASE_DIR / "classifiers" / "pni_clf_10x.pkl"
NERVE_CLF_10X = joblib.load(_10x_nerve) if _10x_nerve.exists() else None
PNI_CLF_10X = joblib.load(_10x_pni) if _10x_pni.exists() else None
if NERVE_CLF_10X is not None:
print(" 10x classifiers loaded.")
else:
print(" 10x classifiers not found β€” will use 20x for all magnifications.")
# Macenko stain normaliser (optional β€” for correcting lab-to-lab stain variability)
_macenko_path = BASE_DIR / "classifiers" / "macenko_normalizer.pkl"
STAIN_NORMALIZER = joblib.load(_macenko_path) if _macenko_path.exists() else None
if STAIN_NORMALIZER is not None:
print(" Macenko stain normaliser loaded.")
else:
print(" Macenko normaliser not found β€” stain normalisation unavailable.")
print("Ready!\n")
# ── Inference function ──────────────────────────────────────────────────
def analyze_image(
image,
magnification,
crop_top,
crop_bottom,
crop_left,
crop_right,
nerve_threshold,
pni_threshold,
stain_norm,
phone_preprocess,
):
"""Process an uploaded image and return results (5 outputs: img, verdict, table, state, timestamp)."""
if image is None:
return None, "Please upload an image.", pd.DataFrame(), None, None
normalizer = STAIN_NORMALIZER if stain_norm else None
# Apply phone-camera preprocessing (white balance, gamma, CLAHE, eyepiece-border crop)
# before any other step. Designed for phone-through-microscope captures.
if phone_preprocess:
try:
image, _diag = preprocess_phone_image(image)
except Exception as exc:
print(f" Phone preprocessing failed: {exc} β€” using raw image")
annotated, verdict, regions = run_inference(
image=image,
model=MODEL,
processor=PROCESSOR,
nerve_clf=NERVE_CLF,
pni_clf=PNI_CLF,
device=DEVICE,
crop_top=int(crop_top),
crop_bottom=int(crop_bottom),
crop_left=int(crop_left),
crop_right=int(crop_right),
nerve_threshold=nerve_threshold,
pni_threshold=pni_threshold,
magnification=magnification,
nerve_clf_10x=NERVE_CLF_10X,
pni_clf_10x=PNI_CLF_10X,
stain_normalizer=normalizer,
)
if regions:
df = pd.DataFrame([
{
"Region": f"R{r['region_id']}",
"Nerve Confidence": f"{r['nerve_prob']:.1%}",
"PNI Probability": f"{r['pni_prob']:.1%}",
"PNI Status": "POSITIVE" if r["pni_positive"] else "Negative",
"Patches": r["n_patches"],
}
for r in regions
])
else:
df = pd.DataFrame(columns=[
"Region", "Nerve Confidence", "PNI Probability", "PNI Status", "Patches"
])
# 4th value β†’ result_state (verdict), 5th β†’ analysis_time_state (when Analyze was pressed)
analysis_ts = datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")
return annotated, verdict, df, verdict, analysis_ts
# ── Feedback handler ────────────────────────────────────────────────────
def submit_feedback_handler(
name,
photo_type,
magnification,
verdict_raw, # from result_state
analyzed, # from analyzed_state (bool)
analysis_timestamp, # from analysis_time_state
path_nerve_count,
path_pni_count,
comments,
):
"""Guard against pre-analysis submission, then delegate to feedback.py."""
if not analyzed or verdict_raw is None:
return "Please click Analyze first before submitting feedback."
return submit_feedback(
name=name,
photo_type=photo_type,
magnification=magnification,
verdict=verdict_raw,
path_nerve_count=path_nerve_count,
path_pni_count=path_pni_count,
comments=comments,
analysis_was_run=analyzed,
analysis_timestamp=analysis_timestamp,
)
# ── UI strings ──────────────────────────────────────────────────────────
TITLE = "AI-Based PNI Detection in Histopathology"
DESCRIPTION = """
Upload an H&E-stained histopathology image to detect **nerve structures**
and classify **perineural invasion (PNI)**.
The model uses [Phikon-v2](https://huggingface.co/owkin/phikon-v2), a
state-of-the-art pathology foundation model trained on 460 million tiles,
with lightweight classifiers for nerve detection (AUC 0.999) and PNI
classification (AUC 0.979).
**Results:** Green boxes = nerve without PNI. Red boxes = nerve with PNI.
"""
DISCLAIMER = """
---
**Research Use Only.** This tool is intended for research and educational
purposes. It has not been validated for clinical diagnostic use and should
not replace professional pathological assessment. All predictions should
be verified by a qualified pathologist.
**License:** CC BY-NC 4.0 β€” Free for non-commercial research use with attribution.
Consistent with the upstream [Phikon-v2 non-commercial license](https://huggingface.co/owkin/phikon-v2/blob/main/LICENSE.pdf).
"""
# Find example images
example_dir = BASE_DIR / "examples"
examples = []
if example_dir.exists():
for f in sorted(example_dir.glob("*.jpg")):
examples.append([str(f), "20x", 0, 0, 0, 0, 0.7, 0.5])
# ── Gradio layout ───────────────────────────────────────────────────────
with gr.Blocks(title=TITLE) as demo:
# Session state
result_state = gr.State(value=None) # raw verdict string from last analysis
analyzed_state = gr.State(value=False) # True once Analyze has been clicked
analysis_time_state = gr.State(value=None) # UTC timestamp when Analyze was pressed
gr.Markdown(f"# {TITLE}")
gr.Markdown(DESCRIPTION)
# Banner when feedback storage is not configured
if not SHEETS_ENABLED:
gr.Markdown(
"> **Note:** Feedback storage is not configured for this deployment. "
"The AI analysis works normally, but submitted feedback will not be saved."
)
with gr.Row():
# ── Left column: inputs ────────────────────────────────────────
with gr.Column(scale=1):
# Pathologist / user info β€” filled once per session
with gr.Accordion("Your Information", open=True):
gr.Markdown(
"Please fill in once per session. This is stored alongside "
"your feedback to help evaluate and improve the AI system."
)
user_name = gr.Textbox(
label="Name",
placeholder="Dr. Jane Smith",
)
photo_type = gr.Dropdown(
choices=[
"Phone Camera",
"DSLR / Microscope Camera",
"Whole Slide Image (WSI)",
"Other",
],
label="Image Source / Photo Type",
value="DSLR / Microscope Camera",
info="How was this image captured?",
)
magnification = gr.Radio(
choices=["10x", "20x"],
value="20x",
label="Image Magnification",
info="Select the magnification at which the image was captured",
)
input_image = gr.Image(
type="numpy",
label="Upload H&E Image",
height=400,
)
with gr.Accordion("Advanced Settings", open=False):
gr.Markdown(
"**Microscope UI Overlay Removal** β€” Set pixel values to "
"crop if your microscope adds a scale bar or metadata overlay."
)
with gr.Row():
crop_top = gr.Number(label="Crop Top (px)", value=0, minimum=0, maximum=500)
crop_bottom = gr.Number(label="Crop Bottom (px)", value=0, minimum=0, maximum=500)
with gr.Row():
crop_left = gr.Number(label="Crop Left (px)", value=0, minimum=0, maximum=500)
crop_right = gr.Number(label="Crop Right (px)", value=0, minimum=0, maximum=500)
gr.Markdown("**Stain Normalisation**")
stain_norm_cb = gr.Checkbox(
label="Apply Macenko stain normalisation",
value=False,
info=(
"Recommended when images come from a different lab or "
"scanner than the training data. Normalises H&E stain "
"colour to a reference slide before feature extraction."
),
interactive=STAIN_NORMALIZER is not None,
)
gr.Markdown("**Phone-Camera Preprocessing**")
phone_preprocess_cb = gr.Checkbox(
label="Apply phone-camera preprocessing",
value=False,
info=(
"Enable when the image was captured with a phone through "
"a microscope eyepiece. Runs eyepiece-border crop, "
"gray-world white balance, auto-gamma exposure correction "
"and CLAHE local-contrast enhancement before classification. "
"Best used together with Macenko stain normalisation."
),
)
gr.Markdown("**Detection Thresholds**")
nerve_thresh = gr.Slider(
minimum=0.5, maximum=0.95, value=0.7, step=0.05,
label="Nerve Detection Threshold",
info="Higher = fewer but more confident detections",
)
pni_thresh = gr.Slider(
minimum=0.3, maximum=0.8, value=0.5, step=0.05,
label="PNI Classification Threshold",
info="Higher = more specific, lower = more sensitive",
)
analyze_btn = gr.Button("Analyze", variant="primary", size="lg")
# ── Right column: outputs + feedback ──────────────────────────
with gr.Column(scale=1):
output_image = gr.Image(label="Detection Results", height=400)
verdict_text = gr.Textbox(label="Verdict", lines=2)
regions_table = gr.Dataframe(
label="Region Details",
headers=["Region", "Nerve Confidence", "PNI Probability",
"PNI Status", "Patches"],
)
# Pathologist assessment β€” filled after each analysis
with gr.Accordion("Pathologist Feedback", open=False):
gr.Markdown(
"After reviewing the AI result and the original image, "
"enter your own assessment below. "
"Your input is cross-referenced with the AI output to "
"measure real-world accuracy and drive future retraining."
)
with gr.Row():
path_nerve_count = gr.Number(
label="Nerves you identified",
value=0, minimum=0, maximum=50, step=1,
info="Total nerve profiles visible in this image",
)
path_pni_count = gr.Number(
label="Nerves with PNI+",
value=0, minimum=0, maximum=50, step=1,
info="Of those, how many show perineural invasion?",
)
path_comments = gr.Textbox(
label="Comments",
lines=3,
placeholder=(
"Optional β€” note any disagreements with the AI, image "
"quality issues, or unusual features."
),
)
submit_btn = gr.Button("Submit Feedback", variant="secondary")
feedback_status = gr.Textbox(
label="Submission Status",
interactive=False,
value="",
lines=1,
)
# Examples
if examples:
gr.Examples(
examples=examples,
inputs=[
input_image, magnification,
crop_top, crop_bottom, crop_left, crop_right,
nerve_thresh, pni_thresh,
],
outputs=[output_image, verdict_text, regions_table],
fn=lambda img, mag, ct, cb, cl, cr, nt, pt: analyze_image(
img, mag, ct, cb, cl, cr, nt, pt, False
)[:3], # examples only need the first 3 outputs (img, verdict, table)
cache_examples=False,
label="Example Images (click to try)",
)
gr.Markdown(DISCLAIMER)
# ── Event wiring ───────────────────────────────────────────────────
# Analyze β†’ run inference, store verdict + analysis timestamp, mark analyzed=True
analyze_btn.click(
fn=analyze_image,
inputs=[
input_image, magnification,
crop_top, crop_bottom, crop_left, crop_right,
nerve_thresh, pni_thresh,
stain_norm_cb, phone_preprocess_cb,
],
outputs=[output_image, verdict_text, regions_table, result_state, analysis_time_state],
).then(
fn=lambda: True,
inputs=[],
outputs=[analyzed_state],
)
# New image uploaded β†’ reset analyzed flag, clear analysis time, clear feedback status
input_image.change(
fn=lambda: (False, None, ""),
inputs=[],
outputs=[analyzed_state, analysis_time_state, feedback_status],
)
# Submit Feedback button
submit_btn.click(
fn=submit_feedback_handler,
inputs=[
user_name, photo_type, magnification,
result_state, analyzed_state, analysis_time_state,
path_nerve_count, path_pni_count, path_comments,
],
outputs=[feedback_status],
)
# Launch
if __name__ == "__main__":
demo.queue()
demo.launch(server_name="0.0.0.0", server_port=7860, share=True)