ArchEnhancer / backend /services /diffusion_pipeline.py
Aguilar Elizondo
Fix: Correct relative imports in diffusion_pipeline.py
544d65f
Raw
History Blame Contribute Delete
11.4 kB
"""
Diffusion Pipeline Service
Orchestrates the complete image enhancement pipeline:
1. Load and preprocess input image
2. Run Stable Diffusion img2img with optional LoRA
3. Apply upscaling
4. Apply post-processing effects
"""
import logging
from pathlib import Path
from typing import Optional, Callable
import torch
from PIL import Image
from diffusers import (
StableDiffusionImg2ImgPipeline,
AutoencoderKL,
DPMSolverMultistepScheduler
)
from backend.config import settings, get_lora_path
from backend.services.upscaler import upscale_image
from backend.services.postprocess import postprocess_image
logger = logging.getLogger(__name__)
class DiffusionPipelineManager:
"""
Manages the Stable Diffusion pipeline for image enhancement
This class handles lazy loading of models and provides a clean
interface for image-to-image enhancement.
"""
def __init__(self):
self.pipe = None
self.lora_loaded = False
self.device = "cuda" if torch.cuda.is_available() else "cpu"
logger.info(f"Pipeline will use device: {self.device}")
logger.info(f"CUDA available: {torch.cuda.is_available()}")
if torch.cuda.is_available():
logger.info(f"CUDA device: {torch.cuda.get_device_name(0)}")
def load_pipeline(self):
"""
Load the Stable Diffusion img2img pipeline with optimizations
This method initializes the Stable Diffusion pipeline with:
- Mixed precision for faster inference
- Memory-efficient attention
- Optional VAE for better image quality
"""
if self.pipe is not None:
return
logger.info("Loading Stable Diffusion 1.5 pipeline...")
try:
# Determine dtype based on device
dtype = torch.float32 # Always use float32 for CPU compatibility
logger.info(f"Loading model with dtype: {dtype}, device: {self.device}")
# Initialize the pipeline (SD 1.5) - let diffusers handle all components
self.pipe = StableDiffusionImg2ImgPipeline.from_pretrained(
settings.BASE_MODEL,
torch_dtype=dtype,
safety_checker=None,
requires_safety_checker=False,
low_cpu_mem_usage=False # Load everything properly
)
# Optimize scheduler
self.pipe.scheduler = DPMSolverMultistepScheduler.from_config(
self.pipe.scheduler.config
)
# Move to device
self.pipe = self.pipe.to(self.device)
# Enable memory optimizations
if self.device == "cuda":
if settings.ENABLE_ATTENTION_SLICING:
self.pipe.enable_attention_slicing()
logger.info("Enabled attention slicing")
if settings.ENABLE_VAE_SLICING:
self.pipe.enable_vae_slicing()
logger.info("Enabled VAE slicing")
logger.info("Pipeline loaded successfully")
logger.info(f"Model: {settings.BASE_MODEL}")
logger.info(f"Device: {self.device}")
except Exception as e:
logger.error(f"Failed to load pipeline: {e}", exc_info=True)
raise
def load_lora(self, lora_path: Path):
"""
Load a custom LoRA adapter into the pipeline
Args:
lora_path: Path to the LoRA weights file (.safetensors)
"""
if self.pipe is None:
self.load_pipeline()
try:
logger.info(f"Loading LoRA from: {lora_path}")
self.pipe.load_lora_weights(str(lora_path.parent), weight_name=lora_path.name)
self.lora_loaded = True
logger.info("LoRA loaded successfully")
except Exception as e:
logger.warning(f"Failed to load LoRA (continuing without it): {e}")
self.lora_loaded = False
def enhance(
self,
image: Image.Image,
prompt: str,
negative_prompt: str,
strength: float,
guidance_scale: float,
num_inference_steps: int = 30,
progress_callback: Optional[Callable[[int, str], None]] = None
) -> Image.Image:
"""
Run image-to-image enhancement
Args:
image: Input PIL Image
prompt: Positive prompt
negative_prompt: Negative prompt
strength: Denoising strength (0-1)
guidance_scale: CFG scale
num_inference_steps: Number of diffusion steps
progress_callback: Optional callback for progress updates
Returns:
Enhanced PIL Image
"""
if self.pipe is None:
if progress_callback:
progress_callback(30, "Loading AI model...")
self.load_pipeline()
if progress_callback:
progress_callback(40, "Preprocessing image...")
logger.info(f"Running enhancement - Strength: {strength}, Guidance: {guidance_scale}")
# Ensure image is in RGB mode
if image.mode != "RGB":
image = image.convert("RGB")
# Resize if too large
max_size = settings.MAX_IMAGE_SIZE
if max(image.size) > max_size:
ratio = max_size / max(image.size)
new_size = tuple(int(dim * ratio) for dim in image.size)
image = image.resize(new_size, Image.LANCZOS)
logger.info(f"Resized input to: {new_size}")
if progress_callback:
progress_callback(50, f"Starting generation ({num_inference_steps} steps)...")
# Callback to track diffusion progress
def step_callback(pipe, step_index, timestep, callback_kwargs):
if progress_callback:
# Calculate progress: 50-85% range mapped to steps
progress = 50 + int((step_index / num_inference_steps) * 35)
progress_callback(progress, f"Step {step_index+1}/{num_inference_steps}")
return callback_kwargs
# Run inference
with torch.inference_mode():
result = self.pipe(
prompt=prompt,
negative_prompt=negative_prompt,
image=image,
strength=strength,
guidance_scale=guidance_scale,
num_inference_steps=num_inference_steps,
callback_on_step_end=step_callback if progress_callback else None
).images[0]
if progress_callback:
progress_callback(85, "Finalizing...")
return result
def unload(self):
"""
Unload the pipeline to free memory
"""
if self.pipe is not None:
del self.pipe
self.pipe = None
self.lora_loaded = False
torch.cuda.empty_cache()
logger.info("Pipeline unloaded")
# Global pipeline instance (singleton pattern)
_pipeline_manager = None
def get_pipeline_manager() -> DiffusionPipelineManager:
"""
Get or create the global pipeline manager instance
Returns:
DiffusionPipelineManager instance
"""
global _pipeline_manager
if _pipeline_manager is None:
_pipeline_manager = DiffusionPipelineManager()
return _pipeline_manager
def preprocess_image(image_path: Path) -> Image.Image:
"""
Load and preprocess an input image
This function:
- Loads the image
- Converts to RGB
- Normalizes exposure/gamma if needed
Args:
image_path: Path to input image
Returns:
Preprocessed PIL Image
"""
logger.info(f"Loading image from: {image_path}")
image = Image.open(image_path)
# Convert to RGB if necessary
if image.mode != "RGB":
image = image.convert("RGB")
# Optional: Add gamma/exposure normalization here
# This would analyze the histogram and adjust brightness
return image
def enhance_image(
input_image_path: Path,
output_path: Path,
strength: float = settings.IMG2IMG_STRENGTH,
guidance_scale: float = settings.GUIDANCE_SCALE,
custom_prompt: Optional[str] = None,
use_upscaler: bool = True,
use_postprocess: bool = True,
progress_callback: Optional[Callable[[int, str], None]] = None
) -> Path:
"""
Complete image enhancement pipeline
This is the main entry point for enhancing images. It orchestrates:
1. Image preprocessing
2. Stable Diffusion img2img with optional LoRA
3. Upscaling (optional)
4. Post-processing (optional)
Args:
input_image_path: Path to input image
output_path: Path to save enhanced image
strength: Image-to-image strength
guidance_scale: CFG scale
custom_prompt: Optional custom prompt
use_upscaler: Whether to apply upscaling
use_postprocess: Whether to apply post-processing
progress_callback: Optional callback for progress updates
Returns:
Path to the saved enhanced image
"""
try:
# Step 1: Load and preprocess
if progress_callback:
progress_callback(25, "Loading and preprocessing image...")
logger.info("Step 1/4: Preprocessing image")
image = preprocess_image(input_image_path)
# Step 2: Get pipeline and load LoRA if available
logger.info("Step 2/4: Running diffusion model")
pipeline = get_pipeline_manager()
lora_path = get_lora_path()
if lora_path and not pipeline.lora_loaded:
pipeline.load_lora(lora_path)
# Determine prompt
prompt = custom_prompt if custom_prompt else settings.DEFAULT_PROMPT
# Run enhancement
enhanced = pipeline.enhance(
image=image,
prompt=prompt,
negative_prompt=settings.NEGATIVE_PROMPT,
strength=strength,
guidance_scale=guidance_scale,
num_inference_steps=settings.NUM_INFERENCE_STEPS,
progress_callback=progress_callback
)
# Step 3: Upscaling
if use_upscaler:
if progress_callback:
progress_callback(92, "Upscaling image...")
logger.info("Step 3/4: Upscaling image")
enhanced = upscale_image(enhanced, scale=settings.UPSCALE_FACTOR)
else:
logger.info("Step 3/4: Skipping upscaling")
# Step 4: Post-processing
if use_postprocess:
if progress_callback:
progress_callback(94, "Applying final touches...")
logger.info("Step 4/4: Applying post-processing")
enhanced = postprocess_image(enhanced)
else:
logger.info("Step 4/4: Skipping post-processing")
# Save result
enhanced.save(output_path, quality=95, optimize=True)
logger.info(f"Enhanced image saved to: {output_path}")
return output_path
except Exception as e:
logger.error(f"Enhancement pipeline failed: {e}", exc_info=True)
raise