import base64 import io import torch from pathlib import Path from PIL import Image from huggingface_hub import hf_hub_download from refiners.foundationals.latent_diffusion import solvers from enhancer import ESRGANUpscaler, ESRGANUpscalerCheckpoints class EndpointHandler: def __init__(self, path=""): """Initialize the handler with model checkpoints""" # Download model checkpoints self.checkpoints = ESRGANUpscalerCheckpoints( unet=Path( hf_hub_download( repo_id="refiners/juggernaut.reborn.sd1_5.unet", filename="model.safetensors", revision="347d14c3c782c4959cc4d1bb1e336d19f7dda4d2", ) ), clip_text_encoder=Path( hf_hub_download( repo_id="refiners/juggernaut.reborn.sd1_5.text_encoder", filename="model.safetensors", revision="744ad6a5c0437ec02ad826df9f6ede102bb27481", ) ), lda=Path( hf_hub_download( repo_id="refiners/juggernaut.reborn.sd1_5.autoencoder", filename="model.safetensors", revision="3c1aae3fc3e03e4a2b7e0fa42b62ebb64f1a4c19", ) ), controlnet_tile=Path( hf_hub_download( repo_id="refiners/controlnet.sd1_5.tile", filename="model.safetensors", revision="48ced6ff8bfa873a8976fa467c3629a240643387", ) ), esrgan=Path( hf_hub_download( repo_id="philz1337x/upscaler", filename="4x-UltraSharp.pth", revision="011deacac8270114eb7d2eeff4fe6fa9a837be70", ) ), negative_embedding=Path( hf_hub_download( repo_id="philz1337x/embeddings", filename="JuggernautNegative-neg.pt", revision="203caa7e9cc2bc225031a4021f6ab1ded283454a", ) ), negative_embedding_key="string_to_param.*", loras={ "more_details": Path( hf_hub_download( repo_id="philz1337x/loras", filename="more_details.safetensors", revision="a3802c0280c0d00c2ab18d37454a8744c44e474e", ) ), "sdxl_render": Path( hf_hub_download( repo_id="philz1337x/loras", filename="SDXLrender_v2.0.safetensors", revision="a3802c0280c0d00c2ab18d37454a8744c44e474e", ) ), }, ) # Initialize device and dtype self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") self.dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float32 # Initialize the enhancer self.enhancer = ESRGANUpscaler( checkpoints=self.checkpoints, device=self.device, dtype=self.dtype ) def __call__(self, data): """ Process the input data and return enhanced image Args: data (dict): Input data containing: - image (str): Base64 encoded input image - prompt (str, optional): Enhancement prompt - negative_prompt (str, optional): Negative prompt - seed (int, optional): Random seed - upscale_factor (float, optional): Upscale factor - controlnet_scale (float, optional): ControlNet scale - controlnet_decay (float, optional): ControlNet decay - condition_scale (int, optional): Condition scale - tile_width (int, optional): Tile width - tile_height (int, optional): Tile height - denoise_strength (float, optional): Denoise strength - num_inference_steps (int, optional): Number of inference steps - solver (str, optional): Solver type Returns: dict: Contains enhanced image as base64 string """ try: # Extract and decode input image image_data = data.get("image") if not image_data: return {"error": "No image provided"} # Decode base64 image if image_data.startswith('data:image'): # Remove data:image/...;base64, prefix if present image_data = image_data.split(',')[1] image_bytes = base64.b64decode(image_data) input_image = Image.open(io.BytesIO(image_bytes)).convert('RGB') # Extract parameters with defaults prompt = data.get("prompt", "masterpiece, best quality, highres") negative_prompt = data.get("negative_prompt", "worst quality, low quality, normal quality") seed = data.get("seed", 42) upscale_factor = data.get("upscale_factor", 2) controlnet_scale = data.get("controlnet_scale", 0.6) controlnet_decay = data.get("controlnet_decay", 1.0) condition_scale = data.get("condition_scale", 6) tile_width = data.get("tile_width", 112) tile_height = data.get("tile_height", 144) denoise_strength = data.get("denoise_strength", 0.35) num_inference_steps = data.get("num_inference_steps", 18) solver_name = data.get("solver", "DDIM") # Get solver type solver_type = getattr(solvers, solver_name) # Set up generator generator = torch.Generator(device=self.device) generator.manual_seed(seed) # Resize input image if too large to avoid VRAM issues side_size = min(input_image.size) if side_size > 768: scale = 768 / side_size new_size = (int(input_image.width * scale), int(input_image.height * scale)) resized_image = input_image.resize(new_size, resample=Image.Resampling.LANCZOS) else: resized_image = input_image # Enhance the image enhanced_image = self.enhancer.upscale( image=resized_image, prompt=prompt, negative_prompt=negative_prompt, upscale_factor=upscale_factor, controlnet_scale=controlnet_scale, controlnet_scale_decay=controlnet_decay, condition_scale=condition_scale, tile_size=(tile_height, tile_width), denoise_strength=denoise_strength, num_inference_steps=num_inference_steps, loras_scale={"more_details": 0.5, "sdxl_render": 1.0}, solver_type=solver_type, generator=generator, ) # Convert enhanced image to base64 buffered = io.BytesIO() enhanced_image.save(buffered, format="PNG") enhanced_base64 = base64.b64encode(buffered.getvalue()).decode('utf-8') return { "enhanced_image": enhanced_base64, "original_size": input_image.size, "enhanced_size": enhanced_image.size } except Exception as e: return {"error": f"Enhancement failed: {str(e)}"}