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 осталась непроверенной прогоном.
849 lines
29 KiB
Python
849 lines
29 KiB
Python
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
|
|
|
|
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
|
|
|
|
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)
|
|
|
|
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 env_text(name, default=''):
|
|
value = os.environ.get(name)
|
|
if value is None:
|
|
return default
|
|
value = value.strip()
|
|
return value or default
|
|
|
|
|
|
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=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():
|
|
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': 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(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)})
|