feat(photo-ai): официальный CodeFormer вендорен, веса на этапе сборки

CodeFormer был написан по API сторонней обёртки с PyPI (rohitkhatri):
CodeFormer(..., device=..., fp16=False) и net.device официальный класс
не принимает — вызов падал бы с TypeError. В photo-ai/vendor/codeformer
вендорен официальный sczhou/CodeFormer (codeformer_arch.py +
vqgan_arch.py, b33cc7d, лицензия S-Lab 1.0), устройство передаётся
через net.to(device), чекпоинт читается из params_ema.

Веса больше не качаются лениво при первом запросе: fetch-weights.py на
сборке образа кладёт их в /opt/photo-ai-seed (build-arg
PHOTO_AI_PREFETCH, дефолт codeformer), при старте seed_weights()
переносит их в том photo-ai-models:/models/weights. Том переживает
пересборку образа, BuildKit-кэш не даёт качать повторно, недоступная
сеть на сборке не роняет образ.
This commit is contained in:
dev
2026-10-04 21:05:15 +03:00
parent 1d71e249e4
commit 19be1cc9ef
12 changed files with 955 additions and 21 deletions
+50 -3
View File
@@ -80,6 +80,7 @@ if DEVICE_PREF not in ('auto', 'cuda', 'cpu', 'mps'):
DEVICE_PREF = 'auto'
MODELS_DIR = env_text('PHOTO_AI_MODELS_DIR', '/models')
WEIGHTS_DIR = os.path.join(MODELS_DIR, 'weights')
SEED_DIR = env_text('PHOTO_AI_SEED_DIR', '/opt/photo-ai-seed')
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)
@@ -320,10 +321,48 @@ def ensure_weight(name, spec, alias=''):
link_or_copy(alias, target)
log.info('веса %s взяты из существующего файла %s', name, alias)
return target
seed = os.path.join(SEED_DIR, name)
if file_ok(seed, spec['min_bytes']):
link_or_copy(seed, target)
log.info('веса %s взяты из слоя образа %s', name, seed)
return target
download_weight(spec['url'], target, spec['min_bytes'])
return target
def seed_weights():
names = [spec['file'] for spec in MODEL_REGISTRY.values()]
names += [spec['file'] for spec in FACE_REGISTRY.values()]
names += [spec['dni']['file'] for spec in MODEL_REGISTRY.values() if spec.get('dni')]
names += list(FACEXLIB_WEIGHTS)
seeded = 0
for name in names:
source = os.path.join(SEED_DIR, name)
spec = weight_spec(name)
if spec is None or not file_ok(source, spec['min_bytes']):
continue
try:
ensure_weight(name, spec)
seeded += 1
except EnhanceError as err:
log.warning('не удалось перенести веса %s из образа: %s', name, err)
if seeded:
log.info('перенесено весов из слоя образа: %d', seeded)
return seeded
def weight_spec(name):
for spec in MODEL_REGISTRY.values():
if spec['file'] == name:
return spec
if spec.get('dni') and spec['dni']['file'] == name:
return spec['dni']
for spec in FACE_REGISTRY.values():
if spec['file'] == name:
return spec
return FACEXLIB_WEIGHTS.get(name)
def ensure_facexlib_weights():
for name, spec in FACEXLIB_WEIGHTS.items():
ensure_weight(name, spec)
@@ -552,10 +591,14 @@ def build_face_codeformer(spec, outscale, bg):
from torchvision.transforms.functional import normalize
ensure_facexlib_weights()
path = ensure_weight(spec['file'], spec)
device = torch.device(state['device'])
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))
connect_list=['32', '64', '128', '256'])
checkpoint = torch.load(path, map_location='cpu', weights_only=False)
state_dict = checkpoint.get('params_ema', checkpoint) if isinstance(checkpoint, dict) else checkpoint
net.load_state_dict(state_dict)
net.eval()
net.to(device)
guard_forward(net)
helper = face_helper(outscale)
@@ -567,7 +610,7 @@ def build_face_codeformer(spec, outscale, bg):
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)
tensor = tensor.unsqueeze(0).to(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))
@@ -733,6 +776,10 @@ def validate(model_name, face, face_name, strength, quality):
def preload():
try:
seed_weights()
except Exception as err:
log.warning('перенос весов из образа не удался: %s', err)
names = list(MODEL_REGISTRY) if LOAD_ALL else [DEFAULT_MODEL]
for name in names:
try: