cxr-report-generation / src /clients /medsam_client.py
adhisetiawan's picture
Sync latest application changes
21245a4
Raw
History Blame Contribute Delete
5.01 kB
from __future__ import annotations
import numpy as np
import torch
import torch.nn.functional as F
from PIL import Image
from skimage import transform
from skimage.measure import find_contours
from src.config.endpoints import MEDSAM_CHECKPOINT, MEDSAM_DEVICE, MEDSAM_IMAGE_SIZE
from src.interfaces.segmenter import Segmenter
from src.schemas.segmentation import MaskResult
def _preprocess_image(image: Image.Image, target_size: int) -> tuple[torch.Tensor, int, int]:
"""Resize to target_size × target_size, normalize to [0,1], return tensor + original H, W."""
img_np = np.asarray(image.convert("RGB"), dtype=np.uint8)
H, W = img_np.shape[:2]
img_resized = transform.resize(
img_np,
(target_size, target_size),
order=3,
preserve_range=True,
anti_aliasing=True,
).astype(np.float32)
img_normalized = (img_resized - img_resized.min()) / np.clip(
img_resized.max() - img_resized.min(), a_min=1e-8, a_max=None
)
# (H, W, 3) → (1, 3, H, W)
tensor = torch.from_numpy(img_normalized).permute(2, 0, 1).unsqueeze(0).float()
return tensor, H, W
def _bbox_normalized_to_medsam(
bbox_normalized: list[int],
H: int,
W: int,
target_size: int,
) -> np.ndarray:
"""Convert MedGemma [y0, x0, y1, x1] in [0,1000] to MedSAM [x0,y0,x1,y1] in target_size scale."""
y0_n, x0_n, y1_n, x1_n = bbox_normalized
# to pixel coords in original image
x0_px = x0_n / 1000.0 * W
y0_px = y0_n / 1000.0 * H
x1_px = x1_n / 1000.0 * W
y1_px = y1_n / 1000.0 * H
# scale to target_size
box_scaled = np.array([[x0_px, y0_px, x1_px, y1_px]]) / np.array([W, H, W, H]) * target_size
return box_scaled
def _extract_polygon(mask: np.ndarray) -> list[list[int]] | None:
"""Return the largest contour of a binary mask as [[x, y], ...] pixel coords.
find_contours returns (row, col) = (y, x), so we swap to [x, y].
"""
contours = find_contours(mask.astype(float), level=0.5)
if not contours:
return None
largest = max(contours, key=len)
return [[int(round(pt[1])), int(round(pt[0]))] for pt in largest]
@torch.no_grad()
def _medsam_inference(
model: object,
img_embed: torch.Tensor,
box_scaled: np.ndarray,
H: int,
W: int,
) -> np.ndarray:
"""Run MedSAM mask decoder and return a binary H×W uint8 mask."""
box_torch = torch.as_tensor(box_scaled, dtype=torch.float, device=img_embed.device)
if box_torch.dim() == 2:
box_torch = box_torch[:, None, :] # (B, 1, 4)
sparse_emb, dense_emb = model.prompt_encoder(points=None, boxes=box_torch, masks=None)
low_res_logits, _ = model.mask_decoder(
image_embeddings=img_embed,
image_pe=model.prompt_encoder.get_dense_pe(),
sparse_prompt_embeddings=sparse_emb,
dense_prompt_embeddings=dense_emb,
multimask_output=False,
)
low_res_pred = F.interpolate(
torch.sigmoid(low_res_logits),
size=(H, W),
mode="bilinear",
align_corners=False,
)
mask = (low_res_pred.squeeze().cpu().numpy() > 0.5).astype(np.uint8)
return mask
class MedSAMClient(Segmenter):
"""Wraps MedSAM (ViT-B) for prompt-guided segmentation (Stage 2)."""
def __init__(
self,
checkpoint_path: str = MEDSAM_CHECKPOINT,
device: str = MEDSAM_DEVICE,
) -> None:
from segment_anything import sam_model_registry # import guarded for GPU-only envs
self.device = device
self._model = sam_model_registry["vit_b"](checkpoint=checkpoint_path)
self._model = self._model.to(device)
self._model.eval()
def segment(
self,
image: Image.Image,
bbox_normalized: list[int],
finding_label: str,
image_path: str,
image_index: int,
) -> MaskResult:
try:
img_tensor, H, W = _preprocess_image(image, MEDSAM_IMAGE_SIZE)
img_tensor = img_tensor.to(self.device)
with torch.no_grad():
img_embed = self._model.image_encoder(img_tensor)
box_scaled = _bbox_normalized_to_medsam(bbox_normalized, H, W, MEDSAM_IMAGE_SIZE)
mask = _medsam_inference(self._model, img_embed, box_scaled, H, W)
polygon = _extract_polygon(mask)
return MaskResult(
finding_label=finding_label,
image_path=image_path,
image_index=image_index,
bbox_normalized=bbox_normalized,
mask=mask,
polygon=polygon,
status="success",
)
except Exception as exc:
return MaskResult(
finding_label=finding_label,
image_path=image_path,
image_index=image_index,
bbox_normalized=bbox_normalized,
mask=None,
polygon=None,
status="failed",
error=str(exc),
)