- new photo-ai service: FastAPI + Real-ESRGAN x2plus on CPU, internal only - async job queue in server (POST enhance-ai / GET status), 5min timeout, sequential processing - keep original photo backup (entries.photo_original_path, uploads/.originals), restore-original endpoint - UI: AI button and restore-original button in enhance modal - db: photo_original_path column (init.sql, migration.sql, runtime ensure)
74 lines
2.1 KiB
Python
74 lines
2.1 KiB
Python
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')
|