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:
+50
-3
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user