Team Ai
Modelpublic

mrcuddle/URPM-Inpaint-Hyper-SDXL

sourceHugging Faceotherupdated 2y agoView on Hugging Face
2likes34downloads
handler.py90 linesDownload Raw Back to root
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)