cxr-report-generation / app_old.py
adhisetiawan's picture
Deploy CXR report generation demo
16d6749
Raw
History Blame Contribute Delete
9.81 kB
from __future__ import annotations
import argparse
import json
from pathlib import Path
from typing import Any
from dotenv import load_dotenv
from PIL import Image
load_dotenv(Path(__file__).resolve().parent / ".env")
from src.clients.medgemma_client import MedGemmaClient, MedGemmaDetector, MedGemmaReporter
from src.clients.medsam_client import MedSAMClient
from src.config.endpoints import (
DEFAULT_MAX_FINDINGS,
MEDSAM_CHECKPOINT,
MEDSAM_DEVICE,
MEDGEMMA_MODEL_ID,
)
from src.dataset import get_demo_case, load_demo_cases, resolve_case_images
from src.pipeline import CXRPipeline
from src.schemas.detection import DetectionResult
from src.schemas.report import PipelineOutput
from src.localization_quality import calculate_bbox_area_ratio
from src.visualization import draw_view_localizations
# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description=(
"Three-stage chest X-ray analysis: "
"(1) detect findings, (2) segment them, (3) generate a report."
)
)
input_group = parser.add_mutually_exclusive_group(required=True)
input_group.add_argument("--image", type=Path, help="Path to one chest X-ray image.")
input_group.add_argument("--case-id", type=str, help="Case ID from demo_cases.jsonl.")
parser.add_argument(
"--demo-cases",
type=Path,
default=Path("data/demo_cases.jsonl"),
help="Path to the demo JSONL file.",
)
parser.add_argument(
"--dataset-root",
type=Path,
help="Dataset root for resolving JSONL image paths. Required with --case-id.",
)
parser.add_argument(
"--output-dir",
type=Path,
default=Path("outputs"),
help="Directory for outputs.",
)
parser.add_argument(
"--max-findings",
type=int,
default=DEFAULT_MAX_FINDINGS,
help=f"Max findings to localize. Default: {DEFAULT_MAX_FINDINGS}.",
)
parser.add_argument(
"--model-id",
type=str,
default=MEDGEMMA_MODEL_ID,
help=f"MedGemma model ID. Default: {MEDGEMMA_MODEL_ID}.",
)
parser.add_argument(
"--segment",
action="store_true",
help="Run Stage 2: segment each localized finding with MedSAM.",
)
parser.add_argument(
"--medsam-checkpoint",
type=str,
default=MEDSAM_CHECKPOINT,
help=f"MedSAM checkpoint path. Default: {MEDSAM_CHECKPOINT}.",
)
parser.add_argument(
"--medsam-device",
type=str,
default=MEDSAM_DEVICE,
help=f"Device for MedSAM. Default: {MEDSAM_DEVICE}.",
)
parser.add_argument(
"--report",
action="store_true",
help="Run Stage 3: generate a radiology report with MedGemma.",
)
parser.add_argument(
"--show",
action="store_true",
help="Open each annotated view image after processing.",
)
return parser
# ---------------------------------------------------------------------------
# Input resolution
# ---------------------------------------------------------------------------
def resolve_input(
args: argparse.Namespace,
) -> tuple[list[Path], str | None, dict[str, Any] | None]:
if args.image is not None:
if not args.image.exists():
raise FileNotFoundError(f"Input image not found: {args.image}")
return [args.image], None, None
if args.dataset_root is None:
raise ValueError("--dataset-root is required when using --case-id.")
cases = load_demo_cases(args.demo_cases)
case = get_demo_case(cases, args.case_id)
image_paths = resolve_case_images(case, args.dataset_root)
return image_paths, args.case_id, case
# ---------------------------------------------------------------------------
# Console output helpers
# ---------------------------------------------------------------------------
def print_study_input(image_paths: list[Path], case_id: str | None) -> None:
print("\n=== Study Input ===")
if case_id is not None:
print(f"Case ID : {case_id}")
print(f"Views : {len(image_paths)}")
for i, p in enumerate(image_paths):
print(f" [{i}] {p}")
def print_detection_result(result: DetectionResult) -> None:
print("\n=== Stage 1 – Detection Results ===")
if not result.findings:
print("No positive abnormal radiographic findings returned.")
return
for i, finding in enumerate(result.findings, start=1):
print(f"\n[{i}] {finding.finding}")
print(f" Location : {finding.anatomical_location}")
print(f" Certainty : {finding.certainty}")
for view in finding.localizations:
print(f"\n View [{view.image_index}]")
print(f" Image : {view.image_path}")
print(f" Status : {view.status}")
if view.status == "localized":
for j, box in enumerate(view.boxes, start=1):
area = calculate_bbox_area_ratio(box)
print(f" BBox {j}: {box.box_2d} (label={box.label})")
print(f" Area : {area:.4f}")
elif view.status == "rejected_by_quality_gate":
for reason in view.rejection_reasons:
print(f" Reason : {reason}")
elif view.status == "parser_error":
for reason in view.rejection_reasons:
print(f" Reason : {reason}")
def print_segmentation_summary(output: PipelineOutput) -> None:
if not output.masks:
return
print("\n=== Stage 2 – Segmentation Summary ===")
for mask in output.masks:
shape = mask.mask.shape if mask.mask is not None else "N/A"
print(f" [{mask.image_index}] {mask.finding_label} | status={mask.status} | mask shape={shape}")
def print_report(output: PipelineOutput) -> None:
if output.report is None:
return
print("\n=== Stage 3 – Radiology Report ===")
if output.report.status == "failed":
print(f"Report generation failed: {output.report.error}")
else:
print(output.report.report_text)
print("===================================")
def print_hidden_reference(case: dict[str, Any] | None) -> None:
if case is None:
return
print("\n=== Hidden IU X-Ray Reference (not sent to model) ===")
for i, finding in enumerate(case.get("reference_findings", []), start=1):
print(f"[{i}] {finding['finding']} | {finding['location']} | {finding['certainty']}")
# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
def main() -> None:
parser = build_parser()
args = parser.parse_args()
image_paths, case_id, demo_case = resolve_input(args)
print_study_input(image_paths, case_id)
images = [Image.open(p).convert("RGB") for p in image_paths]
input_images = [str(p) for p in image_paths]
# -- build pipeline components ------------------------------------------
print(f"\nLoading MedGemma: {args.model_id}")
medgemma = MedGemmaClient(model_id=args.model_id)
detector = MedGemmaDetector(client=medgemma, max_findings=args.max_findings)
segmenter = None
if args.segment:
print(f"Loading MedSAM : {args.medsam_checkpoint} on {args.medsam_device}")
segmenter = MedSAMClient(
checkpoint_path=args.medsam_checkpoint,
device=args.medsam_device,
)
reporter = None
if args.report:
reporter = MedGemmaReporter(client=medgemma)
# -- run pipeline --------------------------------------------------------
pipeline = CXRPipeline(detector=detector, segmenter=segmenter, reporter=reporter)
output = pipeline.run(images=images, input_images=input_images, case_id=case_id)
# -- save outputs --------------------------------------------------------
run_name = case_id or image_paths[0].stem
run_dir = args.output_dir / run_name
run_dir.mkdir(parents=True, exist_ok=True)
# Stage 1 detection JSON
detection_json = output.detection.model_dump()
(run_dir / "result.json").write_text(
json.dumps(detection_json, indent=2, ensure_ascii=False), encoding="utf-8"
)
# Stage 3 report JSON (if generated)
if output.report is not None:
report_json = output.report.model_dump()
(run_dir / "report.json").write_text(
json.dumps(report_json, indent=2, ensure_ascii=False), encoding="utf-8"
)
# Annotated images
annotated_paths: list[Path] = []
for image_index, processed_image in enumerate(output.processed_images):
annotated = draw_view_localizations(
image=processed_image,
image_index=image_index,
findings=output.detection.findings,
)
annotated_path = run_dir / f"annotated_{image_index}.png"
annotated.save(annotated_path)
annotated_paths.append(annotated_path)
if args.show:
annotated.show()
# -- console summary -----------------------------------------------------
print_detection_result(output.detection)
print_segmentation_summary(output)
print_report(output)
print_hidden_reference(demo_case)
print("\n=== Saved Outputs ===")
print(f"Detection JSON : {run_dir / 'result.json'}")
if output.report is not None:
print(f"Report JSON : {run_dir / 'report.json'}")
for p in annotated_paths:
print(f"Annotated : {p}")
if __name__ == "__main__":
main()