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