File size: 4,352 Bytes
f66bbd0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
"""Gaze-mask-guided image correction with an image-to-image diffusion model.

The segmentation mask is deliberately used only to localise corruption: pixels
outside the attended region are preserved in the source image before the
diffusion model is invoked.  This is a research workflow, not a diagnostic or
clinical image reconstruction method.
"""
from __future__ import annotations

from functools import lru_cache
import os

import numpy as np
import torch
from PIL import Image, ImageFilter


ROENTGEN_MODEL = "stanfordmimi/RoentGen-v2"


def make_attention_noise(
    image: Image.Image,
    mask: Image.Image | np.ndarray,
    noise_strength: float = 0.35,
    blur_radius: float = 4.0,
    seed: int | None = None,
) -> tuple[Image.Image, Image.Image]:
    """Return an image with Gaussian noise inside a feathered attention mask.

    ``noise_strength`` is the standard deviation on the 0--255 image scale.
    The returned mask is the feathered mask actually used for compositing.
    """
    source = np.asarray(image.convert("RGB"), dtype=np.float32)
    if isinstance(mask, Image.Image):
        mask_image = mask.convert("L")
    else:
        values = np.asarray(mask, dtype=np.float32)
        if values.max(initial=0.0) <= 1.0:
            values = values * 255.0
        mask_image = Image.fromarray(np.clip(values, 0, 255).astype(np.uint8), "L")
    mask_image = mask_image.resize(image.size, Image.Resampling.BILINEAR)
    if blur_radius > 0:
        mask_image = mask_image.filter(ImageFilter.GaussianBlur(radius=blur_radius))

    alpha = np.asarray(mask_image, dtype=np.float32)[..., None] / 255.0
    rng = np.random.default_rng(seed)
    noise = rng.normal(0.0, 255.0 * float(noise_strength), size=source.shape)
    noised = np.clip(source + alpha * noise, 0, 255).astype(np.uint8)
    return Image.fromarray(noised, "RGB"), mask_image


def _device_and_dtype() -> tuple[str, torch.dtype]:
    if torch.cuda.is_available():
        return "cuda", torch.float16
    if getattr(torch.backends, "mps", None) and torch.backends.mps.is_available():
        return "mps", torch.float32
    return "cpu", torch.float32


@lru_cache(maxsize=1)
def _load_roentgen_pipeline():
    """Load lazily so segmentation remains usable without diffusers installed."""
    try:
        from diffusers import StableDiffusionImg2ImgPipeline
    except ImportError as exc:
        raise RuntimeError(
            "Diffusion correction requires `diffusers`, `transformers`, and "
            "`accelerate`. Install the project's updated requirements first."
        ) from exc

    device, dtype = _device_and_dtype()
    token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACEHUB_API_TOKEN")
    # RoentGen-v2 is a Stable Diffusion 2.1 fine-tune.  Its published example
    # shows text-to-image with DiffusionPipeline; this image-to-image variant
    # loads the same weights but accepts the noised attention image as `image`.
    pipe = StableDiffusionImg2ImgPipeline.from_pretrained(
        ROENTGEN_MODEL, torch_dtype=dtype, token=token
    )
    return pipe.to(device)


def correct_with_roentgen(
    noised_image: Image.Image,
    prompt: str = "chest x-ray",
    strength: float = 0.35,
    guidance_scale: float = 3.5,
    steps: int = 30,
    seed: int | None = None,
) -> Image.Image:
    """Use RoentGen-v2 image-to-image inference to restore a noised region."""
    if not 0.0 < strength <= 1.0:
        raise ValueError("Diffusion strength must be in (0, 1].")
    pipe = _load_roentgen_pipeline()
    device, _ = _device_and_dtype()
    generator = None if seed is None else torch.Generator(device=device).manual_seed(int(seed))
    result = pipe(
        prompt=prompt,
        image=noised_image.convert("RGB"),
        strength=float(strength),
        guidance_scale=float(guidance_scale),
        num_inference_steps=int(steps),
        generator=generator,
    )
    return result.images[0].convert("RGB")


def apply_correction_only_to_attention(
    source_image: Image.Image, generated_image: Image.Image, feathered_mask: Image.Image
) -> Image.Image:
    """Keep un-attended anatomy exactly from the input image after diffusion."""
    generated = generated_image.convert("RGB").resize(source_image.size, Image.Resampling.LANCZOS)
    return Image.composite(generated, source_image.convert("RGB"), feathered_mask)