import numpy as np import torch import torch.nn.functional as F from PIL import Image from transformers.pipelines.base import Pipeline from .image_processing_cond_unet import CondUNetImageProcessor class CondUNetImageSegmentationPipeline(Pipeline): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) if self.image_processor is None: self.image_processor = CondUNetImageProcessor( image_size=self.model.config.image_size, keep_aspect_ratio=self.model.config.keep_aspect_ratio, self_normalize=self.model.config.self_normalize, ) def _sanitize_parameters(self, organ_id=None, threshold=None, **kwargs): preprocess_kwargs = {} postprocess_kwargs = {} if organ_id is not None: preprocess_kwargs["organ_id"] = organ_id if threshold is not None: postprocess_kwargs["threshold"] = threshold return preprocess_kwargs, {}, postprocess_kwargs def preprocess(self, image, organ_id=None, **kwargs): if not isinstance(image, Image.Image): image = Image.open(image).convert("RGB") else: image = image.convert("RGB") width, height = image.size inputs = self.image_processor(images=image, return_tensors="pt") inputs["original_size"] = (height, width) if organ_id is not None: inputs["organ_id"] = torch.tensor([organ_id], dtype=torch.long) return inputs def _forward(self, model_inputs, **kwargs): original_size = model_inputs.pop("original_size") outputs = self.model(**model_inputs) return {"logits": outputs.logits, "original_size": original_size} def postprocess(self, model_outputs, threshold=0.7, **kwargs): logits = model_outputs["logits"] height, width = model_outputs["original_size"] probabilities = torch.sigmoid( F.interpolate(logits, size=(height, width), mode="nearest") )[0, 0] mask = (probabilities >= threshold).to(torch.uint8).cpu().numpy() * 255 return {"label": "foreground", "mask": Image.fromarray(mask), "score": float(probabilities.mean())}