Files
WhatIDo/photo-ai/app.py
T
dev e8c15cc425 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
осталась непроверенной прогоном.
2026-09-29 00:36:16 +03:00

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)})