Download handler.py from TBurdairon/finegrain-image-enhancer-endpoint: direct link, hf CLI and curl.
- Browser
- Download file 8.87 kB
-
https://huggingface.co/TBurdairon/finegrain-image-enhancer-endpoint/resolve/main/handler.py
- Command line
-
hf download hf://TBurdairon/finegrain-image-enhancer-endpoint/handler.py
-
curl -L -o handler.py https://huggingface.co/TBurdairon/finegrain-image-enhancer-endpoint/resolve/main/handler.py
8.87 kB
| 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() | |
| } |