File size: 7,452 Bytes
f5d3bf3
0cda03f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f5d3bf3
0cda03f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f5d3bf3
0cda03f
f5d3bf3
0cda03f
f5d3bf3
0cda03f
f5d3bf3
0cda03f
 
 
 
 
 
 
 
 
 
 
 
 
 
f5d3bf3
0cda03f
 
f5d3bf3
0cda03f
 
f5d3bf3
 
0cda03f
 
 
 
 
 
 
 
f5d3bf3
0cda03f
f5d3bf3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0cda03f
f5d3bf3
314ac36
 
f5d3bf3
 
5d4f88d
e815955
 
 
f5d3bf3
 
314ac36
 
 
f5d3bf3
0cda03f
 
 
 
 
f5d3bf3
 
 
0cda03f
 
 
 
 
f5d3bf3
 
 
 
 
314ac36
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f5d3bf3
 
314ac36
 
 
 
 
 
0cda03f
 
 
 
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
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
"""Lazy FLUX.2 multi-reference editor for GPU-backed Hugging Face Spaces."""

from __future__ import annotations

import os
import threading

import numpy as np
from PIL import Image, ImageFilter

try:
    import spaces

    gpu_task = spaces.GPU(duration=300)
except ImportError:
    # The `spaces` package is injected by the ZeroGPU runtime. Local and
    # classical-GPU environments should keep working without it.
    def gpu_task(function):
        return function


MODEL_ID = os.getenv("EDIT_MODEL_ID", "black-forest-labs/FLUX.2-klein-4B")
_PIPE = None
_LOCK = threading.Lock()


def _model_roi(image: Image.Image, mask: Image.Image, max_side: int = 1280):
    """Crop to mask context and resize to model-friendly multiples of 32."""
    bbox = mask.getbbox()
    if bbox is None:
        return None
    left, top, right, bottom = bbox
    margin = max(64, round(max(right - left, bottom - top) * 0.22))
    crop_box = (
        max(0, left - margin),
        max(0, top - margin),
        min(image.width, right + margin),
        min(image.height, bottom + margin),
    )
    crop_image = image.crop(crop_box).convert("RGB")
    crop_mask = mask.crop(crop_box).convert("L")
    scale = min(1.0, max_side / max(crop_image.size))
    width = max(64, int(np.ceil(crop_image.width * scale / 32)) * 32)
    height = max(64, int(np.ceil(crop_image.height * scale / 32)) * 32)
    model_image = crop_image.resize((width, height), Image.Resampling.LANCZOS)
    model_mask = crop_mask.resize((width, height), Image.Resampling.NEAREST)
    return crop_box, crop_image, crop_mask, model_image, model_mask


def _load_pipeline():
    global _PIPE
    if _PIPE is not None:
        return _PIPE
    with _LOCK:
        if _PIPE is not None:
            return _PIPE
        try:
            import torch
            from diffusers import Flux2KleinPipeline
        except ImportError as exc:
            raise RuntimeError("FLUX.2 editing dependencies are not installed.") from exc
        if not torch.cuda.is_available():
            raise RuntimeError("AI replacement requires a CUDA GPU. Select ZeroGPU or GPU hardware for the Space.")

        pipe = Flux2KleinPipeline.from_pretrained(
            MODEL_ID,
            torch_dtype=torch.bfloat16,
            low_cpu_mem_usage=True,
        )
        if os.getenv("LOW_VRAM", "0") == "1":
            pipe.enable_model_cpu_offload()
        else:
            pipe.to("cuda")
        pipe.set_progress_bar_config(disable=True)
        _PIPE = pipe
    return _PIPE


