import base64 import io import torch import json 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""" try: # DEBUG: Log what we received print(f"DEBUG: Received data type: {type(data)}") print(f"DEBUG: Data keys: {list(data.keys()) if isinstance(data, dict) else 'Not a dict'}") # Handle different input formats if "inputs" in data: inputs = data["inputs"] print("DEBUG: Found 'inputs' key") else: inputs = data print("DEBUG: Using data directly as inputs") print(f"DEBUG: Inputs type: {type(inputs)}") print(f"DEBUG: Inputs keys: {list(inputs.keys()) if isinstance(inputs, dict) else 'Not a dict'}") # Try to find image in different possible locations image_data = None if isinstance(inputs, dict): image_data = inputs.get("image") if not image_data: # Try other possible keys for key in ["img", "input_image", "data"]: if key in inputs: image_data = inputs[key] print(f"DEBUG: Found image data in '{key}' key") break print(f"DEBUG: Image data found: {image_data is not None}") if image_data: print(f"DEBUG: Image data length: {len(image_data)}") print(f"DEBUG: Image data starts with: {image_data[:50] if len(image_data) > 50 else image_data}") if not image_data: return { "error": "No image provided", "debug_info": { "data_type": str(type(data)), "data_keys": list(data.keys()) if isinstance(data, dict) else None, "inputs_type": str(type(inputs)), "inputs_keys": list(inputs.keys()) if isinstance(inputs, dict) else None } } # Decode base64 image if image_data.startswith('data:image'): 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 = inputs.get("prompt", "masterpiece, best quality, highres") negative_prompt = inputs.get("negative_prompt", "worst quality, low quality, normal quality") seed = inputs.get("seed", 42) upscale_factor = inputs.get("upscale_factor", 2) controlnet_scale = inputs.get("controlnet_scale", 0.6) controlnet_decay = inputs.get("controlnet_decay", 1.0) condition_scale = inputs.get("condition_scale", 6) tile_width = inputs.get("tile_width", 112) tile_height = inputs.get("tile_height", 144) denoise_strength = inputs.get("denoise_strength", 0.35) num_inference_steps = inputs.get("num_inference_steps", 18) solver_name = inputs.get("solver", "DDIM") print(f"DEBUG: Processing image of size: {input_image.size}") # 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 print(f"DEBUG: About to enhance image of size: {resized_image.size}") # 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: import traceback return { "error": f"Enhancement failed: {str(e)}", "traceback": traceback.format_exc() }