cxr-report-generation / tests /test_multiview_pipeline.py
adhisetiawan's picture
Deploy CXR report generation demo
16d6749
Raw
History Blame Contribute Delete
15.1 kB
from __future__ import annotations
import unittest
from collections.abc import Sequence
from PIL import Image
from src.clients.medgemma_client import MedGemmaDetector
from src.schemas.detection import (
BBox as LocalizationBox,
DetectionResult as PipelineResult,
LocalizedFinding,
ViewLocalization,
)
from src.visualization import draw_view_localizations
from tests.test_parsing import REAL_MEDGEMMA_DISCOVERY_RESPONSE
def make_test_image(color: tuple[int, int, int] = (128, 128, 128)) -> Image.Image:
return Image.new("RGB", (64, 48), color)
def discovery_response(findings: list[dict[str, str]]) -> str:
import json
return f"```json\n{json.dumps({'findings': findings})}\n```"
def localization_response(boxes: list[dict[str, object]]) -> str:
import json
return f"```json\n{json.dumps(boxes)}\n```"
# ---------------------------------------------------------------------------
# Mock clients
# ---------------------------------------------------------------------------
class MockMedGemmaClient:
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
def generate(self, *, images: Sequence[Image.Image], prompt: str, max_new_tokens: int) -> str:
self.calls.append({"images": list(images), "prompt": prompt, "max_new_tokens": max_new_tokens})
call_index = len(self.calls) - 1
if call_index == 0:
return discovery_response([
{"finding": "pneumothorax", "anatomical_location": "right apex", "certainty": "positive"},
{"finding": "rib fracture", "anatomical_location": "right lateral 8th rib", "certainty": "positive"},
])
localization_index = call_index - 1
if localization_index == 0:
return localization_response([{"box_2d": [100, 200, 150, 250], "label": "pneumothorax"}])
if localization_index == 1:
return localization_response([])
if localization_index == 2:
return localization_response([{"box_2d": [300, 400, 350, 450], "label": "rib fracture"}])
return localization_response([])
class SingleFindingMockMedGemmaClient:
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
def generate(self, *, images: Sequence[Image.Image], prompt: str, max_new_tokens: int) -> str:
self.calls.append({"images": list(images), "prompt": prompt, "max_new_tokens": max_new_tokens})
if len(self.calls) == 1:
return discovery_response([
{"finding": "pulmonary nodule", "anatomical_location": "right upper lobe", "certainty": "positive"}
])
return localization_response([{"box_2d": [120, 610, 290, 790], "label": "pulmonary nodule"}])
class RealMedGemmaDiscoveryMockClient:
def generate(self, *, images: Sequence[Image.Image], prompt: str, max_new_tokens: int) -> str:
return REAL_MEDGEMMA_DISCOVERY_RESPONSE
class BroadFractureMockMedGemmaClient:
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
def generate(self, *, images: Sequence[Image.Image], prompt: str, max_new_tokens: int) -> str:
self.calls.append({"images": list(images), "prompt": prompt, "max_new_tokens": max_new_tokens})
if len(self.calls) == 1:
return discovery_response([
{"finding": "fracture", "anatomical_location": "right humerus", "certainty": "positive"}
])
return localization_response([{"box_2d": [100, 100, 700, 500], "label": "fracture"}])
class MalformedLocalizationMockMedGemmaClient:
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
def generate(self, *, images: Sequence[Image.Image], prompt: str, max_new_tokens: int) -> str:
self.calls.append({"images": list(images), "prompt": prompt, "max_new_tokens": max_new_tokens})
if len(self.calls) == 1:
return discovery_response([
{"finding": "nodule", "anatomical_location": "right upper lobe", "certainty": "positive"}
])
return "not valid json at all"
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
class TestMultiViewDiscovery(unittest.TestCase):
def test_discover_findings_normalizes_real_medgemma_response(self) -> None:
detector = MedGemmaDetector(client=RealMedGemmaDiscoveryMockClient()) # type: ignore[arg-type]
discovery = detector.discover_findings(
[make_test_image(), make_test_image(color=(64, 64, 64))]
)
self.assertEqual(len(discovery.findings), 2)
self.assertEqual(discovery.findings[0].finding, "nodule")
self.assertEqual(discovery.findings[0].anatomical_location, "right upper lobe")
self.assertEqual(discovery.findings[1].anatomical_location, "left upper lobe")
def test_a_single_multi_view_discovery_call(self) -> None:
mock_client = MockMedGemmaClient()
detector = MedGemmaDetector(client=mock_client) # type: ignore[arg-type]
detector.detect(
[make_test_image(), make_test_image(color=(64, 64, 64))],
["/study/0.png", "/study/1.png"],
)
discovery_calls = [c for c in mock_client.calls if len(c["images"]) > 1] # type: ignore[arg-type]
self.assertEqual(len(discovery_calls), 1)
self.assertEqual(len(discovery_calls[0]["images"]), 2) # type: ignore[arg-type]
def test_b_localization_call_count(self) -> None:
mock_client = MockMedGemmaClient()
detector = MedGemmaDetector(client=mock_client) # type: ignore[arg-type]
detector.detect(
[make_test_image(), make_test_image(color=(64, 64, 64))],
["/study/0.png", "/study/1.png"],
)
localization_calls = [c for c in mock_client.calls if len(c["images"]) == 1] # type: ignore[arg-type]
self.assertEqual(len(localization_calls), 4)
def test_c_view_localization_metadata(self) -> None:
mock_client = MockMedGemmaClient()
detector = MedGemmaDetector(client=mock_client) # type: ignore[arg-type]
result, _ = detector.detect(
[make_test_image(), make_test_image(color=(64, 64, 64))],
["/study/0.png", "/study/1.png"],
)
for finding in result.findings:
self.assertEqual(len(finding.localizations), 2)
for view in finding.localizations:
self.assertIsInstance(view.image_index, int)
self.assertTrue(view.image_path)
self.assertIsInstance(view.status, str)
self.assertIsInstance(view.boxes, list)
self.assertIsInstance(view.candidate_boxes, list)
def test_d_finding_visible_in_one_view_only(self) -> None:
mock_client = MockMedGemmaClient()
detector = MedGemmaDetector(client=mock_client) # type: ignore[arg-type]
result, _ = detector.detect(
[make_test_image(), make_test_image(color=(64, 64, 64))],
["/study/0.png", "/study/1.png"],
)
pneumothorax = result.findings[0]
view_0 = pneumothorax.localizations[0]
view_1 = pneumothorax.localizations[1]
self.assertEqual(view_0.status, "localized")
self.assertEqual(len(view_0.boxes), 1)
self.assertEqual(view_0.boxes[0].label, "pneumothorax")
self.assertEqual(view_1.status, "abstained")
self.assertEqual(view_1.boxes, [])
def test_f_direct_image_mode(self) -> None:
mock_client = SingleFindingMockMedGemmaClient()
detector = MedGemmaDetector(client=mock_client) # type: ignore[arg-type]
result, processed_images = detector.detect(
[make_test_image()],
["/single/0.png"],
)
discovery_calls = [
c for c in mock_client.calls
if len(c["images"]) == 1 and c is mock_client.calls[0] # type: ignore[arg-type]
]
self.assertEqual(len(discovery_calls), 1)
self.assertEqual(len(mock_client.calls) - 1, 1)
self.assertEqual(len(processed_images), 1)
self.assertEqual(len(result.input_images), 1)
self.assertEqual(len(result.findings[0].localizations), 1)
def test_g_hidden_reference_not_passed_to_generate(self) -> None:
mock_client = MockMedGemmaClient()
detector = MedGemmaDetector(client=mock_client) # type: ignore[arg-type]
detector.detect(
[make_test_image(), make_test_image(color=(64, 64, 64))],
["/study/0.png", "/study/1.png"],
)
for call in mock_client.calls:
prompt = str(call["prompt"])
self.assertNotIn("reference_report", prompt)
self.assertNotIn("reference_findings", prompt)
self.assertNotIn("Hidden IU X-Ray Reference", prompt)
class TestLocalizationGateStates(unittest.TestCase):
def test_d_explicit_medgemma_abstention(self) -> None:
mock_client = MockMedGemmaClient()
detector = MedGemmaDetector(client=mock_client) # type: ignore[arg-type]
result, _ = detector.detect(
[make_test_image(), make_test_image(color=(64, 64, 64))],
["/study/0.png", "/study/1.png"],
)
pneumothorax_view_1 = result.findings[0].localizations[1]
self.assertEqual(pneumothorax_view_1.status, "abstained")
self.assertEqual(pneumothorax_view_1.boxes, [])
self.assertEqual(pneumothorax_view_1.candidate_boxes, [])
def test_e_focal_candidate_fully_rejected(self) -> None:
mock_client = BroadFractureMockMedGemmaClient()
detector = MedGemmaDetector(client=mock_client) # type: ignore[arg-type]
result, _ = detector.detect([make_test_image()], ["/study/0.png"])
view = result.findings[0].localizations[0]
self.assertEqual(view.status, "rejected_by_quality_gate")
self.assertEqual(view.boxes, [])
self.assertEqual(len(view.candidate_boxes), 1)
self.assertTrue(view.rejection_reasons)
def test_f_parser_failure(self) -> None:
mock_client = MalformedLocalizationMockMedGemmaClient()
detector = MedGemmaDetector(client=mock_client) # type: ignore[arg-type]
result, _ = detector.detect([make_test_image()], ["/study/0.png"])
view = result.findings[0].localizations[0]
self.assertEqual(view.status, "parser_error")
self.assertEqual(view.boxes, [])
self.assertEqual(view.candidate_boxes, [])
self.assertTrue(view.rejection_reasons)
class TestAnnotatedImageConsistency(unittest.TestCase):
def test_e_annotated_image_contains_only_matching_view_boxes(self) -> None:
findings = [
LocalizedFinding(
finding="pneumothorax",
anatomical_location="right apex",
certainty="positive",
localizations=[
ViewLocalization(
image_index=0,
image_path="/study/0.png",
status="localized",
boxes=[LocalizationBox(box_2d=[100, 100, 200, 200], label="pneumothorax")],
),
ViewLocalization(
image_index=1, image_path="/study/1.png", status="abstained", boxes=[]
),
],
),
LocalizedFinding(
finding="rib fracture",
anatomical_location="right lateral 8th rib",
certainty="positive",
localizations=[
ViewLocalization(
image_index=0, image_path="/study/0.png", status="abstained", boxes=[]
),
ViewLocalization(
image_index=1,
image_path="/study/1.png",
status="localized",
boxes=[LocalizationBox(box_2d=[300, 300, 400, 400], label="rib fracture")],
),
],
),
]
blank = Image.new("RGB", (1000, 1000), (255, 255, 255))
annotated_0 = draw_view_localizations(blank.copy(), 0, findings)
annotated_1 = draw_view_localizations(blank.copy(), 1, findings)
self.assertNotEqual(list(annotated_0.getdata()), list(blank.getdata()))
self.assertNotEqual(list(annotated_1.getdata()), list(blank.getdata()))
view_1_only = draw_view_localizations(blank.copy(), 1, findings[:1])
self.assertEqual(list(view_1_only.getdata()), list(blank.getdata()))
view_0_only_other = draw_view_localizations(blank.copy(), 0, findings[1:])
self.assertEqual(list(view_0_only_other.getdata()), list(blank.getdata()))
def test_g_visualizer_uses_accepted_boxes_only(self) -> None:
accepted_box = LocalizationBox(box_2d=[100, 100, 200, 200], label="nodule")
rejected_candidate = LocalizationBox(box_2d=[100, 100, 700, 500], label="fracture")
findings = [
LocalizedFinding(
finding="nodule",
anatomical_location="right upper lobe",
certainty="positive",
localizations=[
ViewLocalization(
image_index=0,
image_path="/study/0.png",
status="localized",
boxes=[accepted_box],
candidate_boxes=[accepted_box],
)
],
),
LocalizedFinding(
finding="fracture",
anatomical_location="right humerus",
certainty="positive",
localizations=[
ViewLocalization(
image_index=0,
image_path="/study/0.png",
status="rejected_by_quality_gate",
boxes=[],
candidate_boxes=[rejected_candidate],
rejection_reasons=["bbox_area_ratio=0.2400 exceeds focal max_area_ratio=0.1200"],
)
],
),
]
blank = Image.new("RGB", (1000, 1000), (255, 255, 255))
annotated = draw_view_localizations(blank.copy(), 0, findings)
self.assertNotEqual(list(annotated.getdata()), list(blank.getdata()))
rejected_only = draw_view_localizations(blank.copy(), 0, [findings[1]])
self.assertEqual(list(rejected_only.getdata()), list(blank.getdata()))
class TestPipelineResultSchema(unittest.TestCase):
def test_pipeline_result_uses_input_images(self) -> None:
result = PipelineResult(
case_id="CXR3281_IM-1562",
input_images=["/study/0.png", "/study/1.png"],
findings=[],
)
dumped = result.model_dump()
self.assertIn("input_images", dumped)
self.assertNotIn("input_image", dumped)
if __name__ == "__main__":
unittest.main()