Stage 3 закрыт. Face-режим (off/face/all, strength, warnings) был написан в Stage 1;
здесь доведена приёмка и найден баг, который невозможно было увидеть без GPU-прогона.
Главное: realesные OOM никогда не доходили до run_guarded. И realessrgan/utils.py, и
gfpgan/utils.py ловят RuntimeError вокруг вызова сети и идут дальше
(`except RuntimeError as error: print('Error', error)`), поэтому на нехватку памяти
realesrgan падал уже не RuntimeError, а UnboundLocalError на присваивании
output_tile. Наружу уходил голый 500 «Internal Server Error»: ни лестницы тайлов,
ни деградации на CPU, ни внятного текста. В gfpgan это было тихое ухудшение —
при OOM лицо молча оставалось исходным, а задание уходило в «успех».
Починка: guard_forward() оборачивает forward сетей, которые строим мы
(RealESRGANer.model, restorer.gfpgan, CodeFormer net) и превращает OOM-RuntimeError
в TileOOM. TileOOM не наследует RuntimeError, поэтому проглатывающие except его
пропускают; is_oom и обе точки run_guarded ловят его явно. Обёртка вешается на
экземпляр, идемпотентна по флагу _photo_ai_guarded и не трогает класс.
Также добавлены два предупреждения из чек-листа, которых в коде не было: вход меньше
320×320 (лица могут не найтись) и CodeFormer на не-CUDA. Предупреждение «лица не
найдены» и деградация на CPU были на месте и не менялись.
Проверено на RTX 3050 Laptop (4096 МБ, драйвер 615.71.09, CUDA 12.6):
- tile=2048, x2plus face=all, полное фото: до правки 500 + UnboundLocalError,
после 200 за 28.3 с с единственным warning «не хватило памяти при tile=2048»;
- инъекция OOM: лестница 2048 → 1024 → 512, затем переход на CPU (device: cpu,
half: false) и успешный повтор; при повторе уже на CPU — честная 500 с подсказкой
про PHOTO_AI_TILE / PHOTO_AI_MAX_PIXELS;
- 640×480: off 1.0 с / face 3.8 с / all 2.2 с, faces_found=6; фото без лиц даёт
faces_found=0 и байты, равные face=off;
- I1: 6/6 MATCH байт-в-байт против Stage 2 на CPU (jpg/.jpeg/png+70/webp+100/x4v3/anime).
docker-compose.gpu.yml: GPU выдаётся через CDI (device_ids nvidia.com/gpu=all) —
не требует правки /etc/docker/daemon.json и перезапуска демона, в отличие от
классического резервирования driver: nvidia. TORCH_VARIANT cu124 → cu126: в индексе
cu124 последний torch 2.6.0, а cu126 даёт те же 2.14.0/0.29.0, что и CPU-образ, так
что варианты сборки отличаются только CUDA-библиотеками. Образ тегируется отдельно
(whatido-photo-ai:cu126), чтобы сборка GPU-варианта не перетирала CPU-образ
whatido-photo-ai:latest — откат остаётся обычным docker compose up -d photo-ai.
README: раздел «Запуск на NVIDIA GPU» с установкой NVIDIA Container Toolkit и генерацией
CDI-спеки, оговорками про 4 ГБ VRAM (x2plus + gfpgan влезают, general-x4v3 тяжелее,
LOAD_ALL=1 лучше не включать) и описанием параметров /enhance. .env.example: команда
GPU-запуска и рекомендация по PHOTO_AI_TILE. Журналы раздела 3 и GPU-прогона — в
TODO_PHOTO_FACE_AI.md.
worker.js, server.js, схема БД и фронтенд не тронуты — они в Stage 4…6.
880 lines
30 KiB
Python
880 lines
30 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
|
||
MIN_FACE_SIDE = 320
|
||
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
|
||
|
||
|
||
class TileOOM(Exception):
|
||
pass
|
||
|
||
|
||
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, TileOOM)):
|
||
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 guard_forward(model):
|
||
if getattr(model, '_photo_ai_guarded', False):
|
||
return model
|
||
forward = model.forward
|
||
|
||
def guarded(*args, **kwargs):
|
||
try:
|
||
return forward(*args, **kwargs)
|
||
except RuntimeError as err:
|
||
if not is_oom(err):
|
||
raise
|
||
raise TileOOM(str(err)) from err
|
||
|
||
model.forward = guarded
|
||
model._photo_ai_guarded = True
|
||
return model
|
||
|
||
|
||
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=guard_forward(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']),
|
||
)
|
||
guard_forward(restorer.gfpgan)
|
||
|
||
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()
|
||
guard_forward(net)
|
||
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, TileOOM) 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, TileOOM) 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 = []
|
||
if face_mode != 'off':
|
||
if min(img.shape[0], img.shape[1]) < MIN_FACE_SIDE:
|
||
warnings.append('вход меньше %d×%d — лица могут не найтись'
|
||
% (MIN_FACE_SIDE, MIN_FACE_SIDE))
|
||
if face_name == 'codeformer' and not state['device'].startswith('cuda'):
|
||
warnings.append('CodeFormer на %s медленнее GFPGAN' % state['device'])
|
||
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)})
|