import io import os import asyncio import threading import urllib.request from fastapi import FastAPI, File, Form, UploadFile from fastapi.responses import Response import numpy as np import cv2 import torch from basicsr.archs.rrdbnet_arch import RRDBNet from realesrgan import RealESRGANer MODEL_PATH = os.environ.get('MODEL_PATH', '/models/RealESRGAN_x2plus.pth') MODEL_URL = 'https://github.com/xinntao/Real-ESRGAN/releases/download/v0.2.1/RealESRGAN_x2plus.pth' MAX_PIXELS = int(os.environ.get('MAX_INPUT_PIXELS', str(4_000_000))) lock = threading.Lock() upsampler = None app = FastAPI() def ensure_model(): if not os.path.exists(MODEL_PATH): os.makedirs(os.path.dirname(MODEL_PATH), exist_ok=True) tmp = MODEL_PATH + '.tmp' urllib.request.urlretrieve(MODEL_URL, tmp) os.replace(tmp, MODEL_PATH) def load_model(): global upsampler ensure_model() model = RRDBNet(num_in_ch=3, num_out_ch=3, scale=2, num_feat=64, num_block=23, num_grow_ch=32) upsampler = RealESRGANer( scale=2, model_path=MODEL_PATH, model=model, tile=256, tile_pad=10, pre_pad=0, half=False, device='cpu', ) @app.on_event('startup') async def startup(): await asyncio.to_thread(load_model) @app.get('/health') def health(): return {'ok': upsampler is not None} @app.post('/enhance') async def enhance(image: UploadFile = File(...), scale: int = Form(2)): data = await image.read() img = cv2.imdecode(np.frombuffer(data, np.uint8), cv2.IMREAD_COLOR) if img is None: return Response('bad image', status_code=400) if img.shape[0] * img.shape[1] > MAX_PIXELS: r = (MAX_PIXELS / (img.shape[0] * img.shape[1])) ** 0.5 img = cv2.resize(img, (int(img.shape[1] * r), int(img.shape[0] * r)), interpolation=cv2.INTER_AREA) outscale = min(max(int(scale), 2), 4) with lock: out, _ = upsampler.enhance(img, outscale=outscale) ok, enc = cv2.imencode('.jpg', out, [int(cv2.IMWRITE_JPEG_QUALITY), 92]) if not ok: return Response('encode failed', status_code=500) return Response(enc.tobytes(), media_type='image/jpeg')