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-кэш не даёт качать повторно, недоступная сеть на сборке не роняет образ.
85 lines
2.5 KiB
Python
85 lines
2.5 KiB
Python
import os
|
|
import sys
|
|
|
|
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
|
|
|
import app
|
|
|
|
WEIGHTS = {}
|
|
MODELS = {}
|
|
for _name, _spec in app.MODEL_REGISTRY.items():
|
|
WEIGHTS[_spec['file']] = _spec
|
|
MODELS[_name] = _spec
|
|
if _spec.get('dni'):
|
|
WEIGHTS[_spec['dni']['file']] = _spec['dni']
|
|
for _name, _spec in app.FACE_REGISTRY.items():
|
|
WEIGHTS[_spec['file']] = _spec
|
|
MODELS[_name] = _spec
|
|
for _name, _spec in app.FACEXLIB_WEIGHTS.items():
|
|
WEIGHTS[_name] = _spec
|
|
|
|
|
|
def wanted():
|
|
raw = os.environ.get('PHOTO_AI_PREFETCH', '').strip()
|
|
if not raw or raw.lower() == 'all':
|
|
return dict(WEIGHTS)
|
|
picked = {}
|
|
for key in raw.split(','):
|
|
name = key.strip()
|
|
if not name:
|
|
continue
|
|
if name.lower() == 'face':
|
|
for face_name, face_spec in app.FACE_REGISTRY.items():
|
|
picked[face_spec['file']] = face_spec
|
|
continue
|
|
spec = MODELS.get(name) or WEIGHTS.get(name)
|
|
if spec is None:
|
|
sys.stderr.write('неизвестная модель для предзагрузки: %s\n' % name)
|
|
continue
|
|
picked[spec['file']] = spec
|
|
return picked
|
|
|
|
|
|
def fetch(name, spec, target, cache_dir):
|
|
url = spec['url']
|
|
if cache_dir:
|
|
cached = os.path.join(cache_dir, name)
|
|
if app.file_ok(cached, spec['min_bytes']):
|
|
app.link_or_copy(cached, target)
|
|
print('веса %s взяты из кэша сборки %s' % (name, cached))
|
|
return
|
|
try:
|
|
app.download_weight(url, cached, spec['min_bytes'])
|
|
except Exception:
|
|
if os.path.isfile(cached):
|
|
os.unlink(cached)
|
|
raise
|
|
app.link_or_copy(cached, target)
|
|
return
|
|
app.download_weight(url, target, spec['min_bytes'])
|
|
|
|
|
|
def main():
|
|
target_dir = sys.argv[1] if len(sys.argv) > 1 else app.SEED_DIR
|
|
cache_dir = sys.argv[2] if len(sys.argv) > 2 else ''
|
|
if cache_dir:
|
|
os.makedirs(cache_dir, exist_ok=True)
|
|
os.makedirs(target_dir, exist_ok=True)
|
|
failed = []
|
|
for name, spec in wanted().items():
|
|
target = os.path.join(target_dir, name)
|
|
if app.file_ok(target, spec['min_bytes']):
|
|
continue
|
|
try:
|
|
fetch(name, spec, target, cache_dir)
|
|
except Exception as err:
|
|
failed.append('%s: %s' % (name, err))
|
|
continue
|
|
for line in failed:
|
|
sys.stderr.write('не удалось скачать %s\n' % line)
|
|
return 1 if failed else 0
|
|
|
|
|
|
if __name__ == '__main__':
|
|
sys.exit(main())
|