@gpu_task
def replace_object(
    image: Image.Image,
    mask: Image.Image,
    product_reference: Image.Image,
    prompt: str,
    seed: int,
    steps: int = 4,
    strength: float = 1.0,
) -> Image.Image:
    if mask.getbbox() is None:
        return image
    import torch

    roi = _model_roi(image, mask)
    if roi is None:
        return image
    crop_box, crop_image, crop_mask, model_image, _ = roi
    pipe = _load_pipeline()
    product_reference = product_reference.convert("RGBA")
    ref_bbox = product_reference.getchannel("A").getbbox()
    if ref_bbox:
        product_reference = product_reference.crop(ref_bbox)
    max_ref = 768
    ref_scale = min(1.0, max_ref / max(product_reference.size))
    ref_size = (
        max(1, round(product_reference.width * ref_scale)),
        max(1, round(product_reference.height * ref_scale)),
    )
    product_reference = product_reference.resize(ref_size, Image.Resampling.LANCZOS)
    pad = max(32, round(max(ref_size) * 0.12))
    product_card = Image.new(
        "RGBA",
        (product_reference.width + pad * 2, product_reference.height + pad * 2),
        (128, 128, 128, 255),
    )
    product_card.alpha_composite(product_reference, (pad, pad))

    full_prompt = (
        "Image 1 is the destination scene with an approximate replacement already positioned. "
        "Image 2 is the exact identity and appearance reference for the replacement object; "
        "its neutral isolation background is disposable and must never appear in the result. "
        "Create a photorealistic edit of image 1: replace the previous object with the object "
        "from image 2 at the approximate size, position, orientation and support contact shown "
        "in image 1. Do not enlarge the replacement beyond the guide shown in image 1. Infer the "
        "existing scene interaction only from visible evidence in image 1. Preserve real foreground "
        "occluders and existing contact geometry, but never invent hands, holders, people, stands, "
        "pedestals or supports that are not already present. Remove all remnants of the previous "
        "object. Preserve the replacement's distinctive shape, proportions, colors, material, "
        "label and text. Create physically coherent perspective, lighting, contact shadow, "
        "reflection and foreground occlusion. Preserve all unrelated background and scene pixels "
        "from image 1. Do not add a frame, card, rectangle, halo, backdrop, duplicate product or "
        "extra label around the replacement. "
        f"User direction: {prompt.strip() or 'natural context-aware object replacement'}"
    )
    generator = torch.Generator(device="cuda").manual_seed(int(seed))
    with _LOCK, torch.inference_mode():
        result = pipe(
            prompt=full_prompt,
            image=[model_image, product_card.convert("RGB")],
            height=model_image.height,
            width=model_image.width,
            num_inference_steps=int(steps),
            generator=generator,
        ).images[0].convert("RGB")

    generated_crop = result.resize(crop_image.size, Image.Resampling.LANCZOS)
    generated_crop = Image.blend(
        crop_image,
        generated_crop,
        float(np.clip(strength, 0.35, 1.0)),
    )
    # Match the generated crop to the untouched scene at the inner mask border.
    # This removes exposure/color shifts that otherwise reveal a rectangular or
    # brush-shaped patch around the replacement.
    mask_array = np.asarray(crop_mask, dtype=np.uint8) > 127
    mask_bbox = crop_mask.getbbox()
    span = max(mask_bbox[2] - mask_bbox[0], mask_bbox[3] - mask_bbox[1])
    erosion_size = max(3, min(21, round(span * 0.035)))
    if erosion_size % 2 == 0:
        erosion_size += 1
    eroded = np.asarray(crop_mask.filter(ImageFilter.MinFilter(erosion_size))) > 127
    inner_border = mask_array & ~eroded
    if np.count_nonzero(inner_border) >= 16:
        original_array = np.asarray(crop_image, dtype=np.float32)
        generated_array = np.asarray(generated_crop, dtype=np.float32)
        correction = np.median(
            original_array[inner_border] - generated_array[inner_border], axis=0
        )
        generated_array = np.clip(generated_array + correction, 0, 255).astype(np.uint8)
        generated_crop = Image.fromarray(generated_array, "RGB")
    # The edit model is free to reason globally inside its ROI, but only user-mask
    # pixels are committed to the destination scene.
    feather = float(np.clip(span * 0.025, 2.0, 12.0))
    soft_inside = crop_mask.filter(ImageFilter.GaussianBlur(radius=feather))
    blend_mask = Image.fromarray(
        np.minimum(np.asarray(crop_mask), np.asarray(soft_inside)).astype(np.uint8),
        "L",
    )
    merged_crop = Image.composite(generated_crop, crop_image, blend_mask)
    output = image.copy().convert("RGB")
    output.paste(merged_crop, crop_box[:2])
    return output