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
+5 -1
View File
@@ -69,8 +69,12 @@ PHOTO_AI_DEVICE=auto
# На GPU с 4 ГБ VRAM держите 256: при нехватке сервис сам пройдёт лестницу тайлов и уйдёт на CPU.
PHOTO_AI_TILE=256
# Модель восстановления лиц: gfpgan (по умолчанию) или codeformer.
# codeformer доступен только если модуль вендорен в photo-ai/vendor.
# Официальный модуль codeformer вендорен в photo-ai/vendor, доступен сразу.
PHOTO_AI_FACE_MODEL=gfpgan
# Какие веса photo-ai скачивает на этапе сборки образа (build-arg, не переменная контейнера):
# codeformer (по умолчанию) | face (codeformer + gfpgan) | all | none.
# Веса попадают в слой образа, при старте переносятся в том photo-ai-models и живут там.
PHOTO_AI_PREFETCH=codeformer
# Предзагрузка всех моделей при старте (1) или ленивая загрузка по требованию (0).
PHOTO_AI_LOAD_ALL=0
# Качество JPEG результата (70..100).
+4
View File
@@ -30,5 +30,9 @@ Thumbs.db
# --- Playwright test artifacts ---
.playwright-mcp/
# --- Python bytecode (photo-ai) ---
__pycache__/
*.pyc
# --- Internal audit (not for commit) ---
SECURITY_AUDIT.md
+28 -4
View File
@@ -132,7 +132,7 @@ REDIS_PASSWORD=случайная-длинная-строка
| `PHOTO_AI_MAX_PIXELS` | `4000000` | Максимум пикселей входного изображения, вход большего размера уменьшается |
| `PHOTO_AI_DEVICE` | `auto` | Устройство инференса: `auto` (CUDA, если контейнеру выдан GPU, иначе CPU), `cuda`, `cpu`. Явный `cuda` без CUDA не роняет сервис: WARN в лог и работа на CPU |
| `PHOTO_AI_TILE` | `256` | Размер тайла инференса (`0` — без тайлов): меньше тайл — меньше памяти, но медленнее |
| `PHOTO_AI_FACE_MODEL` | `gfpgan` | Модель восстановления лиц: `gfpgan`; `codeformer` доступен только при вендоринге модуля в `photo-ai/vendor` |
| `PHOTO_AI_FACE_MODEL` | `gfpgan` | Модель восстановления лиц: `gfpgan` или `codeformer` (официальный модуль вендорен в `photo-ai/vendor/codeformer`, доступен сразу) |
| `PHOTO_AI_LOAD_ALL` | `0` | Загружать все модели при старте (`1`) или лениво по требованию (`0`) |
| `PHOTO_AI_JPEG_QUALITY` | `92` | Качество JPEG результата, 70..100 |
| `PHOTO_AI_FACE_TIMEOUT_MS` | `600000` | Таймаут заданий с восстановлением лиц (мс) |
@@ -140,6 +140,32 @@ REDIS_PASSWORD=случайная-длинная-строка
| `PHOTO_AI_SOFT_BACKOFF_MS` | `10000` | Первая пауза перед мягким повтором |
| `PHOTO_AI_SOFT_BACKOFF_MAX_MS` | `300000` | Потолок паузы (задержка растёт вдвое) |
### Модели лиц и где лежат веса
Face-модели доступны обе: `gfpgan` (дефолт) и `codeformer`. Модуль `codeformer` — официальный
`sczhou/CodeFormer` (`b33cc7d`), вендорен в `photo-ai/vendor/codeformer/` (`codeformer_arch.py`
+ `vqgan_arch.py`, лицензия S-Lab 1.0 лежит рядом). Отдельный пакет с PyPI не используется: там лежит
сторонняя обёртка `rohitkhatri`, которая тянет свой `facelib` и `lpips`. Модуль попадает в образ
через `COPY vendor/` и подхватывается `sys.path` в `app.py` — правки Dockerfile не требуется.
Веса **не** лежат в репозитории и **не** скачиваются при первом запросе: на сборке образа
`fetch-weights.py` кладёт их в `/opt/photo-ai-seed`, а при старте сервис переносит их в том
`photo-ai-models:/models/weights` (`seed_weights()`). Дальше модель живёт в томе и переживает
пересборку образа; если её нет ни в томе, ни в образе, работает старый ленивый заозагрузчик.
Сам файл качается в BuildKit-кэш `/var/cache/photo-ai-weights`, поэтому повторная сборка
(и сборка с другим `PHOTO_AI_PREFETCH`) берёт его оттуда и заново не качает.
| Сборка | Что скачает |
|---|---|
| `docker compose build photo-ai` | `codeformer` (дефолт `PHOTO_AI_PREFETCH=codeformer`) |
| `PHOTO_AI_PREFETCH=face docker compose build photo-ai` | `codeformer` + `gfpgan` |
| `PHOTO_AI_PREFETCH=all docker compose build photo-ai` | всё: апскейлы, face-модели, веса facexlib |
| `PHOTO_AI_PREFETCH=none docker compose build photo-ai` | ничего, веса качаются лениво в том |
Переменная `PHOTO_AI_PREFETCH` — именно build-arg, он читается при сборке образа, а не контейнера
(в `.env.example` она есть, чтобы задать значение один раз). Скачивание при сборке не роняет образ:
при недоступной сети шаг пишет предупреждение, сервис докачает веса при первом использовании.
### Запуск на NVIDIA GPU
GPU не обязателен: без него сервис работает на CPU. Чтобы включить GPU-вариант, нужен драйвер NVIDIA
@@ -168,9 +194,7 @@ GPU выдаётся контейнеру ключом `gpus: all`, поэтом
генерации, поэтому после переподключения видеокарты или смены порта она начинает ссылаться на
несуществующий узел, и контейнер не стартует с `CDI device injection failed: failed to stat CDI host
device /dev/dri/cardN`. Перегенерация спеки требует sudo и теряется при каждой перегенерации;
`gpus: all` от этого свободен. Если nvidia-runtime уже зарегистрирован в
демоне (`nvidia-ctk runtime configure --runtime=docker`), в оверрайде можно вместо этого указать
`deploy.resources.reservations.devices` с `driver: nvidia, count: 1` — результат тот же.
`gpus: all` от этого свободен.
Проверка результата: в `/health` должны быть `device: cuda:0`, `half: true`, непустые
`vram_total_mb`/`vram_free_mb`. На 4 ГБ (например, RTX 3050 Laptop) реально держатся одновременно
+21 -12
View File
@@ -171,20 +171,29 @@
сейчас скачиваются мимо тома и теряются при пересборке. Свой `FaceRestoreHelper`/каталог
`/models/weights` — обязательное требование этапа.
- [x] **D5. CodeFormer — опционален по умолчанию.** Проверено 2026-09-28: на PyPI есть только
сторонняя обёртка `codeformer 0.0.11` (`github.com/rohitkhatri/codeformer`, тянет `lpips`) —
это не официальный `sczhou/CodeFormer`, использовать его не будем. Официальные веса доступны.
Решение: официальный модуль вендорится в `photo-ai/vendor/codeformer/` **только если** face-режим
CodeFormer реально понадобится; в рамках текущего плана не вендорим, дефолт
`PHOTO_AI_FACE_MODEL=gfpgan` (Stage 1–3), чтобы дефолтный путь работал без вендоринга.
- `FACE_REGISTRY` всегда содержит обе записи, но запись `codeformer` активна только если
`import codeformer` (с `photo-ai/vendor` в `sys.path`) успешен.
- Нет модуля → запись не попадает в `face_models` в `/health` и в список `/models`, UI её не показывает,
`POST /enhance` с `face_model=codeformer` → `400` с текстом «CodeFormer не установлен в образ,
доступен gfpgan». Сборка и CPU-режим не падают.
- [x] **D5. CodeFormer — вендорен, веса тянутся на сборке.** Первоначально (проверено 2026-09-28) на PyPI
есть только сторонняя обёртка `codeformer 0.0.11` (`github.com/rohitkhatri/codeformer`, тянет
`lpips`) — это не официальный `sczhou/CodeFormer`, использовать его не будем.
**Обновлено 2026-10-04:** официальный модуль вендорен в `photo-ai/vendor/codeformer/`
(`codeformer_arch.py` + `vqgan_arch.py`, коммит `b33cc7d`, лицензия S-Lab 1.0 рядом).
`basicsr==1.4.2` с PyPI не содержит ни `vqgan_arch.py`, ни `codeformer_arch.py`, поэтому
вендорится **оба** файла, а импорт в `codeformer_arch.py` переведён на относительный.
Дефолтом остаётся `PHOTO_AI_FACE_MODEL=gfpgan`, но `codeformer` теперь работает сразу.
- `FACE_REGISTRY` всегда содержит обе записи, запись `codeformer` активна, если `import codeformer`
(с `photo-ai/vendor` в `sys.path`) успешен.
- **`build_face_codeformer` пришлось чинить.** Код был написан по API обёртки rohitkhatri:
`CodeFormer(..., device=..., fp16=False)` и `net.device` официальный класс не принимает —
вызов падал бы с `TypeError`/`AttributeError`. Официальная сигнатура:
`(dim_embd, n_head, n_layers, codebook_size, latent_size, connect_list, fix_modules, vqgan_path)`.
Теперь устройство передаётся через `net.to(device)`. Чекпоунт официальных весов лежит под
ключом `params_ema` — учтено в `build_face_codeformer`.
- `strength` валиден только для CodeFormer: при `face_model=gfpgan` и `strength`, отличном от 0.7,
→ `400` с пояснением, иначе параметр молча игнорировался бы.
- Веса CodeFormer качаются тем же загрузчиком в `/models/weights/codeformer.pth`.
- **Веса больше не качаются лениво.** `photo-ai/fetch-weights.py` на этапе сборки кладёт их в
`/opt/photo-ai-seed` (build-arg `PHOTO_AI_PREFETCH`, дефолт `codeformer`; `face`/`all`/`none`),
при старте `seed_weights()` переносит их в том `photo-ai-models:/models/weights`. Том
переживает пересборку образа, BuildKit-кэш не даёт перекачивать. Если весов нет ни в томе,
ни в seed — работает прежний ленивый загрузчик.
- [x] **D6. Поведение без `photo-ai` в compose.** Сейчас `PHOTO_AI_URL` по умолчанию
`http://photo-ai:8080` (`docker-compose.yml:79`) — то есть «ИИ» включён по умолчанию.
+4 -1
View File
@@ -201,7 +201,10 @@ services:
# Порт 8081 пробрасывается только на loopback хоста — наружу ничего не публикуется,
# хостовый 8080 уже занят text-corrector. Ручные проверки: curl http://127.0.0.1:8081/health
photo-ai:
build: ./photo-ai
build:
context: ./photo-ai
args:
PHOTO_AI_PREFETCH: ${PHOTO_AI_PREFETCH:-codeformer}
container_name: photo-ai
restart: unless-stopped
environment:
+13
View File
@@ -1,3 +1,4 @@
# syntax=docker/dockerfile:1
FROM python:3.10-slim
ARG TORCH_VARIANT=cpu
@@ -25,10 +26,22 @@ RUN BASICSR_DEG=$(python -c "import basicsr; import os; print(os.path.join(os.pa
fi
COPY app.py ./
COPY fetch-weights.py ./
COPY vendor/ ./vendor/
ENV PHOTO_AI_MODELS_DIR=/models
ENV MODEL_PATH=/models/RealESRGAN_x2plus.pth
ENV PHOTO_AI_SEED_DIR=/opt/photo-ai-seed
ARG PHOTO_AI_PREFETCH=codeformer
RUN --mount=type=cache,target=/var/cache/photo-ai-weights,sharing=locked \
mkdir -p "$PHOTO_AI_SEED_DIR" && \
if [ -n "$PHOTO_AI_PREFETCH" ] && [ "$PHOTO_AI_PREFETCH" != "none" ]; then \
python fetch-weights.py "$PHOTO_AI_SEED_DIR" /var/cache/photo-ai-weights || echo "предзагрузка весов не удалась, сервис скачает их при первом запросе"; \
else \
echo "предзагрузка весов отключена (PHOTO_AI_PREFETCH=$PHOTO_AI_PREFETCH)"; \
fi
VOLUME /models
+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:
+84
View File
@@ -0,0 +1,84 @@
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())
+35
View File
@@ -0,0 +1,35 @@
S-Lab License 1.0
Copyright 2022 S-Lab
Redistribution and use for non-commercial purpose in source and
binary forms, with or without modification, are permitted provided
that the following conditions are met:
1. Redistributions of source code must retain the above copyright
notice, this list of conditions and the following disclaimer.
2. Redistributions in binary form must reproduce the above copyright
notice, this list of conditions and the following disclaimer in
the documentation and/or other materials provided with the
distribution.
3. Neither the name of the copyright holder nor the names of its
contributors may be used to endorse or promote products derived
from this software without specific prior written permission.
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS
"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT
LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR
A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT
HOLDER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL,
SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT
LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE,
DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY
THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT
(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
In the event that redistribution and/or use for commercial purpose in
source or binary forms, with or without modification is required,
please contact the contributor(s) of the work.
+3
View File
@@ -0,0 +1,3 @@
from .codeformer_arch import CodeFormer
__all__ = ['CodeFormer']
+277
View File
@@ -0,0 +1,277 @@
import math
import numpy as np
import torch
from torch import nn, Tensor
import torch.nn.functional as F
from typing import Optional, List
from .vqgan_arch import *
def calc_mean_std(feat, eps=1e-5):
"""Calculate mean and std for adaptive_instance_normalization.
Args:
feat (Tensor): 4D tensor.
eps (float): A small value added to the variance to avoid
divide-by-zero. Default: 1e-5.
"""
size = feat.size()
assert len(size) == 4, 'The input feature should be 4D tensor.'
b, c = size[:2]
feat_var = feat.view(b, c, -1).var(dim=2) + eps
feat_std = feat_var.sqrt().view(b, c, 1, 1)
feat_mean = feat.view(b, c, -1).mean(dim=2).view(b, c, 1, 1)
return feat_mean, feat_std
def adaptive_instance_normalization(content_feat, style_feat):
"""Adaptive instance normalization.
Adjust the reference features to have the similar color and illuminations
as those in the degradate features.
Args:
content_feat (Tensor): The reference feature.
style_feat (Tensor): The degradate features.
"""
size = content_feat.size()
style_mean, style_std = calc_mean_std(style_feat)
content_mean, content_std = calc_mean_std(content_feat)
normalized_feat = (content_feat - content_mean.expand(size)) / content_std.expand(size)
return normalized_feat * style_std.expand(size) + style_mean.expand(size)
class PositionEmbeddingSine(nn.Module):
"""
This is a more standard version of the position embedding, very similar to the one
used by the Attention is all you need paper, generalized to work on images.
"""
def __init__(self, num_pos_feats=64, temperature=10000, normalize=False, scale=None):
super().__init__()
self.num_pos_feats = num_pos_feats
self.temperature = temperature
self.normalize = normalize
if scale is not None and normalize is False:
raise ValueError("normalize should be True if scale is passed")
if scale is None:
scale = 2 * math.pi
self.scale = scale
def forward(self, x, mask=None):
if mask is None:
mask = torch.zeros((x.size(0), x.size(2), x.size(3)), device=x.device, dtype=torch.bool)
not_mask = ~mask
y_embed = not_mask.cumsum(1, dtype=torch.float32)
x_embed = not_mask.cumsum(2, dtype=torch.float32)
if self.normalize:
eps = 1e-6
y_embed = y_embed / (y_embed[:, -1:, :] + eps) * self.scale
x_embed = x_embed / (x_embed[:, :, -1:] + eps) * self.scale
dim_t = torch.arange(self.num_pos_feats, dtype=torch.float32, device=x.device)
dim_t = self.temperature ** (2 * (dim_t // 2) / self.num_pos_feats)
pos_x = x_embed[:, :, :, None] / dim_t
pos_y = y_embed[:, :, :, None] / dim_t
pos_x = torch.stack(
(pos_x[:, :, :, 0::2].sin(), pos_x[:, :, :, 1::2].cos()), dim=4
).flatten(3)
pos_y = torch.stack(
(pos_y[:, :, :, 0::2].sin(), pos_y[:, :, :, 1::2].cos()), dim=4
).flatten(3)
pos = torch.cat((pos_y, pos_x), dim=3).permute(0, 3, 1, 2)
return pos
def _get_activation_fn(activation):
"""Return an activation function given a string"""
if activation == "relu":
return F.relu
if activation == "gelu":
return F.gelu
if activation == "glu":
return F.glu
raise RuntimeError(F"activation should be relu/gelu, not {activation}.")
class TransformerSALayer(nn.Module):
def __init__(self, embed_dim, nhead=8, dim_mlp=2048, dropout=0.0, activation="gelu"):
super().__init__()
self.self_attn = nn.MultiheadAttention(embed_dim, nhead, dropout=dropout)
# Implementation of Feedforward model - MLP
self.linear1 = nn.Linear(embed_dim, dim_mlp)
self.dropout = nn.Dropout(dropout)
self.linear2 = nn.Linear(dim_mlp, embed_dim)
self.norm1 = nn.LayerNorm(embed_dim)
self.norm2 = nn.LayerNorm(embed_dim)
self.dropout1 = nn.Dropout(dropout)
self.dropout2 = nn.Dropout(dropout)
self.activation = _get_activation_fn(activation)
def with_pos_embed(self, tensor, pos: Optional[Tensor]):
return tensor if pos is None else tensor + pos
def forward(self, tgt,
tgt_mask: Optional[Tensor] = None,
tgt_key_padding_mask: Optional[Tensor] = None,
query_pos: Optional[Tensor] = None):
# self attention
tgt2 = self.norm1(tgt)
q = k = self.with_pos_embed(tgt2, query_pos)
tgt2 = self.self_attn(q, k, value=tgt2, attn_mask=tgt_mask,
key_padding_mask=tgt_key_padding_mask)[0]
tgt = tgt + self.dropout1(tgt2)
# ffn
tgt2 = self.norm2(tgt)
tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt2))))
tgt = tgt + self.dropout2(tgt2)
return tgt
class Fuse_sft_block(nn.Module):
def __init__(self, in_ch, out_ch):
super().__init__()
self.encode_enc = ResBlock(2*in_ch, out_ch)
self.scale = nn.Sequential(
nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),
nn.LeakyReLU(0.2, True),
nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1))
self.shift = nn.Sequential(
nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),
nn.LeakyReLU(0.2, True),
nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1))
def forward(self, enc_feat, dec_feat, w=1):
enc_feat = self.encode_enc(torch.cat([enc_feat, dec_feat], dim=1))
scale = self.scale(enc_feat)
shift = self.shift(enc_feat)
residual = w * (dec_feat * scale + shift)
out = dec_feat + residual
return out
class CodeFormer(VQAutoEncoder):
def __init__(self, dim_embd=512, n_head=8, n_layers=9,
codebook_size=1024, latent_size=256,
connect_list=['32', '64', '128', '256'],
fix_modules=['quantize','generator'], vqgan_path=None):
super(CodeFormer, self).__init__(512, 64, [1, 2, 2, 4, 4, 8], 'nearest',2, [16], codebook_size)
if vqgan_path is not None:
self.load_state_dict(
torch.load(vqgan_path, map_location='cpu')['params_ema'])
if fix_modules is not None:
for module in fix_modules:
for param in getattr(self, module).parameters():
param.requires_grad = False
self.connect_list = connect_list
self.n_layers = n_layers
self.dim_embd = dim_embd
self.dim_mlp = dim_embd*2
self.position_emb = nn.Parameter(torch.zeros(latent_size, self.dim_embd))
self.feat_emb = nn.Linear(256, self.dim_embd)
# transformer
self.ft_layers = nn.Sequential(*[TransformerSALayer(embed_dim=dim_embd, nhead=n_head, dim_mlp=self.dim_mlp, dropout=0.0)
for _ in range(self.n_layers)])
# logits_predict head
self.idx_pred_layer = nn.Sequential(
nn.LayerNorm(dim_embd),
nn.Linear(dim_embd, codebook_size, bias=False))
self.channels = {
'16': 512,
'32': 256,
'64': 256,
'128': 128,
'256': 128,
'512': 64,
}
# after second residual block for > 16, before attn layer for ==16
self.fuse_encoder_block = {'512':2, '256':5, '128':8, '64':11, '32':14, '16':18}
# after first residual block for > 16, before attn layer for ==16
self.fuse_generator_block = {'16':6, '32': 9, '64':12, '128':15, '256':18, '512':21}
# fuse_convs_dict
self.fuse_convs_dict = nn.ModuleDict()
for f_size in self.connect_list:
in_ch = self.channels[f_size]
self.fuse_convs_dict[f_size] = Fuse_sft_block(in_ch, in_ch)
def _init_weights(self, module):
if isinstance(module, (nn.Linear, nn.Embedding)):
module.weight.data.normal_(mean=0.0, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
module.bias.data.zero_()
elif isinstance(module, nn.LayerNorm):
module.bias.data.zero_()
module.weight.data.fill_(1.0)
def forward(self, x, w=0, detach_16=True, code_only=False, adain=False):
# ################### Encoder #####################
enc_feat_dict = {}
out_list = [self.fuse_encoder_block[f_size] for f_size in self.connect_list]
for i, block in enumerate(self.encoder.blocks):
x = block(x)
if i in out_list:
enc_feat_dict[str(x.shape[-1])] = x.clone()
lq_feat = x
# ################# Transformer ###################
# quant_feat, codebook_loss, quant_stats = self.quantize(lq_feat)
pos_emb = self.position_emb.unsqueeze(1).repeat(1,x.shape[0],1)
# BCHW -> BC(HW) -> (HW)BC
feat_emb = self.feat_emb(lq_feat.flatten(2).permute(2,0,1))
query_emb = feat_emb
# Transformer encoder
for layer in self.ft_layers:
query_emb = layer(query_emb, query_pos=pos_emb)
# output logits
logits = self.idx_pred_layer(query_emb) # (hw)bn
logits = logits.permute(1,0,2) # (hw)bn -> b(hw)n
if code_only: # for training stage II
# logits doesn't need softmax before cross_entropy loss
return logits, lq_feat
# ################# Quantization ###################
# if self.training:
# quant_feat = torch.einsum('btn,nc->btc', [soft_one_hot, self.quantize.embedding.weight])
# # b(hw)c -> bc(hw) -> bchw
# quant_feat = quant_feat.permute(0,2,1).view(lq_feat.shape)
# ------------
soft_one_hot = F.softmax(logits, dim=2)
_, top_idx = torch.topk(soft_one_hot, 1, dim=2)
quant_feat = self.quantize.get_codebook_feat(top_idx, shape=[x.shape[0],16,16,256])
# preserve gradients
# quant_feat = lq_feat + (quant_feat - lq_feat).detach()
if detach_16:
quant_feat = quant_feat.detach() # for training stage III
if adain:
quant_feat = adaptive_instance_normalization(quant_feat, lq_feat)
# ################## Generator ####################
x = quant_feat
fuse_list = [self.fuse_generator_block[f_size] for f_size in self.connect_list]
for i, block in enumerate(self.generator.blocks):
x = block(x)
if i in fuse_list: # fuse after i-th block
f_size = str(x.shape[-1])
if w>0:
x = self.fuse_convs_dict[f_size](enc_feat_dict[f_size].detach(), x, w)
out = x
# logits doesn't need softmax before cross_entropy loss
return out, logits, lq_feat
+431
View File
@@ -0,0 +1,431 @@
'''
VQGAN code, adapted from the original created by the Unleashing Transformers authors:
https://github.com/samb-t/unleashing-transformers/blob/master/models/vqgan.py
'''
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import copy
from basicsr.utils import get_root_logger
def normalize(in_channels):
return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
@torch.jit.script
def swish(x):
return x*torch.sigmoid(x)
# Define VQVAE classes
class VectorQuantizer(nn.Module):
def __init__(self, codebook_size, emb_dim, beta):
super(VectorQuantizer, self).__init__()
self.codebook_size = codebook_size # number of embeddings
self.emb_dim = emb_dim # dimension of embedding
self.beta = beta # commitment cost used in loss term, beta * ||z_e(x)-sg[e]||^2
self.embedding = nn.Embedding(self.codebook_size, self.emb_dim)
self.embedding.weight.data.uniform_(-1.0 / self.codebook_size, 1.0 / self.codebook_size)
def forward(self, z):
# reshape z -> (batch, height, width, channel) and flatten
z = z.permute(0, 2, 3, 1).contiguous()
z_flattened = z.view(-1, self.emb_dim)
# distances from z to embeddings e_j (z - e)^2 = z^2 + e^2 - 2 e * z
d = (z_flattened ** 2).sum(dim=1, keepdim=True) + (self.embedding.weight**2).sum(1) - \
2 * torch.matmul(z_flattened, self.embedding.weight.t())
mean_distance = torch.mean(d)
# find closest encodings
min_encoding_indices = torch.argmin(d, dim=1).unsqueeze(1)
# min_encoding_scores, min_encoding_indices = torch.topk(d, 1, dim=1, largest=False)
# [0-1], higher score, higher confidence
# min_encoding_scores = torch.exp(-min_encoding_scores/10)
min_encodings = torch.zeros(min_encoding_indices.shape[0], self.codebook_size).to(z)
min_encodings.scatter_(1, min_encoding_indices, 1)
# get quantized latent vectors
z_q = torch.matmul(min_encodings, self.embedding.weight).view(z.shape)
# compute loss for embedding
loss = torch.mean((z_q.detach()-z)**2) + self.beta * torch.mean((z_q - z.detach()) ** 2)
# preserve gradients
z_q = z + (z_q - z).detach()
# perplexity
e_mean = torch.mean(min_encodings, dim=0)
perplexity = torch.exp(-torch.sum(e_mean * torch.log(e_mean + 1e-10)))
# reshape back to match original input shape
z_q = z_q.permute(0, 3, 1, 2).contiguous()
return z_q, loss, {
"perplexity": perplexity,
"min_encodings": min_encodings,
"min_encoding_indices": min_encoding_indices,
"mean_distance": mean_distance
}
def get_codebook_feat(self, indices, shape):
# input indices: batch*token_num -> (batch*token_num)*1
# shape: batch, height, width, channel
indices = indices.view(-1,1)
min_encodings = torch.zeros(indices.shape[0], self.codebook_size).to(indices)
min_encodings.scatter_(1, indices, 1)
# get quantized latent vectors
z_q = torch.matmul(min_encodings.float(), self.embedding.weight)
if shape is not None: # reshape back to match original input shape
z_q = z_q.view(shape).permute(0, 3, 1, 2).contiguous()
return z_q
class GumbelQuantizer(nn.Module):
def __init__(self, codebook_size, emb_dim, num_hiddens, straight_through=False, kl_weight=5e-4, temp_init=1.0):
super().__init__()
self.codebook_size = codebook_size # number of embeddings
self.emb_dim = emb_dim # dimension of embedding
self.straight_through = straight_through
self.temperature = temp_init
self.kl_weight = kl_weight
self.proj = nn.Conv2d(num_hiddens, codebook_size, 1) # projects last encoder layer to quantized logits
self.embed = nn.Embedding(codebook_size, emb_dim)
def forward(self, z):
hard = self.straight_through if self.training else True
logits = self.proj(z)
soft_one_hot = F.gumbel_softmax(logits, tau=self.temperature, dim=1, hard=hard)
z_q = torch.einsum("b n h w, n d -> b d h w", soft_one_hot, self.embed.weight)
# + kl divergence to the prior loss
qy = F.softmax(logits, dim=1)
diff = self.kl_weight * torch.sum(qy * torch.log(qy * self.codebook_size + 1e-10), dim=1).mean()
min_encoding_indices = soft_one_hot.argmax(dim=1)
return z_q, diff, {
"min_encoding_indices": min_encoding_indices
}
class Downsample(nn.Module):
def __init__(self, in_channels):
super().__init__()
self.conv = torch.nn.Conv2d(in_channels, in_channels, kernel_size=3, stride=2, padding=0)
def forward(self, x):
pad = (0, 1, 0, 1)
x = torch.nn.functional.pad(x, pad, mode="constant", value=0)
x = self.conv(x)
return x
class Upsample(nn.Module):
def __init__(self, in_channels):
super().__init__()
self.conv = nn.Conv2d(in_channels, in_channels, kernel_size=3, stride=1, padding=1)
def forward(self, x):
x = F.interpolate(x, scale_factor=2.0, mode="nearest")
x = self.conv(x)
return x
class ResBlock(nn.Module):
def __init__(self, in_channels, out_channels=None):
super(ResBlock, self).__init__()
self.in_channels = in_channels
self.out_channels = in_channels if out_channels is None else out_channels
self.norm1 = normalize(in_channels)
self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
self.norm2 = normalize(out_channels)
self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
if self.in_channels != self.out_channels:
self.conv_out = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
def forward(self, x_in):
x = x_in
x = self.norm1(x)
x = swish(x)
x = self.conv1(x)
x = self.norm2(x)
x = swish(x)
x = self.conv2(x)
if self.in_channels != self.out_channels:
x_in = self.conv_out(x_in)
return x + x_in
class AttnBlock(nn.Module):
def __init__(self, in_channels):
super().__init__()
self.in_channels = in_channels
self.norm = normalize(in_channels)
self.q = torch.nn.Conv2d(
in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0
)
self.k = torch.nn.Conv2d(
in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0
)
self.v = torch.nn.Conv2d(
in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0
)
self.proj_out = torch.nn.Conv2d(
in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0
)
def forward(self, x):
h_ = x
h_ = self.norm(h_)
q = self.q(h_)
k = self.k(h_)
v = self.v(h_)
# compute attention
b, c, h, w = q.shape
q = q.reshape(b, c, h*w)
q = q.permute(0, 2, 1)
k = k.reshape(b, c, h*w)
w_ = torch.bmm(q, k)
w_ = w_ * (int(c)**(-0.5))
w_ = F.softmax(w_, dim=2)
# attend to values
v = v.reshape(b, c, h*w)
w_ = w_.permute(0, 2, 1)
h_ = torch.bmm(v, w_)
h_ = h_.reshape(b, c, h, w)
h_ = self.proj_out(h_)
return x+h_
class Encoder(nn.Module):
def __init__(self, in_channels, nf, emb_dim, ch_mult, num_res_blocks, resolution, attn_resolutions):
super().__init__()
self.nf = nf
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.resolution = resolution
self.attn_resolutions = attn_resolutions
curr_res = self.resolution
in_ch_mult = (1,)+tuple(ch_mult)
blocks = []
# initial convultion
blocks.append(nn.Conv2d(in_channels, nf, kernel_size=3, stride=1, padding=1))
# residual and downsampling blocks, with attention on smaller res (16x16)
for i in range(self.num_resolutions):
block_in_ch = nf * in_ch_mult[i]
block_out_ch = nf * ch_mult[i]
for _ in range(self.num_res_blocks):
blocks.append(ResBlock(block_in_ch, block_out_ch))
block_in_ch = block_out_ch
if curr_res in attn_resolutions:
blocks.append(AttnBlock(block_in_ch))
if i != self.num_resolutions - 1:
blocks.append(Downsample(block_in_ch))
curr_res = curr_res // 2
# non-local attention block
blocks.append(ResBlock(block_in_ch, block_in_ch))
blocks.append(AttnBlock(block_in_ch))
blocks.append(ResBlock(block_in_ch, block_in_ch))
# normalise and convert to latent size
blocks.append(normalize(block_in_ch))
blocks.append(nn.Conv2d(block_in_ch, emb_dim, kernel_size=3, stride=1, padding=1))
self.blocks = nn.ModuleList(blocks)
def forward(self, x):
for block in self.blocks:
x = block(x)
return x
class Generator(nn.Module):
def __init__(self, nf, emb_dim, ch_mult, res_blocks, img_size, attn_resolutions):
super().__init__()
self.nf = nf
self.ch_mult = ch_mult
self.num_resolutions = len(self.ch_mult)
self.num_res_blocks = res_blocks
self.resolution = img_size
self.attn_resolutions = attn_resolutions
self.in_channels = emb_dim
self.out_channels = 3
block_in_ch = self.nf * self.ch_mult[-1]
curr_res = self.resolution // 2 ** (self.num_resolutions-1)
blocks = []
# initial conv
blocks.append(nn.Conv2d(self.in_channels, block_in_ch, kernel_size=3, stride=1, padding=1))
# non-local attention block
blocks.append(ResBlock(block_in_ch, block_in_ch))
blocks.append(AttnBlock(block_in_ch))
blocks.append(ResBlock(block_in_ch, block_in_ch))
for i in reversed(range(self.num_resolutions)):
block_out_ch = self.nf * self.ch_mult[i]
for _ in range(self.num_res_blocks):
blocks.append(ResBlock(block_in_ch, block_out_ch))
block_in_ch = block_out_ch
if curr_res in self.attn_resolutions:
blocks.append(AttnBlock(block_in_ch))
if i != 0:
blocks.append(Upsample(block_in_ch))
curr_res = curr_res * 2
blocks.append(normalize(block_in_ch))
blocks.append(nn.Conv2d(block_in_ch, self.out_channels, kernel_size=3, stride=1, padding=1))
self.blocks = nn.ModuleList(blocks)
def forward(self, x):
for block in self.blocks:
x = block(x)
return x
class VQAutoEncoder(nn.Module):
def __init__(self, img_size, nf, ch_mult, quantizer="nearest", res_blocks=2, attn_resolutions=[16], codebook_size=1024, emb_dim=256,
beta=0.25, gumbel_straight_through=False, gumbel_kl_weight=1e-8, model_path=None):
super().__init__()
logger = get_root_logger()
self.in_channels = 3
self.nf = nf
self.n_blocks = res_blocks
self.codebook_size = codebook_size
self.embed_dim = emb_dim
self.ch_mult = ch_mult
self.resolution = img_size
self.attn_resolutions = attn_resolutions
self.quantizer_type = quantizer
self.encoder = Encoder(
self.in_channels,
self.nf,
self.embed_dim,
self.ch_mult,
self.n_blocks,
self.resolution,
self.attn_resolutions
)
if self.quantizer_type == "nearest":
self.beta = beta #0.25
self.quantize = VectorQuantizer(self.codebook_size, self.embed_dim, self.beta)
elif self.quantizer_type == "gumbel":
self.gumbel_num_hiddens = emb_dim
self.straight_through = gumbel_straight_through
self.kl_weight = gumbel_kl_weight
self.quantize = GumbelQuantizer(
self.codebook_size,
self.embed_dim,
self.gumbel_num_hiddens,
self.straight_through,
self.kl_weight
)
self.generator = Generator(
self.nf,
self.embed_dim,
self.ch_mult,
self.n_blocks,
self.resolution,
self.attn_resolutions
)
if model_path is not None:
chkpt = torch.load(model_path, map_location='cpu')
if 'params_ema' in chkpt:
self.load_state_dict(torch.load(model_path, map_location='cpu')['params_ema'])
logger.info(f'vqgan is loaded from: {model_path} [params_ema]')
elif 'params' in chkpt:
self.load_state_dict(torch.load(model_path, map_location='cpu')['params'])
logger.info(f'vqgan is loaded from: {model_path} [params]')
else:
raise ValueError(f'Wrong params!')
def forward(self, x):
x = self.encoder(x)
quant, codebook_loss, quant_stats = self.quantize(x)
x = self.generator(quant)
return x, codebook_loss, quant_stats
# patch based discriminator
class VQGANDiscriminator(nn.Module):
def __init__(self, nc=3, ndf=64, n_layers=4, model_path=None):
super().__init__()
layers = [nn.Conv2d(nc, ndf, kernel_size=4, stride=2, padding=1), nn.LeakyReLU(0.2, True)]
ndf_mult = 1
ndf_mult_prev = 1
for n in range(1, n_layers): # gradually increase the number of filters
ndf_mult_prev = ndf_mult
ndf_mult = min(2 ** n, 8)
layers += [
nn.Conv2d(ndf * ndf_mult_prev, ndf * ndf_mult, kernel_size=4, stride=2, padding=1, bias=False),
nn.BatchNorm2d(ndf * ndf_mult),
nn.LeakyReLU(0.2, True)
]
ndf_mult_prev = ndf_mult
ndf_mult = min(2 ** n_layers, 8)
layers += [
nn.Conv2d(ndf * ndf_mult_prev, ndf * ndf_mult, kernel_size=4, stride=1, padding=1, bias=False),
nn.BatchNorm2d(ndf * ndf_mult),
nn.LeakyReLU(0.2, True)
]
layers += [
nn.Conv2d(ndf * ndf_mult, 1, kernel_size=4, stride=1, padding=1)] # output 1 channel prediction map
self.main = nn.Sequential(*layers)
if model_path is not None:
chkpt = torch.load(model_path, map_location='cpu')
if 'params_d' in chkpt:
self.load_state_dict(torch.load(model_path, map_location='cpu')['params_d'])
elif 'params' in chkpt:
self.load_state_dict(torch.load(model_path, map_location='cpu')['params'])
else:
raise ValueError(f'Wrong params!')
def forward(self, x):
return self.main(x)