Files
WhatIDo/photo-ai/app.py
T
dev 4c63a46d24 fix(photo-ai): лестница OOM на CUDA + приёмка face-режима и GPU-оверрайда
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.
2026-09-29 12:02:03 +03:00

880 lines
30 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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)})