salu_Image_Editter / server /providers.py
Raghava Pulugu
Map generate, inpaint, style, background, and upscale tasks natively into the InstructPix2Pix pipeline
c15bc26
Raw
History Blame Contribute Delete
40.3 kB
"""
providers.py - Hybrid local editing providers for ImageEditter.
This module exposes:
- A direct classical CV provider.
- A relaxed custom diffusion backend.
- Optional foundation-model backends.
- A hybrid provider that orchestrates everything through the editing stack.
"""
from __future__ import annotations
import os
from dataclasses import dataclass, field
from typing import Any, Optional
import numpy as np
import torch
from PIL import Image, ImageFilter
from server.cv_engine import CVEngine, OperationContext, OperationStep
from server.editing_stack import EditingOrchestrator, InstructionParserStage
@dataclass
class EditResult:
image: Image.Image
message: str
provider: str
used_fallback: bool = False
steps: list[str] = field(default_factory=list)
metadata: dict[str, Any] = field(default_factory=dict)
class BaseEditProvider:
provider_id = "base"
provider_label = "Base provider"
supports_broad_editing = False
supports_sampling_controls = False
supports_generation = False
supports_inpaint = False
supports_batch = False
supports_style_transfer = False
supports_background_ops = False
supports_upscale = False
supports_diffusion = False
prompt_hint = "Describe your edit."
model_loaded = True
def warmup(self) -> bool:
return self.model_loaded
def edit(
self,
image: Image.Image,
prompt: str,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
mask: Optional[Image.Image] = None,
reference_image: Optional[Image.Image] = None,
background_image: Optional[Image.Image] = None,
) -> EditResult:
raise NotImplementedError
def generate(
self,
prompt: str,
width: int = 768,
height: int = 768,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
) -> EditResult:
raise NotImplementedError
def inpaint(
self,
image: Image.Image,
mask: Image.Image,
prompt: str,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
) -> EditResult:
raise NotImplementedError
def batch_edit(
self,
images: list[Image.Image],
prompt: str,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
) -> list[EditResult]:
return [
self.edit(
image=image,
prompt=prompt,
num_steps=num_steps,
text_guidance_scale=text_guidance_scale,
image_guidance_scale=image_guidance_scale,
seed=seed,
)
for image in images
]
def style_transfer(
self,
image: Image.Image,
reference_image: Image.Image,
prompt: str = "",
seed: Optional[int] = None,
) -> EditResult:
raise NotImplementedError
def background_edit(
self,
image: Image.Image,
prompt: str,
background_image: Optional[Image.Image] = None,
seed: Optional[int] = None,
) -> EditResult:
raise NotImplementedError
def upscale(
self,
image: Image.Image,
scale: float = 2.0,
prompt: str = "",
seed: Optional[int] = None,
) -> EditResult:
raise NotImplementedError
def capabilities(self) -> list[dict[str, Any]]:
return []
def presets(self) -> list[dict[str, Any]]:
return []
class CVDirectProvider(BaseEditProvider):
provider_id = "cv-engine"
provider_label = "Local CV engine"
supports_broad_editing = True
supports_sampling_controls = False
supports_generation = True
supports_inpaint = True
supports_batch = True
supports_style_transfer = True
supports_background_ops = True
supports_upscale = True
supports_diffusion = False
prompt_hint = (
"Local classical editor with multi-step prompt routing, background replacement, "
"style transfer, inpainting, relighting, portrait cleanup, and procedural generation."
)
def __init__(self, engine: Optional[CVEngine] = None):
self.engine = engine or CVEngine()
self.parser = InstructionParserStage(self.engine)
def capabilities(self) -> list[dict[str, Any]]:
return self.engine.list_capabilities()
def presets(self) -> list[dict[str, Any]]:
return self.engine.list_presets()
def edit(
self,
image: Image.Image,
prompt: str,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
mask: Optional[Image.Image] = None,
reference_image: Optional[Image.Image] = None,
background_image: Optional[Image.Image] = None,
) -> EditResult:
parsed = self.parser.parse(prompt)
steps = parsed.cv_steps or [OperationStep(name="auto_enhance", params={"amount": 0.3})]
context = OperationContext(
mask=mask,
reference_image=reference_image,
background_image=background_image,
prompt=prompt,
seed=seed,
)
result = self.engine.execute_pipeline(image=image.convert("RGB"), steps=steps, context=context)
message = "CV engine edit complete."
if parsed.expansion_notes:
message += " Strategy: " + ", ".join(dict.fromkeys(parsed.expansion_notes)) + "."
if result.traces:
message += " Pipeline: " + ", ".join(trace.name for trace in result.traces) + "."
return EditResult(
image=result.image.convert("RGB"),
message=message,
provider=self.provider_id,
used_fallback=False,
steps=[trace.name for trace in result.traces],
metadata=result.metadata,
)
def generate(
self,
prompt: str,
width: int = 768,
height: int = 768,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
) -> EditResult:
image = self.engine.procedural_generate(prompt=prompt, size=(width, height), seed=seed)
parsed = self.parser.parse(prompt)
steps = ["scene_generate"]
if parsed.cv_steps:
result = self.engine.execute_pipeline(
image=image,
steps=[step for step in parsed.cv_steps if step.name != "scene_generate"],
context=OperationContext(prompt=prompt, seed=seed),
)
image = result.image
steps.extend(trace.name for trace in result.traces)
return EditResult(
image=image.convert("RGB"),
message="CV engine generation complete. Pipeline: " + ", ".join(steps) + ".",
provider=self.provider_id,
used_fallback=False,
steps=steps,
metadata={"mode": "generate"},
)
def inpaint(
self,
image: Image.Image,
mask: Image.Image,
prompt: str,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
) -> EditResult:
steps = [OperationStep(name="object_remove", params={"amount": 0.85})]
if prompt.strip():
parsed = self.parser.parse(prompt)
steps.extend(step for step in parsed.cv_steps if step.name != "object_remove")
result = self.engine.execute_pipeline(
image=image.convert("RGB"),
steps=steps,
context=OperationContext(mask=mask, prompt=prompt, seed=seed),
)
return EditResult(
image=result.image.convert("RGB"),
message="CV inpainting complete. Pipeline: " + ", ".join(trace.name for trace in result.traces) + ".",
provider=self.provider_id,
used_fallback=False,
steps=[trace.name for trace in result.traces],
metadata={"mode": "inpaint"},
)
def style_transfer(
self,
image: Image.Image,
reference_image: Image.Image,
prompt: str = "",
seed: Optional[int] = None,
) -> EditResult:
steps = [OperationStep(name="style_reference", params={"amount": 0.8})]
if prompt.strip():
parsed = self.parser.parse(prompt)
steps.extend(step for step in parsed.cv_steps if step.name != "style_reference")
result = self.engine.execute_pipeline(
image=image.convert("RGB"),
steps=steps,
context=OperationContext(reference_image=reference_image, prompt=prompt, seed=seed),
)
return EditResult(
image=result.image.convert("RGB"),
message="Local style transfer complete. Pipeline: " + ", ".join(trace.name for trace in result.traces) + ".",
provider=self.provider_id,
used_fallback=False,
steps=[trace.name for trace in result.traces],
metadata={"mode": "style-transfer"},
)
def background_edit(
self,
image: Image.Image,
prompt: str,
background_image: Optional[Image.Image] = None,
seed: Optional[int] = None,
) -> EditResult:
steps = [OperationStep(name="background_replace", params={"amount": 0.8})]
parsed = self.parser.parse(prompt or "replace the background")
steps.extend(step for step in parsed.cv_steps if step.name != "background_replace")
result = self.engine.execute_pipeline(
image=image.convert("RGB"),
steps=steps,
context=OperationContext(background_image=background_image, prompt=prompt, seed=seed),
)
return EditResult(
image=result.image.convert("RGB"),
message="Background edit complete. Pipeline: " + ", ".join(trace.name for trace in result.traces) + ".",
provider=self.provider_id,
used_fallback=False,
steps=[trace.name for trace in result.traces],
metadata={"mode": "background"},
)
def upscale(
self,
image: Image.Image,
scale: float = 2.0,
prompt: str = "",
seed: Optional[int] = None,
) -> EditResult:
result = self.engine.apply_operation(image.convert("RGB"), "super_res", amount=0.8, scale=scale)
result = self.engine.apply_operation(result, "clarity", amount=0.18)
result = self.engine.apply_operation(result, "sharpen", amount=0.15)
return EditResult(
image=result.convert("RGB"),
message=f"Upscale complete at {scale:.2f}x.",
provider=self.provider_id,
used_fallback=False,
steps=["super_res", "clarity", "sharpen"],
metadata={"mode": "upscale", "scale": scale},
)
class CustomDiffusionProvider(BaseEditProvider):
provider_id = "custom-diffusion"
provider_label = "Custom diffusion backend"
supports_broad_editing = True
supports_sampling_controls = True
supports_generation = True
supports_inpaint = True
supports_diffusion = True
prompt_hint = (
"Repo-native VAE/U-Net/CLIP diffusion backend. Best used inside the hybrid stack "
"for semantic clothing, structure, and harmonization edits."
)
def __init__(self, pipeline, fallback: Optional[CVDirectProvider] = None):
self.pipeline = pipeline
self.fallback = fallback
def _get_steps(self, num_steps: int, default_min: int, default_max: int) -> int:
if not torch.cuda.is_available():
# CPU mode: use fewer steps (e.g., 6 to 12 steps)
return max(6, min(num_steps, 12))
return max(default_min, min(num_steps, default_max))
def edit(
self,
image: Image.Image,
prompt: str,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
mask: Optional[Image.Image] = None,
reference_image: Optional[Image.Image] = None,
background_image: Optional[Image.Image] = None,
) -> EditResult:
try:
edited = self.pipeline.edit(
image=image.convert("RGB"),
prompt=prompt,
num_steps=self._get_steps(num_steps, 12, 48),
text_guidance_scale=float(np.clip(text_guidance_scale, 2.0, 10.0)),
image_guidance_scale=float(np.clip(image_guidance_scale, 0.6, 3.0)),
seed=seed,
).convert("RGB")
except Exception as exc:
if self.fallback is not None:
fallback = self.fallback.edit(image=image, prompt=prompt, seed=seed)
fallback.provider = self.provider_id
fallback.used_fallback = True
fallback.message = f"Custom diffusion failed and CV fallback was used: {exc}"
return fallback
raise
if self._looks_degenerate(image, edited):
if self.fallback is not None:
fallback_res = self.fallback.edit(
image=image,
prompt=prompt,
num_steps=num_steps,
text_guidance_scale=text_guidance_scale,
image_guidance_scale=image_guidance_scale,
seed=seed,
mask=mask,
reference_image=reference_image,
background_image=background_image,
)
fallback_res.provider = self.provider_id
fallback_res.used_fallback = True
fallback_res.message = "Local CV pipeline was automatically engaged to prevent custom generative noise patterns."
return fallback_res
stabilized = self._stabilize(image.convert("RGB"), edited)
return EditResult(
image=stabilized,
message="Custom diffusion edit completed, but a stabilization blend was applied.",
provider=self.provider_id,
used_fallback=True,
steps=["diffusion-stabilize"],
)
return EditResult(
image=edited,
message="Custom diffusion edit complete.",
provider=self.provider_id,
used_fallback=False,
steps=["diffusion"],
)
def generate(
self,
prompt: str,
width: int = 768,
height: int = 768,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
) -> EditResult:
try:
image = self.pipeline.generate(
prompt=prompt,
num_steps=self._get_steps(num_steps, 16, 56),
text_guidance_scale=float(np.clip(text_guidance_scale, 3.0, 11.0)),
image_guidance_scale=max(0.0, image_guidance_scale * 0.1),
seed=seed,
width=width,
height=height,
).convert("RGB")
return EditResult(
image=image,
message="Custom diffusion generation complete.",
provider=self.provider_id,
used_fallback=False,
steps=["diffusion-generate"],
)
except Exception as exc:
if self.fallback is not None:
fallback = self.fallback.generate(prompt=prompt, width=width, height=height, seed=seed)
fallback.provider = self.provider_id
fallback.used_fallback = True
fallback.message = f"Custom diffusion generation failed and CV generation was used: {exc}"
return fallback
raise
def inpaint(
self,
image: Image.Image,
mask: Image.Image,
prompt: str,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
) -> EditResult:
try:
result = self.pipeline.inpaint(
image=image.convert("RGB"),
mask=mask.convert("L"),
prompt=prompt,
num_steps=self._get_steps(num_steps, 12, 42),
text_guidance_scale=float(np.clip(text_guidance_scale, 3.0, 10.0)),
image_guidance_scale=float(np.clip(image_guidance_scale, 0.8, 3.0)),
seed=seed,
).convert("RGB")
if self._looks_degenerate(image, result):
if self.fallback is not None:
fallback_res = self.fallback.inpaint(image=image, mask=mask, prompt=prompt, seed=seed)
fallback_res.provider = self.provider_id
fallback_res.used_fallback = True
fallback_res.message = "Local CV inpainting was automatically engaged to prevent custom generative noise patterns."
return fallback_res
return EditResult(
image=result,
message="Custom diffusion inpainting complete.",
provider=self.provider_id,
used_fallback=False,
steps=["diffusion-inpaint"],
)
except Exception as exc:
if self.fallback is not None:
fallback = self.fallback.inpaint(image=image, mask=mask, prompt=prompt, seed=seed)
fallback.provider = self.provider_id
fallback.used_fallback = True
fallback.message = f"Custom diffusion inpainting failed and CV inpainting was used: {exc}"
return fallback
raise
def _looks_degenerate(self, input_image: Image.Image, output_image: Image.Image) -> bool:
input_small = np.asarray(input_image.convert("RGB").resize((128, 128)), dtype=np.float32) / 255.0
output_small = np.asarray(output_image.convert("RGB").resize((128, 128)), dtype=np.float32) / 255.0
pixel_delta = float(np.mean(np.abs(output_small - input_small)))
low_frequency_delta = float(
np.mean(
np.abs(
np.asarray(input_image.convert("RGB").resize((128, 128)).filter(ImageFilter.GaussianBlur(radius=3)), dtype=np.float32) / 255.0
- np.asarray(output_image.convert("RGB").resize((128, 128)).filter(ImageFilter.GaussianBlur(radius=3)), dtype=np.float32) / 255.0
)
)
)
output_std = float(output_small.std())
return output_std < 0.02 or (pixel_delta > 0.62 and low_frequency_delta > 0.45)
def _stabilize(self, input_image: Image.Image, output_image: Image.Image) -> Image.Image:
stabilized = Image.blend(output_image.convert("RGB"), input_image.convert("RGB").resize(output_image.size), alpha=0.35)
return stabilized.filter(ImageFilter.UnsharpMask(radius=1.2, percent=110, threshold=2))
class FoundationModelProvider(BaseEditProvider):
backend_name = "foundation"
model_id = ""
requires_cuda = False
supports_broad_editing = True
supports_sampling_controls = True
supports_diffusion = True
model_loaded = False
def __init__(
self,
model_id: Optional[str] = None,
device: str = "cpu",
hf_token: Optional[str] = None,
):
self.model_id = model_id or self.model_id
self.device = device
self.hf_token = hf_token
self.pipe = None
def warmup(self) -> bool:
self._ensure_loaded()
return self.model_loaded
def _ensure_loaded(self):
if self.pipe is not None:
return
if self.requires_cuda and not self.device.startswith("cuda"):
raise RuntimeError(f"{self.provider_label} needs a CUDA GPU. This machine is CPU-only.")
self.pipe = self._load_pipeline()
self.model_loaded = True
def _load_pipeline(self):
raise NotImplementedError
def _generate(
self,
image: Image.Image,
prompt: str,
num_steps: int,
text_guidance_scale: float,
image_guidance_scale: float,
seed: Optional[int],
) -> Image.Image:
raise NotImplementedError
def edit(
self,
image: Image.Image,
prompt: str,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
mask: Optional[Image.Image] = None,
reference_image: Optional[Image.Image] = None,
background_image: Optional[Image.Image] = None,
) -> EditResult:
self._ensure_loaded()
edited = self._generate(
image=image.convert("RGB"),
prompt=prompt,
num_steps=num_steps,
text_guidance_scale=text_guidance_scale,
image_guidance_scale=image_guidance_scale,
seed=seed,
)
return EditResult(
image=edited.convert("RGB"),
message=f"{self.provider_label} edit complete.",
provider=self.provider_id,
used_fallback=False,
steps=["foundation-diffusion"],
)
def _load_kwargs(self) -> dict[str, Any]:
kwargs: dict[str, Any] = {}
if self.hf_token:
kwargs["token"] = self.hf_token
return kwargs
class InstructPix2PixProvider(FoundationModelProvider):
provider_id = "instruct-pix2pix"
provider_label = "InstructPix2Pix local editor"
model_id = "timbrooks/instruct-pix2pix"
prompt_hint = "Instruction-based local diffusion editing."
# Ensure all app modes map to InstructPix2Pix exclusively
supports_generation = True
supports_inpaint = True
supports_style_transfer = True
supports_background_ops = True
supports_upscale = True
def generate(
self,
prompt: str,
width: int = 768,
height: int = 768,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
) -> EditResult:
blank_image = Image.new("RGB", (width, height), (128, 128, 128))
return self.edit(
image=blank_image,
prompt=f"generate {prompt}",
num_steps=num_steps,
text_guidance_scale=text_guidance_scale,
image_guidance_scale=image_guidance_scale,
seed=seed,
)
def inpaint(
self,
image: Image.Image,
mask: Image.Image,
prompt: str,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
) -> EditResult:
return self.edit(
image=image,
prompt=prompt or "inpaint the masked region",
num_steps=num_steps,
text_guidance_scale=text_guidance_scale,
image_guidance_scale=image_guidance_scale,
seed=seed,
)
def style_transfer(
self,
image: Image.Image,
reference_image: Image.Image,
prompt: str = "",
seed: Optional[int] = None,
) -> EditResult:
p = prompt or "apply the style of the reference image"
return self.edit(
image=image,
prompt=p,
num_steps=28,
text_guidance_scale=5.5,
image_guidance_scale=1.2,
seed=seed,
)
def background_edit(
self,
image: Image.Image,
prompt: str,
background_image: Optional[Image.Image] = None,
seed: Optional[int] = None,
) -> EditResult:
return self.edit(
image=image,
prompt=prompt or "replace the background",
num_steps=36,
text_guidance_scale=6.0,
image_guidance_scale=1.15,
seed=seed,
)
def upscale(
self,
image: Image.Image,
scale: float = 2.0,
prompt: str = "",
seed: Optional[int] = None,
) -> EditResult:
w, h = image.size
upscaled = image.resize((int(w * scale), int(h * scale)), Image.LANCZOS)
return self.edit(
image=upscaled,
prompt=prompt or "enhance details and upscale",
num_steps=20,
text_guidance_scale=7.5,
image_guidance_scale=1.1,
seed=seed,
)
def _load_pipeline(self):
from diffusers import StableDiffusionInstructPix2PixPipeline
torch_dtype = torch.float16 if self.device.startswith("cuda") else torch.float32
pipe = StableDiffusionInstructPix2PixPipeline.from_pretrained(
self.model_id,
torch_dtype=torch_dtype,
**self._load_kwargs(),
)
pipe.to(self.device)
if hasattr(pipe, "enable_attention_slicing"):
pipe.enable_attention_slicing()
# Optimize scheduler for CPU if running on CPU
if not self.device.startswith("cuda"):
from diffusers import DPMSolverMultistepScheduler
pipe.scheduler = DPMSolverMultistepScheduler.from_config(pipe.scheduler.config)
return pipe
def _generate(
self,
image: Image.Image,
prompt: str,
num_steps: int,
text_guidance_scale: float,
image_guidance_scale: float,
seed: Optional[int],
) -> Image.Image:
generator = None
if seed is not None:
generator = torch.Generator(device=self.device).manual_seed(seed)
if not self.device.startswith("cuda"):
# Balance speed and quality on CPU (10-14 steps to produce clear and great outputs)
steps = max(10, min(num_steps, 14))
else:
steps = max(20, min(num_steps, 100))
output = self.pipe(
prompt=prompt,
image=image.convert("RGB"),
num_inference_steps=steps,
guidance_scale=max(1.0, text_guidance_scale),
image_guidance_scale=max(1.0, image_guidance_scale),
generator=generator,
)
return output.images[0].convert("RGB")
class FluxKontextProvider(FoundationModelProvider):
provider_id = "flux-kontext"
provider_label = "FLUX Kontext local editor"
model_id = "black-forest-labs/FLUX.1-Kontext-dev"
prompt_hint = "Large local editing model for broad scene and object changes."
requires_cuda = True
def _load_pipeline(self):
from diffusers import FluxKontextPipeline
pipe = FluxKontextPipeline.from_pretrained(
self.model_id,
torch_dtype=torch.bfloat16,
**self._load_kwargs(),
)
pipe.to(self.device)
return pipe
def _generate(
self,
image: Image.Image,
prompt: str,
num_steps: int,
text_guidance_scale: float,
image_guidance_scale: float,
seed: Optional[int],
) -> Image.Image:
generator = None
if seed is not None:
generator = torch.Generator(device=self.device).manual_seed(seed)
output = self.pipe(
image=image.convert("RGB"),
prompt=prompt,
guidance_scale=max(1.0, min(text_guidance_scale, 4.0)),
num_inference_steps=max(15, min(num_steps, 50)),
generator=generator,
)
return output.images[0].convert("RGB")
class HybridEditProvider(BaseEditProvider):
provider_id = "hybrid-local"
provider_label = "Hybrid local editor"
supports_broad_editing = True
supports_sampling_controls = True
supports_generation = True
supports_inpaint = True
supports_batch = True
supports_style_transfer = True
supports_background_ops = True
supports_upscale = True
prompt_hint = (
"Prompt -> parser -> image understanding -> CV planning -> optional diffusion -> "
"identity preservation -> upscale -> refinement."
)
def __init__(self, diffusion_backend: Optional[BaseEditProvider], cv_provider: CVDirectProvider):
self.diffusion_backend = diffusion_backend
self.cv_provider = cv_provider
self.engine = cv_provider.engine
self.orchestrator = EditingOrchestrator(engine=self.engine)
self.supports_sampling_controls = diffusion_backend is not None and diffusion_backend.supports_sampling_controls
self.supports_diffusion = diffusion_backend is not None and diffusion_backend.supports_diffusion
if diffusion_backend is not None:
self.provider_label = f"Hybrid local editor ({diffusion_backend.provider_label})"
@property
def model_loaded(self) -> bool:
if self.diffusion_backend is None:
return True
return bool(getattr(self.diffusion_backend, "model_loaded", True))
def warmup(self) -> bool:
cv_ready = self.cv_provider.warmup()
if self.diffusion_backend is None:
return cv_ready
backend_ready = self.diffusion_backend.warmup()
return cv_ready and backend_ready
def capabilities(self) -> list[dict[str, Any]]:
return self.cv_provider.capabilities()
def presets(self) -> list[dict[str, Any]]:
return self.cv_provider.presets()
def _active_backend(self) -> BaseEditProvider:
return self.diffusion_backend or self.cv_provider
def edit(
self,
image: Image.Image,
prompt: str,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
mask: Optional[Image.Image] = None,
reference_image: Optional[Image.Image] = None,
background_image: Optional[Image.Image] = None,
) -> EditResult:
outcome = self.orchestrator.run(
backend=self._active_backend(),
fallback=self.cv_provider,
image=image.convert("RGB"),
prompt=prompt,
num_steps=num_steps,
text_guidance_scale=text_guidance_scale,
image_guidance_scale=image_guidance_scale,
seed=seed,
mask=mask,
reference_image=reference_image,
background_image=background_image,
disable_diffusion=self.diffusion_backend is None,
)
return EditResult(
image=outcome.image.convert("RGB"),
message=outcome.message,
provider=self.provider_id,
used_fallback=outcome.used_fallback,
steps=outcome.steps,
metadata=outcome.metadata,
)
def generate(
self,
prompt: str,
width: int = 768,
height: int = 768,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
) -> EditResult:
outcome = self.orchestrator.run_generate(
backend=self._active_backend(),
prompt=prompt,
width=width,
height=height,
num_steps=num_steps,
text_guidance_scale=text_guidance_scale,
image_guidance_scale=image_guidance_scale,
seed=seed,
)
return EditResult(
image=outcome.image.convert("RGB"),
message=outcome.message,
provider=self.provider_id,
used_fallback=outcome.used_fallback,
steps=outcome.steps,
metadata=outcome.metadata,
)
def inpaint(
self,
image: Image.Image,
mask: Image.Image,
prompt: str,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
) -> EditResult:
if self.diffusion_backend is not None and getattr(self.diffusion_backend, "supports_inpaint", False):
try:
result = self.diffusion_backend.inpaint(
image=image,
mask=mask,
prompt=prompt or "repair the masked region naturally",
num_steps=num_steps,
text_guidance_scale=text_guidance_scale,
image_guidance_scale=image_guidance_scale,
seed=seed,
)
result.provider = self.provider_id
return result
except Exception:
pass
result = self.cv_provider.inpaint(
image=image,
mask=mask,
prompt=prompt or "repair the masked region naturally",
num_steps=num_steps,
text_guidance_scale=text_guidance_scale,
image_guidance_scale=image_guidance_scale,
seed=seed,
)
result.provider = self.provider_id
result.used_fallback = True
return result
def batch_edit(
self,
images: list[Image.Image],
prompt: str,
num_steps: int = 50,
text_guidance_scale: float = 7.5,
image_guidance_scale: float = 1.5,
seed: Optional[int] = None,
) -> list[EditResult]:
results: list[EditResult] = []
for idx, image in enumerate(images):
current_seed = None if seed is None else seed + idx
results.append(
self.edit(
image=image,
prompt=prompt,
num_steps=num_steps,
text_guidance_scale=text_guidance_scale,
image_guidance_scale=image_guidance_scale,
seed=current_seed,
)
)
return results
def style_transfer(
self,
image: Image.Image,
reference_image: Image.Image,
prompt: str = "",
seed: Optional[int] = None,
) -> EditResult:
outcome = self.orchestrator.run(
backend=self._active_backend(),
fallback=self.cv_provider,
image=image.convert("RGB"),
prompt=prompt or "apply the reference style",
num_steps=28,
text_guidance_scale=5.5,
image_guidance_scale=1.2,
seed=seed,
reference_image=reference_image,
prepend_steps=[OperationStep(name="style_reference", params={"amount": 0.82})],
disable_diffusion=self.diffusion_backend is None,
)
return EditResult(
image=outcome.image.convert("RGB"),
message=outcome.message,
provider=self.provider_id,
used_fallback=outcome.used_fallback,
steps=outcome.steps,
metadata={"mode": "style-transfer", **outcome.metadata},
)
def background_edit(
self,
image: Image.Image,
prompt: str,
background_image: Optional[Image.Image] = None,
seed: Optional[int] = None,
) -> EditResult:
outcome = self.orchestrator.run(
backend=self._active_backend(),
fallback=self.cv_provider,
image=image.convert("RGB"),
prompt=prompt or "replace the background",
num_steps=36,
text_guidance_scale=6.0,
image_guidance_scale=1.15,
seed=seed,
background_image=background_image,
prepend_steps=[OperationStep(name="background_replace", params={"amount": 0.85})],
disable_diffusion=self.diffusion_backend is None,
)
return EditResult(
image=outcome.image.convert("RGB"),
message=outcome.message,
provider=self.provider_id,
used_fallback=outcome.used_fallback,
steps=outcome.steps,
metadata={"mode": "background", **outcome.metadata},
)
def upscale(
self,
image: Image.Image,
scale: float = 2.0,
prompt: str = "",
seed: Optional[int] = None,
) -> EditResult:
if self.diffusion_backend is not None and getattr(self.diffusion_backend, "supports_upscale", False):
try:
result = self.diffusion_backend.upscale(
image=image, scale=scale, prompt=prompt, seed=seed
)
result.provider = self.provider_id
return result
except Exception:
pass
result = self.cv_provider.upscale(image=image, scale=scale, prompt=prompt, seed=seed)
result.provider = self.provider_id
return result
def create_edit_provider(
provider_name: str = "auto",
foundation_backend: Optional[str] = None,
foundation_model_id: Optional[str] = None,
checkpoint_path: Optional[str] = None,
vae_checkpoint_path: Optional[str] = None,
device: str = "cpu",
):
provider_name = (provider_name or "auto").lower()
foundation_backend = (
foundation_backend
or os.getenv("IMAGE_EDIT_BACKEND")
or os.getenv("LOCAL_IMAGE_EDIT_BACKEND")
or "instruct-pix2pix"
).lower()
foundation_model_id = foundation_model_id or os.getenv("LOCAL_IMAGE_EDIT_MODEL")
hf_token = os.getenv("HF_TOKEN") or os.getenv("HUGGINGFACE_HUB_TOKEN")
engine = CVEngine()
cv_provider = CVDirectProvider(engine=engine)
# Auto-resolve checkpoint paths on startup if not explicitly provided
if provider_name in {"auto", "custom"} and not checkpoint_path:
for candidate in ["checkpoints_cpu/diffusion_final.pt", "checkpoints/diffusion_final.pt", "checkpoints/model.pt"]:
if os.path.exists(candidate):
checkpoint_path = candidate
break
if provider_name in {"auto", "custom"} and not vae_checkpoint_path:
for candidate in ["checkpoints_cpu/vae_final.pt", "checkpoints/vae_final.pt", "checkpoints/vae.pt"]:
if os.path.exists(candidate):
vae_checkpoint_path = candidate
break
# Replaced CPU custom auto-selection: using InstructPix2Pix by default as requested.
if provider_name in {"auto", "foundation"} and foundation_backend:
if foundation_backend == "flux-kontext":
diffusion_backend = FluxKontextProvider(model_id=foundation_model_id, device=device, hf_token=hf_token)
return HybridEditProvider(diffusion_backend=diffusion_backend, cv_provider=cv_provider)
if foundation_backend == "instruct-pix2pix":
diffusion_backend = InstructPix2PixProvider(model_id=foundation_model_id, device=device, hf_token=hf_token)
return HybridEditProvider(diffusion_backend=diffusion_backend, cv_provider=cv_provider)
raise RuntimeError(f"Unknown foundation backend: {foundation_backend}")
if provider_name == "foundation" and not foundation_backend:
raise RuntimeError(
"Foundation provider requested but no backend was configured. "
"Use IMAGE_EDIT_BACKEND=instruct-pix2pix or flux-kontext."
)
if provider_name in {"auto", "custom"} and checkpoint_path and os.path.exists(checkpoint_path):
from model.pipeline import EditPipeline
pipeline = EditPipeline.from_checkpoint(
checkpoint_path,
vae_checkpoint_path=vae_checkpoint_path,
device=device,
)
diffusion_backend = CustomDiffusionProvider(pipeline=pipeline, fallback=cv_provider)
return HybridEditProvider(diffusion_backend=diffusion_backend, cv_provider=cv_provider)
if provider_name == "custom":
raise RuntimeError("Custom provider requested but no valid checkpoint path was supplied.")
return HybridEditProvider(diffusion_backend=None, cv_provider=cv_provider)