mrcuddle/URPM-Inpaint-Hyper-SDXL
234
1import torch2import json3import base644import io5from PIL import Image6from diffusers import DPMSolverMultistepScheduler, StableDiffusionXLInpaintPipeline7 8# Set device9device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')10 11if device.type != 'cuda':12 raise ValueError("Need to run on GPU")13 14class EndpointHandler:15 def __init__(self, path="mrcuddle/URPM-Inpaint-Hyper-SDXL"):16 """Load the SDXL Inpainting model."""17 self.pipeline = StableDiffusionXLInpaintPipeline.from_pretrained(18 path, torch_dtype=torch.float1619 )20 self.pipeline.scheduler = DPMSolverMultistepScheduler.from_config(self.pipeline.scheduler.config)21 self.pipeline = self.pipeline.to(device)22 23 def __call__(self, data: dict):24 """Custom call function for Hugging Face Inference Endpoints."""25 try:26 # Extract inputs from JSON payload27 inputs = data.get("inputs", "")28 encoded_image = data.get("image", None)29 encoded_mask_image = data.get("mask_image", None)30 31 # Extract optional parameters with default values32 num_inference_steps = data.get("num_inference_steps", 25)33 guidance_scale = data.get("guidance_scale", 7.5)34 negative_prompt = data.get("negative_prompt", None)35 height = data.get("height", None)36 width = data.get("width", None)37 38 # Ensure both images are provided39 if not encoded_image or not encoded_mask_image:40 raise ValueError("Both 'image' and 'mask_image' are required in base64 format.")41 42 # Decode base64 images43 image = self.decode_base64_image(encoded_image)44 mask_image = self.decode_base64_image(encoded_mask_image)45 46 print("\n--- Running Inference ---")47 print(f"Prompt: {inputs}")48 print(f"Steps: {num_inference_steps}, Guidance Scale: {guidance_scale}")49 print(f"Negative Prompt: {negative_prompt}")50 print(f"Image Size: {image.size}, Mask Size: {mask_image.size}")51 52 # Run inference53 output_image = self.pipeline(54 prompt=inputs,55 image=image,56 mask_image=mask_image,57 num_inference_steps=num_inference_steps,58 guidance_scale=guidance_scale,59 num_images_per_prompt=1,60 negative_prompt=negative_prompt,61 height=height,62 width=width63 ).images[0]64 65 # Return base64-encoded image66 return json.dumps({"output": self.encode_base64_image(output_image)})67 68 except Exception as e:69 return json.dumps({"error": str(e)})70 71 def decode_base64_image(self, image_string):72 """Decode base64-encoded image to a PIL Image."""73 try:74 base64_image = base64.b64decode(image_string)75 buffer = io.BytesIO(base64_image)76 return Image.open(buffer).convert("RGB")77 except Exception as e:78 raise ValueError(f"Failed to decode base64 image: {e}")79 80 def encode_base64_image(self, image):81 """Encode PIL image to base64."""82 buffered = io.BytesIO()83 image.save(buffered, format="PNG")84 return base64.b64encode(buffered.getvalue()).decode("utf-8")85 86# Create an instance of EndpointHandler87handler = EndpointHandler()88 89def handle(data: dict):90 return handler(data)