feat(photo-ai): реестр моделей и авто-выбор устройства
Stage 1: переписан photo-ai/app.py. Реестр моделей вместо одной модели, авто-выбор cuda/mps/cpu, ModelPool с ленивой загрузкой, прогревом и LRU на 2 записи, OOM-деградация по лестнице тайлов, расширенные /health и /models, /enhance с выбором модели, режима лиц и качества JPEG. Инвариант I1 держится: запрос только с image + scale=2 по-прежнему даёт байт-в-байт тот же JPEG (5/5 MATCH против эталона Stage 0). Расширение выводится по имени файла, а не по content_type, и .jpeg нормализуется в .jpg — этого хватает для совместимости с воркером, который шлёт image/jpeg для всего. /enhance отдаёт сырой JPEG по умолчанию и JSON при Accept: application/json. Пока модель грузится — 503 с Retry-After: 5; worker.js относит status >= 500 к мягким, поэтому попытка не тратится (I5). Новые переменные (PHOTO_AI_DEVICE, PHOTO_AI_MODELS_DIR, PHOTO_AI_TILE, PHOTO_AI_LOAD_ALL, PHOTO_AI_WARMUP, PHOTO_AI_FACE_MODEL, PHOTO_AI_JPEG_QUALITY) пока не документированы в .env.example: compose их ещё не подставляет, это Stage 2. Найдено: на хосте есть GPU (nvidia-smi, драйвер 615.71.09), не хватает только nvidia-container-toolkit. Установка — решение оператора, ветка CUDA осталась непроверенной прогоном.
This commit is contained in:
+821
-46
@@ -1,73 +1,848 @@
|
||||
import io
|
||||
import os
|
||||
import asyncio
|
||||
import base64
|
||||
import importlib.util
|
||||
import logging
|
||||
import os
|
||||
import platform
|
||||
import re
|
||||
import shutil
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import urllib.request
|
||||
from collections import OrderedDict
|
||||
from contextlib import contextmanager
|
||||
|
||||
from fastapi import FastAPI, File, Form, UploadFile
|
||||
from fastapi.responses import Response
|
||||
import numpy as np
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from basicsr.archs.rrdbnet_arch import RRDBNet
|
||||
from basicsr.archs.srvgg_arch import SRVGGNetCompact
|
||||
from fastapi import FastAPI, File, Form, Request, UploadFile
|
||||
from fastapi.responses import JSONResponse, Response
|
||||
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)))
|
||||
VENDOR_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'vendor')
|
||||
if os.path.isdir(VENDOR_DIR) and VENDOR_DIR not in sys.path:
|
||||
sys.path.append(VENDOR_DIR)
|
||||
|
||||
lock = threading.Lock()
|
||||
upsampler = None
|
||||
app = FastAPI()
|
||||
log = logging.getLogger('photo-ai')
|
||||
if not log.handlers:
|
||||
handler = logging.StreamHandler(sys.stderr)
|
||||
handler.setFormatter(logging.Formatter('%(asctime)s %(levelname)s %(name)s %(message)s'))
|
||||
log.addHandler(handler)
|
||||
log.setLevel(logging.INFO)
|
||||
log.propagate = False
|
||||
|
||||
FACE_MODES = ('off', 'face', 'all')
|
||||
DEFAULT_MODEL = 'x2plus'
|
||||
DEFAULT_STRENGTH = 0.7
|
||||
MIN_TILE = 64
|
||||
POOL_LIMIT = 2
|
||||
WARMUP_SIZE = 64
|
||||
RETRY_AFTER_SEC = 5
|
||||
|
||||
|
||||
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 env_text(name, default=''):
|
||||
value = os.environ.get(name)
|
||||
if value is None:
|
||||
return default
|
||||
value = value.strip()
|
||||
return value or default
|
||||
|
||||
|
||||
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)
|
||||
def env_flag(name, default):
|
||||
value = env_text(name).lower()
|
||||
if not value:
|
||||
return default
|
||||
return value not in ('0', 'false', 'no', 'off')
|
||||
|
||||
|
||||
def env_int(name, default):
|
||||
value = env_text(name)
|
||||
if not value:
|
||||
return default
|
||||
try:
|
||||
return int(value)
|
||||
except ValueError:
|
||||
log.warning('%s=%r не число, беру %s', name, value, default)
|
||||
return default
|
||||
|
||||
|
||||
def clamp(value, low, high):
|
||||
return max(low, min(high, value))
|
||||
|
||||
|
||||
DEVICE_PREF = env_text('PHOTO_AI_DEVICE', 'auto').lower()
|
||||
if DEVICE_PREF not in ('auto', 'cuda', 'cpu', 'mps'):
|
||||
log.warning('PHOTO_AI_DEVICE=%r неизвестно, беру auto', DEVICE_PREF)
|
||||
DEVICE_PREF = 'auto'
|
||||
MODELS_DIR = env_text('PHOTO_AI_MODELS_DIR', '/models')
|
||||
WEIGHTS_DIR = os.path.join(MODELS_DIR, 'weights')
|
||||
LEGACY_MODEL_PATH = env_text('MODEL_PATH')
|
||||
MAX_PIXELS = max(env_int('PHOTO_AI_MAX_PIXELS', env_int('MAX_INPUT_PIXELS', 4000000)), 1)
|
||||
BASE_TILE = max(env_int('PHOTO_AI_TILE', 256), 0)
|
||||
TILE_PAD = 10
|
||||
PRE_PAD = 0
|
||||
LOAD_ALL = env_flag('PHOTO_AI_LOAD_ALL', False)
|
||||
WARMUP = env_flag('PHOTO_AI_WARMUP', True)
|
||||
DEFAULT_FACE_MODEL = env_text('PHOTO_AI_FACE_MODEL', 'gfpgan').lower()
|
||||
DEFAULT_JPEG_QUALITY = clamp(env_int('PHOTO_AI_JPEG_QUALITY', 92), 70, 100)
|
||||
|
||||
OUTPUT_FORMATS = {
|
||||
'jpg': ('jpg', 'image/jpeg', [int(cv2.IMWRITE_JPEG_QUALITY)]),
|
||||
'png': ('png', 'image/png', [int(cv2.IMWRITE_PNG_COMPRESSION), 3]),
|
||||
'webp': ('webp', 'image/webp', [int(cv2.IMWRITE_WEBP_QUALITY), 92]),
|
||||
}
|
||||
|
||||
|
||||
def rrdb_x2():
|
||||
return RRDBNet(num_in_ch=3, num_out_ch=3, scale=2, num_feat=64, num_block=23, num_grow_ch=32)
|
||||
|
||||
|
||||
def srvgg(num_conv):
|
||||
return lambda: SRVGGNetCompact(num_in_ch=3, num_out_ch=3, num_feat=64, num_conv=num_conv,
|
||||
upscale=4, act_type='prelu')
|
||||
|
||||
|
||||
MODEL_REGISTRY = OrderedDict([
|
||||
('x2plus', {
|
||||
'scale': 2,
|
||||
'arch': rrdb_x2,
|
||||
'file': 'RealESRGAN_x2plus.pth',
|
||||
'url': 'https://github.com/xinntao/Real-ESRGAN/releases/download/v0.2.1/RealESRGAN_x2plus.pth',
|
||||
'min_bytes': 60_000_000,
|
||||
'alias': True,
|
||||
'denoise': False,
|
||||
'title': 'Универсальный апскейл x2, дефолт',
|
||||
}),
|
||||
('general-x4v3', {
|
||||
'scale': 4,
|
||||
'arch': srvgg(32),
|
||||
'file': 'realesr-general-x4v3.pth',
|
||||
'url': 'https://github.com/xinntao/Real-ESRGAN/releases/download/v0.2.5.0/realesr-general-x4v3.pth',
|
||||
'min_bytes': 2_000_000,
|
||||
'alias': False,
|
||||
'denoise': True,
|
||||
'title': 'Быстрый апскейл x4 с денойзом',
|
||||
'dni': {
|
||||
'file': 'realesr-general-wdn-x4v3.pth',
|
||||
'url': 'https://github.com/xinntao/Real-ESRGAN/releases/download/v0.2.5.0/realesr-general-wdn-x4v3.pth',
|
||||
'min_bytes': 2_000_000,
|
||||
'weight': 0.5,
|
||||
},
|
||||
}),
|
||||
('animevideo-v3', {
|
||||
'scale': 4,
|
||||
'arch': srvgg(16),
|
||||
'file': 'realesr-animevideov3.pth',
|
||||
'url': 'https://github.com/xinntao/Real-ESRGAN/releases/download/v0.2.5.0/realesr-animevideov3.pth',
|
||||
'min_bytes': 1_000_000,
|
||||
'alias': False,
|
||||
'denoise': False,
|
||||
'title': 'Быстрый апскейл x4 для скриншотов и иллюстраций',
|
||||
}),
|
||||
])
|
||||
|
||||
FACE_REGISTRY = OrderedDict([
|
||||
('gfpgan', {
|
||||
'module': 'gfpgan',
|
||||
'file': 'GFPGANv1.4.pth',
|
||||
'url': 'https://github.com/TencentARC/GFPGAN/releases/download/v1.3.0/GFPGANv1.4.pth',
|
||||
'min_bytes': 300_000_000,
|
||||
'arch': 'clean',
|
||||
'channel_multiplier': 2,
|
||||
'strength': False,
|
||||
'title': 'GFPGAN v1.4, восстановление лиц, дефолт',
|
||||
}),
|
||||
('codeformer', {
|
||||
'module': 'codeformer',
|
||||
'file': 'codeformer.pth',
|
||||
'url': 'https://github.com/sczhou/CodeFormer/releases/download/v0.1.0/codeformer.pth',
|
||||
'min_bytes': 300_000_000,
|
||||
'strength': True,
|
||||
'title': 'CodeFormer, восстановление лиц с регулируемой силой',
|
||||
}),
|
||||
])
|
||||
|
||||
FACEXLIB_WEIGHTS = OrderedDict([
|
||||
('detection_Resnet50_Final.pth', {
|
||||
'url': 'https://github.com/xinntao/facexlib/releases/download/v0.1.0/detection_Resnet50_Final.pth',
|
||||
'min_bytes': 90_000_000,
|
||||
}),
|
||||
('parsing_parsenet.pth', {
|
||||
'url': 'https://github.com/xinntao/facexlib/releases/download/v0.2.2/parsing_parsenet.pth',
|
||||
'min_bytes': 70_000_000,
|
||||
}),
|
||||
])
|
||||
|
||||
|
||||
def cuda_ready():
|
||||
try:
|
||||
return bool(torch.cuda.is_available()) and torch.cuda.device_count() > 0
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def mps_ready():
|
||||
try:
|
||||
return bool(torch.backends.mps.is_available())
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def cpu_device_name():
|
||||
return 'CPU (' + (platform.machine() or 'unknown') + ')'
|
||||
|
||||
|
||||
def pick_device():
|
||||
if DEVICE_PREF == 'cpu':
|
||||
return 'cpu', False, cpu_device_name()
|
||||
if DEVICE_PREF == 'mps':
|
||||
if mps_ready():
|
||||
return 'mps', False, 'Apple Silicon (MPS)'
|
||||
log.warning('PHOTO_AI_DEVICE=mps, но MPS недоступен — работаю на CPU')
|
||||
return 'cpu', False, cpu_device_name()
|
||||
if DEVICE_PREF == 'cuda':
|
||||
if not cuda_ready():
|
||||
log.warning('PHOTO_AI_DEVICE=cuda, но CUDA недоступна — работаю на CPU')
|
||||
return 'cpu', False, cpu_device_name()
|
||||
return 'cuda:0', True, torch.cuda.get_device_name(0)
|
||||
if cuda_ready():
|
||||
return 'cuda:0', True, torch.cuda.get_device_name(0)
|
||||
if mps_ready():
|
||||
return 'mps', False, 'Apple Silicon (MPS)'
|
||||
return 'cpu', False, cpu_device_name()
|
||||
|
||||
|
||||
state = {'device': 'cpu', 'half': False, 'device_name': 'CPU', 'tile': BASE_TILE, 'degraded': False}
|
||||
state['device'], state['half'], state['device_name'] = pick_device()
|
||||
INFER_LOCK = threading.RLock()
|
||||
|
||||
|
||||
class EnhanceError(Exception):
|
||||
def __init__(self, message, status_code=500):
|
||||
super().__init__(message)
|
||||
self.message = message
|
||||
self.status_code = status_code
|
||||
|
||||
|
||||
class ModelNotReady(EnhanceError):
|
||||
def __init__(self, label):
|
||||
super().__init__('модель %s ещё загружается, повторите позже' % label, 503)
|
||||
self.label = label
|
||||
|
||||
|
||||
def error_response(status_code, message, headers=None):
|
||||
return JSONResponse(status_code=status_code, content={'ok': False, 'error': message},
|
||||
headers=headers)
|
||||
|
||||
|
||||
def is_oom(err):
|
||||
if isinstance(err, torch.cuda.OutOfMemoryError):
|
||||
return True
|
||||
text = str(err).lower()
|
||||
return any(mark in text for mark in ('out of memory', 'not enough memory',
|
||||
'alloc_cpu', "can't allocate memory"))
|
||||
|
||||
|
||||
def tile_ladder(base):
|
||||
if base <= 0:
|
||||
return [base]
|
||||
ladder = []
|
||||
for tile in (base, base // 2, base // 4):
|
||||
if tile >= MIN_TILE and (not ladder or ladder[-1] != tile):
|
||||
ladder.append(tile)
|
||||
return ladder or [base]
|
||||
|
||||
|
||||
def file_ok(path, min_bytes):
|
||||
return os.path.isfile(path) and os.path.getsize(path) >= min_bytes
|
||||
|
||||
|
||||
def link_or_copy(src, dst):
|
||||
os.makedirs(os.path.dirname(dst), exist_ok=True)
|
||||
try:
|
||||
os.link(src, dst)
|
||||
except OSError:
|
||||
shutil.copyfile(src, dst)
|
||||
|
||||
|
||||
def download_weight(url, target, min_bytes):
|
||||
os.makedirs(os.path.dirname(target), exist_ok=True)
|
||||
tmp = target + '.tmp'
|
||||
log.info('качаю веса %s -> %s', url, target)
|
||||
try:
|
||||
urllib.request.urlretrieve(url, tmp)
|
||||
except Exception as err:
|
||||
if os.path.exists(tmp):
|
||||
os.unlink(tmp)
|
||||
raise EnhanceError('не удалось скачать веса %s: %s' % (os.path.basename(target), err), 500) from err
|
||||
size = os.path.getsize(tmp) if os.path.isfile(tmp) else 0
|
||||
if size < min_bytes:
|
||||
os.unlink(tmp)
|
||||
raise EnhanceError('веса %s повреждены: %d байт, минимум %d'
|
||||
% (os.path.basename(target), size, min_bytes), 500)
|
||||
os.replace(tmp, target)
|
||||
|
||||
|
||||
def ensure_weight(name, spec, alias=''):
|
||||
target = os.path.join(WEIGHTS_DIR, name)
|
||||
if file_ok(target, spec['min_bytes']):
|
||||
return target
|
||||
if os.path.isfile(target):
|
||||
log.warning('удаляю битые веса %s (%d байт)', target, os.path.getsize(target))
|
||||
os.unlink(target)
|
||||
if alias and file_ok(alias, spec['min_bytes']):
|
||||
link_or_copy(alias, target)
|
||||
log.info('веса %s взяты из существующего файла %s', name, alias)
|
||||
return target
|
||||
download_weight(spec['url'], target, spec['min_bytes'])
|
||||
return target
|
||||
|
||||
|
||||
def ensure_facexlib_weights():
|
||||
for name, spec in FACEXLIB_WEIGHTS.items():
|
||||
ensure_weight(name, spec)
|
||||
|
||||
|
||||
def upsamplers_of(obj):
|
||||
found = []
|
||||
if isinstance(obj, RealESRGANer):
|
||||
found.append(obj)
|
||||
inner = getattr(obj, 'upsampler', None)
|
||||
if isinstance(inner, RealESRGANer):
|
||||
found.append(inner)
|
||||
bg = getattr(obj, 'bg_upsampler', None)
|
||||
if isinstance(bg, RealESRGANer):
|
||||
found.append(bg)
|
||||
elif isinstance(getattr(bg, 'upsampler', None), RealESRGANer):
|
||||
found.append(bg.upsampler)
|
||||
return found
|
||||
|
||||
|
||||
def apply_tile(obj, tile):
|
||||
for upsampler in upsamplers_of(obj):
|
||||
upsampler.tile_size = tile
|
||||
|
||||
|
||||
def warm_up(runner):
|
||||
if not WARMUP:
|
||||
return
|
||||
noise = np.random.default_rng(0).integers(0, 256, (WARMUP_SIZE, WARMUP_SIZE, 3), dtype=np.uint8)
|
||||
with INFER_LOCK:
|
||||
runner.enhance(noise)
|
||||
|
||||
|
||||
class Upscaler:
|
||||
def __init__(self, upsampler, outscale):
|
||||
self.upsampler = upsampler
|
||||
self.outscale = outscale
|
||||
|
||||
def enhance(self, img, outscale=None):
|
||||
output, _mode = self.upsampler.enhance(img, outscale=outscale or self.outscale)
|
||||
return output
|
||||
|
||||
|
||||
class FaceRunner:
|
||||
def __init__(self, label, outscale, bg, restore):
|
||||
self.label = label
|
||||
self.outscale = outscale
|
||||
self.bg_upsampler = bg
|
||||
self.restore = restore
|
||||
|
||||
def enhance(self, img, strength=DEFAULT_STRENGTH):
|
||||
output, faces_found = self.restore(img, strength)
|
||||
return output, faces_found
|
||||
|
||||
|
||||
class ModelPool:
|
||||
def __init__(self, limit=POOL_LIMIT):
|
||||
self.lock = threading.Lock()
|
||||
self.entries = OrderedDict()
|
||||
self.loading = OrderedDict()
|
||||
self.in_use = {}
|
||||
self.limit = limit
|
||||
|
||||
def get(self, key, label, factory):
|
||||
with self.lock:
|
||||
entry = self.entries.get(key)
|
||||
if entry is not None:
|
||||
self.entries.move_to_end(key)
|
||||
self.in_use[key] = self.in_use.get(key, 0) + 1
|
||||
return entry['obj']
|
||||
if key in self.loading:
|
||||
raise ModelNotReady(label)
|
||||
self.loading[key] = label
|
||||
try:
|
||||
obj = factory()
|
||||
warm_up(obj)
|
||||
except BaseException:
|
||||
with self.lock:
|
||||
self.loading.pop(key, None)
|
||||
raise
|
||||
with self.lock:
|
||||
self.loading.pop(key, None)
|
||||
self.entries[key] = {'obj': obj, 'label': label}
|
||||
self.in_use[key] = self.in_use.get(key, 0) + 1
|
||||
apply_tile(obj, state['tile'])
|
||||
self.evict_locked()
|
||||
return obj
|
||||
|
||||
def release(self, key):
|
||||
with self.lock:
|
||||
if self.in_use.get(key):
|
||||
self.in_use[key] -= 1
|
||||
|
||||
def evict_locked(self):
|
||||
while len(self.entries) > self.limit:
|
||||
for key in list(self.entries):
|
||||
if not self.in_use.get(key):
|
||||
self.entries.pop(key, None)
|
||||
self.in_use.pop(key, None)
|
||||
log.info('выгружаю из кэша модель %s (LRU, лимит %d)', key, self.limit)
|
||||
break
|
||||
else:
|
||||
break
|
||||
|
||||
@contextmanager
|
||||
def acquire(self, key, label, factory):
|
||||
obj = self.get(key, label, factory)
|
||||
try:
|
||||
yield obj
|
||||
finally:
|
||||
self.release(key)
|
||||
|
||||
def contains(self, key):
|
||||
with self.lock:
|
||||
return key in self.entries
|
||||
|
||||
def clear(self):
|
||||
with self.lock:
|
||||
self.entries.clear()
|
||||
self.in_use.clear()
|
||||
|
||||
def set_tile(self, tile):
|
||||
with self.lock:
|
||||
for entry in self.entries.values():
|
||||
apply_tile(entry['obj'], tile)
|
||||
|
||||
def loaded_labels(self):
|
||||
with self.lock:
|
||||
labels = []
|
||||
for entry in self.entries.values():
|
||||
if entry['label'] not in labels:
|
||||
labels.append(entry['label'])
|
||||
return labels
|
||||
|
||||
def loading_labels(self):
|
||||
with self.lock:
|
||||
return list(self.loading.values())
|
||||
|
||||
|
||||
pool = ModelPool()
|
||||
|
||||
|
||||
def build_esrgan(name, denoise):
|
||||
spec = MODEL_REGISTRY[name]
|
||||
alias = LEGACY_MODEL_PATH if spec['alias'] else ''
|
||||
path = ensure_weight(spec['file'], spec, alias)
|
||||
model_path = path
|
||||
dni_weight = None
|
||||
if denoise and spec.get('dni'):
|
||||
dni_spec = spec['dni']
|
||||
dni_path = ensure_weight(dni_spec['file'], dni_spec)
|
||||
weight = float(dni_spec.get('weight', 0.5))
|
||||
model_path = [path, dni_path]
|
||||
dni_weight = (1.0 - weight, weight)
|
||||
upsampler = RealESRGANer(
|
||||
scale=2,
|
||||
model_path=MODEL_PATH,
|
||||
model=model,
|
||||
tile=256,
|
||||
tile_pad=10,
|
||||
pre_pad=0,
|
||||
half=False,
|
||||
device='cpu',
|
||||
scale=spec['scale'],
|
||||
model_path=model_path,
|
||||
dni_weight=dni_weight,
|
||||
model=spec['arch'](),
|
||||
tile=state['tile'],
|
||||
tile_pad=TILE_PAD,
|
||||
pre_pad=PRE_PAD,
|
||||
half=state['half'],
|
||||
device=torch.device(state['device']),
|
||||
)
|
||||
return Upscaler(upsampler, spec['scale'])
|
||||
|
||||
|
||||
def face_helper(outscale):
|
||||
from facexlib.utils.face_restoration_helper import FaceRestoreHelper
|
||||
return FaceRestoreHelper(
|
||||
upscale_factor=outscale,
|
||||
face_size=512,
|
||||
crop_ratio=(1, 1),
|
||||
det_model='retinaface_resnet50',
|
||||
save_ext='png',
|
||||
use_parse=True,
|
||||
device=torch.device(state['device']),
|
||||
model_rootpath=WEIGHTS_DIR,
|
||||
)
|
||||
|
||||
|
||||
def link_default_facexlib_dir():
|
||||
default_dir = os.path.abspath('gfpgan/weights')
|
||||
try:
|
||||
if os.path.realpath(default_dir) == os.path.realpath(WEIGHTS_DIR):
|
||||
return
|
||||
if os.path.isdir(default_dir) and not os.path.islink(default_dir):
|
||||
shutil.rmtree(default_dir, ignore_errors=True)
|
||||
os.makedirs(os.path.dirname(default_dir), exist_ok=True)
|
||||
if not os.path.lexists(default_dir):
|
||||
os.symlink(WEIGHTS_DIR, default_dir)
|
||||
log.info('каталог facexlib %s смотрит в том моделей', default_dir)
|
||||
except OSError as err:
|
||||
log.warning('не удалось направить %s в %s: %s', default_dir, WEIGHTS_DIR, err)
|
||||
|
||||
|
||||
def build_face_gfpgan(spec, outscale, bg):
|
||||
from gfpgan import GFPGANer
|
||||
link_default_facexlib_dir()
|
||||
ensure_facexlib_weights()
|
||||
path = ensure_weight(spec['file'], spec)
|
||||
restorer = GFPGANer(
|
||||
model_path=path,
|
||||
upscale=outscale,
|
||||
arch=spec['arch'],
|
||||
channel_multiplier=spec['channel_multiplier'],
|
||||
bg_upsampler=bg.upsampler,
|
||||
device=torch.device(state['device']),
|
||||
)
|
||||
|
||||
def restore(img, strength):
|
||||
cropped, _restored, output = restorer.enhance(img, has_aligned=False, only_center_face=False,
|
||||
paste_back=True)
|
||||
if output is None:
|
||||
output = bg.enhance(img, outscale)
|
||||
return output, len(cropped)
|
||||
|
||||
return FaceRunner('gfpgan', outscale, bg, restore)
|
||||
|
||||
|
||||
def build_face_codeformer(spec, outscale, bg):
|
||||
from codeformer import CodeFormer
|
||||
from gfpgan.utils import img2tensor, tensor2img
|
||||
from torchvision.transforms.functional import normalize
|
||||
ensure_facexlib_weights()
|
||||
path = ensure_weight(spec['file'], spec)
|
||||
net = CodeFormer(dim_embd=512, codebook_size=1024, n_head=8, n_layers=9,
|
||||
connect_list=['32', '64', '128', '256'], device=state['device'], fp16=False)
|
||||
net.load_state_dict(torch.load(path, map_location=lambda storage, loc: storage))
|
||||
net.eval()
|
||||
helper = face_helper(outscale)
|
||||
|
||||
def restore(img, strength):
|
||||
helper.clean_all()
|
||||
helper.read_image(img)
|
||||
helper.get_face_landmarks_5(only_center_face=False, eye_dist_threshold=5)
|
||||
helper.align_warp_face()
|
||||
for cropped in helper.cropped_faces:
|
||||
tensor = img2tensor(cropped / 255., bgr2rgb=True, float32=True)
|
||||
normalize(tensor, (0.5, 0.5, 0.5), (0.5, 0.5, 0.5), inplace=True)
|
||||
tensor = tensor.unsqueeze(0).to(net.device)
|
||||
with torch.no_grad():
|
||||
output = net(tensor, w=clamp(float(strength), 0.0, 1.0))[0]
|
||||
restored = tensor2img(output.squeeze(0), rgb2bgr=True, min_max=(-1, 1))
|
||||
helper.add_restored_face(restored.astype('uint8'))
|
||||
bg_img = bg.enhance(img, outscale)
|
||||
helper.get_inverse_affine(None)
|
||||
return helper.paste_faces_to_input_image(upsample_img=bg_img), len(helper.cropped_faces)
|
||||
|
||||
return FaceRunner('codeformer', outscale, bg, restore)
|
||||
|
||||
|
||||
FACE_BUILDERS = {'gfpgan': build_face_gfpgan, 'codeformer': build_face_codeformer}
|
||||
|
||||
|
||||
def face_installed(name):
|
||||
module = FACE_REGISTRY[name]['module']
|
||||
try:
|
||||
return importlib.util.find_spec(module) is not None
|
||||
except (ImportError, ValueError, AttributeError):
|
||||
return False
|
||||
|
||||
|
||||
def available_face_models():
|
||||
return [name for name in FACE_REGISTRY if face_installed(name)]
|
||||
|
||||
|
||||
def build_face(name, outscale, bg):
|
||||
spec = FACE_REGISTRY[name]
|
||||
return FACE_BUILDERS[name](spec, outscale, bg)
|
||||
|
||||
|
||||
def process_image(img, outscale, model_name, face, face_name, strength):
|
||||
denoise = face == 'all' and bool(MODEL_REGISTRY[model_name].get('denoise'))
|
||||
if face == 'off':
|
||||
suffix = ':wdn' if denoise else ''
|
||||
with pool.acquire('esrgan:' + model_name + suffix, model_name,
|
||||
lambda: build_esrgan(model_name, denoise)) as runner:
|
||||
return runner.enhance(img, outscale), 0
|
||||
key = 'face:%s@%d' % (face_name, outscale)
|
||||
|
||||
def factory():
|
||||
return build_face(face_name, outscale, build_esrgan(model_name, denoise))
|
||||
|
||||
with pool.acquire(key, face_name, factory) as runner:
|
||||
return runner.enhance(img, strength)
|
||||
|
||||
|
||||
def degrade_to_cpu(warnings):
|
||||
if state['degraded']:
|
||||
return
|
||||
log.warning('устройство %s не справилось, переключаюсь на CPU', state['device'])
|
||||
warnings.append('не хватило памяти на %s, обработка переведена на CPU' % state['device'])
|
||||
pool.clear()
|
||||
state['device'] = 'cpu'
|
||||
state['half'] = False
|
||||
state['device_name'] = cpu_device_name()
|
||||
state['degraded'] = True
|
||||
state['tile'] = BASE_TILE
|
||||
|
||||
|
||||
def run_guarded(img, outscale, model_name, face, face_name, strength, warnings):
|
||||
with INFER_LOCK:
|
||||
start_tile = state['tile']
|
||||
try:
|
||||
for tile in tile_ladder(start_tile):
|
||||
state['tile'] = tile
|
||||
pool.set_tile(tile)
|
||||
try:
|
||||
return process_image(img, outscale, model_name, face, face_name, strength)
|
||||
except RuntimeError as err:
|
||||
if not is_oom(err):
|
||||
raise EnhanceError('ошибка модели: %s' % err, 500) from err
|
||||
log.warning('нехватка памяти при tile=%s: %s', tile, err)
|
||||
warnings.append('не хватило памяти при tile=%d' % tile)
|
||||
if not state['device'].startswith('cpu'):
|
||||
degrade_to_cpu(warnings)
|
||||
try:
|
||||
return process_image(img, outscale, model_name, face, face_name, strength)
|
||||
except RuntimeError as err:
|
||||
if is_oom(err):
|
||||
raise EnhanceError('не хватило памяти даже на CPU: %s' % err, 500) from err
|
||||
raise EnhanceError('ошибка модели на CPU: %s' % err, 500) from err
|
||||
raise EnhanceError('не хватило памяти даже при tile=%d: пересмотрите PHOTO_AI_TILE или PHOTO_AI_MAX_PIXELS'
|
||||
% state['tile'], 500)
|
||||
finally:
|
||||
state['tile'] = start_tile
|
||||
pool.set_tile(start_tile)
|
||||
|
||||
|
||||
def encode_image(out, fmt, quality):
|
||||
ext, media_type, params = fmt
|
||||
if ext == 'jpg':
|
||||
params = params + [quality]
|
||||
ok, encoded = cv2.imencode('.' + ext, out, params)
|
||||
if not ok:
|
||||
raise EnhanceError('не удалось закодировать результат как %s' % ext, 500)
|
||||
return encoded.tobytes(), media_type
|
||||
|
||||
|
||||
def output_format(filename):
|
||||
name = (filename or '').rsplit('/', 1)[-1]
|
||||
ext = name.rsplit('.', 1)[-1].lower() if '.' in name else ''
|
||||
if ext == 'jpeg':
|
||||
ext = 'jpg'
|
||||
return OUTPUT_FORMATS.get(ext, OUTPUT_FORMATS['jpg'])
|
||||
|
||||
|
||||
def wants_json(request):
|
||||
return 'application/json' in (request.headers.get('accept') or '').lower()
|
||||
|
||||
|
||||
def driver_version():
|
||||
try:
|
||||
with open('/proc/driver/nvidia/version', 'r') as handle:
|
||||
text = handle.read()
|
||||
except OSError:
|
||||
return None
|
||||
match = re.search(r'\d+\.\d+\.\d+', text)
|
||||
return match.group(0) if match else None
|
||||
|
||||
|
||||
def vram_total_mb():
|
||||
if not state['device'].startswith('cuda'):
|
||||
return None
|
||||
try:
|
||||
return round(torch.cuda.get_device_properties(0).total_memory / (1024 * 1024))
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def vram_free_mb():
|
||||
if not state['device'].startswith('cuda'):
|
||||
return None
|
||||
try:
|
||||
free, _total = torch.cuda.mem_get_info()
|
||||
return round(free / (1024 * 1024))
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def validate(model_name, face, face_name, strength, quality):
|
||||
if model_name not in MODEL_REGISTRY:
|
||||
raise EnhanceError('неизвестная модель %r, доступны: %s'
|
||||
% (model_name, ', '.join(MODEL_REGISTRY)), 400)
|
||||
if face not in FACE_MODES:
|
||||
raise EnhanceError('неизвестный режим лиц %r, доступны: %s' % (face, ', '.join(FACE_MODES)), 400)
|
||||
if quality < 70 or quality > 100:
|
||||
raise EnhanceError('jpeg_quality должен быть 70..100, получено %d' % quality, 400)
|
||||
if face == 'off':
|
||||
return
|
||||
if face_name not in FACE_REGISTRY:
|
||||
raise EnhanceError('неизвестная face-модель %r, доступны: %s'
|
||||
% (face_name, ', '.join(available_face_models()) or 'нет'), 400)
|
||||
if not face_installed(face_name):
|
||||
others = [name for name in available_face_models()]
|
||||
raise EnhanceError('модель лиц %s не установлена в образ, доступен %s'
|
||||
% (face_name, ', '.join(others) or 'ни один'), 400)
|
||||
if strength < 0.0 or strength > 1.0:
|
||||
raise EnhanceError('strength должен быть 0..1, получено %s' % strength, 400)
|
||||
if not FACE_REGISTRY[face_name]['strength'] and abs(strength - DEFAULT_STRENGTH) > 1e-6:
|
||||
raise EnhanceError('strength применяется только к CodeFormer, для %s оставьте %s'
|
||||
% (face_name, DEFAULT_STRENGTH), 400)
|
||||
|
||||
|
||||
def preload():
|
||||
names = list(MODEL_REGISTRY) if LOAD_ALL else [DEFAULT_MODEL]
|
||||
for name in names:
|
||||
try:
|
||||
with pool.acquire('esrgan:' + name, name, lambda n=name: build_esrgan(n, False)):
|
||||
log.info('модель %s готова', name)
|
||||
except Exception as err:
|
||||
log.error('предзагрузка модели %s не удалась: %s', name, err)
|
||||
|
||||
|
||||
app = FastAPI()
|
||||
|
||||
|
||||
@app.on_event('startup')
|
||||
async def startup():
|
||||
await asyncio.to_thread(load_model)
|
||||
log.info('устройство: %s (%s), half=%s, tile=%s, веса: %s',
|
||||
state['device'], state['device_name'], state['half'], state['tile'], WEIGHTS_DIR)
|
||||
threading.Thread(target=preload, name='photo-ai-preload', daemon=True).start()
|
||||
|
||||
|
||||
@app.get('/health')
|
||||
def health():
|
||||
return {'ok': upsampler is not None}
|
||||
return {
|
||||
'ok': True,
|
||||
'ready': pool.contains('esrgan:' + DEFAULT_MODEL),
|
||||
'device': state['device'],
|
||||
'device_name': state['device_name'],
|
||||
'half': state['half'],
|
||||
'tile': state['tile'],
|
||||
'driver': driver_version(),
|
||||
'cuda': torch.version.cuda,
|
||||
'vram_total_mb': vram_total_mb(),
|
||||
'vram_free_mb': vram_free_mb(),
|
||||
'models': list(MODEL_REGISTRY),
|
||||
'face_models': available_face_models(),
|
||||
'loaded': pool.loaded_labels(),
|
||||
'loading': pool.loading_labels(),
|
||||
'max_pixels': MAX_PIXELS,
|
||||
}
|
||||
|
||||
|
||||
@app.get('/models')
|
||||
def models():
|
||||
loaded = pool.loaded_labels()
|
||||
faces = available_face_models()
|
||||
return {
|
||||
'device': state['device'],
|
||||
'device_name': state['device_name'],
|
||||
'half': state['half'],
|
||||
'tile': state['tile'],
|
||||
'max_pixels': MAX_PIXELS,
|
||||
'defaults': {
|
||||
'model': DEFAULT_MODEL,
|
||||
'scale': 2,
|
||||
'face': 'off',
|
||||
'face_model': DEFAULT_FACE_MODEL,
|
||||
'strength': DEFAULT_STRENGTH,
|
||||
'jpeg_quality': DEFAULT_JPEG_QUALITY,
|
||||
},
|
||||
'models': [
|
||||
{
|
||||
'name': name,
|
||||
'title': spec['title'],
|
||||
'scale': spec['scale'],
|
||||
'weights': spec['file'],
|
||||
'denoise': bool(spec.get('denoise')),
|
||||
'loaded': name in loaded,
|
||||
}
|
||||
for name, spec in MODEL_REGISTRY.items()
|
||||
],
|
||||
'face_models': [
|
||||
{
|
||||
'name': name,
|
||||
'title': spec['title'],
|
||||
'weights': spec['file'],
|
||||
'available': name in faces,
|
||||
'strength': bool(spec['strength']),
|
||||
'loaded': name in loaded,
|
||||
}
|
||||
for name, spec in FACE_REGISTRY.items()
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@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')
|
||||
async def enhance(request: Request,
|
||||
image: UploadFile = File(...),
|
||||
scale: int = Form(2),
|
||||
model: str = Form(DEFAULT_MODEL),
|
||||
face: str = Form('off'),
|
||||
face_model: str = Form(''),
|
||||
strength: float = Form(DEFAULT_STRENGTH),
|
||||
jpeg_quality: int = Form(0)):
|
||||
started = time.time()
|
||||
model_name = (model or DEFAULT_MODEL).strip().lower()
|
||||
face_mode = (face or 'off').strip().lower()
|
||||
face_name = (face_model or '').strip().lower() or (
|
||||
DEFAULT_FACE_MODEL if DEFAULT_FACE_MODEL in FACE_REGISTRY else 'gfpgan')
|
||||
quality = jpeg_quality if jpeg_quality else DEFAULT_JPEG_QUALITY
|
||||
try:
|
||||
validate(model_name, face_mode, face_name, strength, quality)
|
||||
data = await image.read()
|
||||
img = cv2.imdecode(np.frombuffer(data, np.uint8), cv2.IMREAD_COLOR)
|
||||
if img is None:
|
||||
raise EnhanceError('bad image', 400)
|
||||
if img.shape[0] * img.shape[1] > MAX_PIXELS:
|
||||
ratio = (MAX_PIXELS / (img.shape[0] * img.shape[1])) ** 0.5
|
||||
img = cv2.resize(img, (int(img.shape[1] * ratio), int(img.shape[0] * ratio)),
|
||||
interpolation=cv2.INTER_AREA)
|
||||
outscale = min(max(int(scale), 2), 4)
|
||||
fmt = output_format(image.filename)
|
||||
warnings = []
|
||||
output, faces_found = await asyncio.to_thread(run_guarded, img, outscale, model_name,
|
||||
face_mode, face_name, float(strength), warnings)
|
||||
payload, media_type = encode_image(output, fmt, int(quality))
|
||||
except EnhanceError as err:
|
||||
if err.status_code == 503:
|
||||
return error_response(503, err.message, headers={'Retry-After': str(RETRY_AFTER_SEC)})
|
||||
return error_response(err.status_code, err.message)
|
||||
if face_mode != 'off' and not faces_found:
|
||||
warnings.append('лица не найдены, фон обработан апскейлом')
|
||||
elapsed_ms = int((time.time() - started) * 1000)
|
||||
if wants_json(request):
|
||||
return {
|
||||
'ok': True,
|
||||
'image_base64': base64.b64encode(payload).decode('ascii'),
|
||||
'image_ext': fmt[0],
|
||||
'model': model_name,
|
||||
'face': face_mode,
|
||||
'face_model': face_name if face_mode != 'off' else None,
|
||||
'faces_found': faces_found,
|
||||
'device': state['device'],
|
||||
'elapsed_ms': elapsed_ms,
|
||||
'warnings': warnings,
|
||||
}
|
||||
return Response(payload, media_type=media_type, headers={'X-Photo-AI-Model': model_name,
|
||||
'X-Photo-AI-Face': face_mode,
|
||||
'X-Photo-AI-Device': state['device'],
|
||||
'X-Photo-AI-Elapsed-Ms': str(elapsed_ms)})
|
||||
|
||||
Reference in New Issue
Block a user