feat: AI photo enhancement (Real-ESRGAN container) + restore original
- 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)
This commit is contained in:
@@ -0,0 +1,73 @@
|
||||
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')
|
||||
Reference in New Issue
Block a user