File size: 2,224 Bytes
8966f37
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
badc3e1
8966f37
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bfeb81e
8966f37
 
bfeb81e
 
 
8966f37
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
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())}