from pathlib import Path import torch from PIL import Image import base64 import io from enhancer import ESRGANUpscaler, ESRGANUpscalerCheckpoints checkpoints = ESRGANUpscalerCheckpoints( esrgan=Path("checkpoints/4x-UltraSharp.pth") ) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float32 enhancer = ESRGANUpscaler( checkpoints=checkpoints, device=device, dtype=dtype ) def inference(inputs: dict) -> dict: if "image" not in inputs: return {"error": "No image provided"} image_data = inputs["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") enhanced_image = enhancer.upscale(input_image) buf = io.BytesIO() enhanced_image.save(buf, format="PNG") b64 = base64.b64encode(buf.getvalue()).decode("utf-8") return { "enhanced_image": b64, "original_size": input_image.size, "enhanced_size": enhanced_image.size }