US_Cond-UNet / pipeline.py
Morelli001's picture
Upload folder using huggingface_hub
badc3e1 verified
Raw
History Blame Contribute Delete
2.22 kB
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())}