File size: 14,855 Bytes
ba5f45d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4b4189f
ba5f45d
 
 
 
 
 
 
ff1dbf3
ba5f45d
4b4189f
 
 
ba5f45d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ce81410
 
 
 
ba5f45d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ff1dbf3
 
 
 
ba5f45d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e279cee
 
 
 
4b4189f
 
 
 
 
 
 
 
 
 
 
ba5f45d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4b4189f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ce81410
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e279cee
 
4b4189f
 
 
e279cee
4b4189f
ce81410
 
 
 
 
e279cee
4b4189f
e279cee
 
 
 
 
 
 
4b4189f
 
 
 
 
e279cee
ba5f45d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4b4189f
ba5f45d
 
 
 
 
 
 
 
 
 
 
 
4b4189f
ba5f45d
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
"""
Custom inference handler для HF Inference Endpoints: Wan 2.2 A14B (Lightning-merged, diffusers).

Один и тот же файл кладётся в оба репозитория — t2v и i2v. Какой пайплайн поднимать,
определяется по model_index.json самого репозитория (_class_name), так что копия
идентичная и не нужно помнить, какой файл куда.

Вход (JSON, принимаются и плоская форма, и вложенная в inputs/parameters):
    {
      "inputs": "промпт",
      "parameters": {
        "negative_prompt": "...",
        "width": 1280, "height": 720, "num_frames": 81, "fps": 16,
        "num_inference_steps": 8, "guidance_scale": 1.0, "guidance_scale_2": 1.0,
        "flow_shift": 5.0, "seed": 123
      },
      "image": "<base64>"        # только для i2v
    }

Выход:
    {"video_base64": "...", "content_type": "video/mp4", "info": {...}}
"""

import base64
import io
import json
import logging
import os
import sys
import tempfile
import time
from pathlib import Path

import torch
from diffusers import AutoencoderKLWan, UniPCMultistepScheduler
from diffusers.utils import export_to_video
from PIL import Image, ImageOps

sys.path.insert(0, str(Path(__file__).resolve().parent))  # wan_budget лежит рядом в репозитории
import wan_budget  # noqa: E402

log = logging.getLogger(__name__)
logging.basicConfig(level=logging.INFO)

# Дефолты под Lightning-дистилляцию: 4 шага на эксперта, без CFG.
DEFAULTS = {
    "width": 1280,
    "height": 720,
    "num_frames": 81,
    "fps": 16,
    "num_inference_steps": 8,
    "guidance_scale": 1.0,
    "guidance_scale_2": 1.0,
    "flow_shift": 5.0,
    "max_sequence_length": 512,
}

# Во сколько секунд обходится переезд весов между CPU и GPU. Замерено грубо:
# первая генерация после смены режима идёт примерно на это дольше модели.
SWITCH_COST_S = 25.0

NEGATIVE_DEFAULT = (
    "色调艳丽, 过曝, 静态, 细节模糊不清, 字幕, 风格, 作品, 画作, 画面, 静止, 整体发灰, "
    "最差质量, 低质量, JPEG压缩残留, 丑陋的, 残缺的, 多余的手指, 画得不好的手部, "
    "画得不好的脸部, 畸形的, 毁容的, 形态畸形的肢体, 手指融合, 静止不动的画面, "
    "杂乱的背景, 三条腿, 背景人很多, 倒着走"
)


def _snap(value, step=16, minimum=16):
    return max(minimum, int(round(float(value) / step)) * step)


def _snap_frames(value):
    """Wan сжимает время в 4 раза — длина обязана быть 4n+1."""
    n = max(5, int(value))
    return ((n - 1) // 4) * 4 + 1


def _decode_image(raw):
    if isinstance(raw, str) and raw.startswith("data:"):
        raw = raw.split(",", 1)[1]
    data = base64.b64decode(raw) if isinstance(raw, str) else raw
    # exif_transpose обязателен: телефонные фото хранятся «боком» с флагом поворота,
    # обычный open() его игнорирует и видео выходит повёрнутым. Клиент это уже
    # чинит у себя, но эндпоинт принимает запросы и напрямую — страхуемся здесь тоже.
    return ImageOps.exif_transpose(Image.open(io.BytesIO(data))).convert("RGB")


class EndpointHandler:
    def __init__(self, path=""):
        path = path or "."
        cls_name = json.loads((Path(path) / "model_index.json").read_text())["_class_name"]
        self.is_i2v = cls_name == "WanImageToVideoPipeline"

        import diffusers
        pipe_cls = getattr(diffusers, cls_name)
        log.info("Загружаю %s из %s", cls_name, path)

        # VAE в fp32 — так рекомендует diffusers для Wan, в bf16 вылезают артефакты декода.
        vae = AutoencoderKLWan.from_pretrained(path, subfolder="vae", torch_dtype=torch.float32)
        self.pipe = pipe_cls.from_pretrained(path, vae=vae, torch_dtype=torch.bfloat16)

        # Обе экспертные модели в bf16 — это ~69 ГБ. На карте от 80 ГБ держим всё
        # резидентно, иначе просим diffusers жонглировать модулями через CPU.
        # ВНИМАНИЕ на единицы: CUDA рапортует байты, и делить их надо на 1024**3 (GiB),
        # а не на 1e9. «A100 80GB» — это 79.25 GiB, но 85.9 «гигабайт» в десятичном
        # счёте, из-за чего порог «меньше 80» не срабатывал и модель ехала на GPU
        # целиком: 77.3 из 79.25 GiB, то есть OOM на любом запросе крупнее 720p/81.
        self.vram = torch.cuda.get_device_properties(0).total_memory / 1024**3 if torch.cuda.is_available() else 0
        # WAN_OFFLOAD: on / off — жёстко зафиксировать режим, auto — выбирать по
        # каждому запросу исходя из того, помещаются ли его активации (см. wan_budget).
        self.forced = os.environ.get("WAN_OFFLOAD", "auto").lower()
        log.info("GPU: %s, %.1f GiB | режим=%s | резидентный потолок=%s токенов",
                 torch.cuda.get_device_name(0) if self.vram else "нет", self.vram,
                 self.forced, wan_budget.max_resident_tokens(self.vram))
        # Стартуем с выгрузки: она всегда безопасна, а поднять веса на GPU
        # дешевле, чем словить OOM на первом же крупном запросе.
        self._offloaded = None
        self._set_mode("offload" if self.forced != "off" else "resident")

        # Тайлинг VAE — декод 81 кадра в 720p иначе даёт пик в несколько десятков ГБ.
        for mod in (self.pipe.vae,):
            if hasattr(mod, "enable_tiling"):
                mod.enable_tiling()

        self._base_scheduler_config = dict(self.pipe.scheduler.config)
        self._flow_shift = self._base_scheduler_config.get("flow_shift")
        log.info("Готов. i2v=%s, flow_shift по умолчанию=%s", self.is_i2v, self._flow_shift)

    def _set_flow_shift(self, shift):
        """flow_shift живёт в планировщике, а не в аргументах вызова — пересобираем при смене."""
        if shift is None or float(shift) == float(self._flow_shift):
            return
        self.pipe.scheduler = UniPCMultistepScheduler.from_config(
            self._base_scheduler_config, flow_shift=float(shift))
        self._flow_shift = float(shift)

    def _set_mode(self, mode):
        """
        Переключает размещение весов. resident — всё на GPU, быстрее примерно
        на четверть; offload — diffusers держит на GPU только активный модуль.
        enable_model_cpu_offload() сам снимает прежние хуки, обратный переход
        делаем вручную через remove_all_hooks() + .to('cuda').
        """
        want_offload = mode == "offload"
        if self._offloaded == want_offload:
            return
        if want_offload:
            self.pipe.enable_model_cpu_offload()
        else:
            self.pipe.remove_all_hooks()
            self.pipe.to("cuda")
        self._offloaded = want_offload
        torch.cuda.empty_cache()
        log.info("Режим весов -> %s", mode)

    def _plan(self, width, height, num_frames):
        """Какой режим брать под этот запрос и почему — уходит в info для UI."""
        if self.forced in ("on", "off"):
            mode = "offload" if self.forced == "on" else "resident"
            info = {"tokens": wan_budget.latent_tokens(width, height, num_frames),
                    "resident_cap": wan_budget.max_resident_tokens(self.vram),
                    "forced": self.forced}
            return mode, info
        return wan_budget.pick_mode(width, height, num_frames, self.vram)

    def _worth_switching(self, target, width, height, num_frames, steps):
        """
        Переезд весов между CPU и GPU — это ~65 GiB по PCIe, десятки секунд.
        На коротком клипе экономия от резидентного режима такой переезд не
        окупает, поэтому вверх переключаемся только когда выигрыш ощутим.
        Вниз (в выгрузку) переключаемся всегда: там вопрос не скорости, а OOM.
        """
        if self._offloaded is None or target == "offload":
            return True
        current = "offload" if self._offloaded else "resident"
        if current == target:
            return False
        gain = (wan_budget.estimate_seconds(width, height, num_frames, "offload", steps)
                - wan_budget.estimate_seconds(width, height, num_frames, "resident", steps))
        return gain > SWITCH_COST_S

    def _run(self, call, width, height, num_frames):
        """
        Запускает пайплайн в выбранном режиме. На OOM один раз переходит на
        выгрузку и повторяет: лучше отдать видео медленнее, чем 400-ю ошибку.
        Если не помогло — объясняем, что уменьшить, вместо простыни от CUDA.
        """
        mode, plan = self._plan(width, height, num_frames)
        if self._worth_switching(mode, width, height, num_frames, call["num_inference_steps"]):
            self._set_mode(mode)
        else:
            mode = "offload" if self._offloaded else "resident"
            plan["kept_mode"] = True   # переключение не окупалось, остались как были
        try:
            return self.pipe(**call).frames[0], mode, plan
        except torch.cuda.OutOfMemoryError:
            if self._offloaded:
                raise RuntimeError(
                    f"Не хватило видеопамяти на {width}×{height}, {num_frames} кадров даже "
                    f"с выгрузкой на CPU. Уменьши число кадров или разрешение "
                    f"(832×480 требует примерно втрое меньше, чем 1280×720)."
                ) from None
            log.warning("OOM в резидентном режиме на %s токенах — выгружаю и повторяю",
                        plan.get("tokens"))
            self._set_mode("offload")
            plan["fell_back"] = True   # значит модель памяти оптимистична, стоит подкрутить
            return self.pipe(**call).frames[0], "offload", plan

    def __call__(self, data):
        started = time.time()
        if not isinstance(data, dict):
            data = {"inputs": data}

        params = dict(DEFAULTS)
        params.update({k: v for k, v in (data.get("parameters") or {}).items() if v is not None})

        inputs = data.get("inputs")
        if isinstance(inputs, dict):  # клиент мог сложить всё внутрь inputs
            params.update({k: v for k, v in inputs.items()
                           if k not in ("prompt", "image") and v is not None})
            prompt = inputs.get("prompt") or ""
            image_raw = inputs.get("image") or data.get("image")
        else:
            prompt = inputs or data.get("prompt") or ""
            image_raw = data.get("image")

        width = _snap(params["width"])
        height = _snap(params["height"])
        num_frames = _snap_frames(params["num_frames"])
        negative = params.get("negative_prompt")
        if negative is None:
            negative = NEGATIVE_DEFAULT

        seed = params.get("seed")
        generator = None
        if seed not in (None, ""):
            generator = torch.Generator(device="cuda").manual_seed(int(seed))

        self._set_flow_shift(params.get("flow_shift"))

        call = dict(
            prompt=prompt,
            negative_prompt=negative,
            width=width,
            height=height,
            num_frames=num_frames,
            num_inference_steps=int(params["num_inference_steps"]),
            guidance_scale=float(params["guidance_scale"]),
            guidance_scale_2=float(params["guidance_scale_2"]),
            max_sequence_length=int(params["max_sequence_length"]),
            generator=generator,
            output_type="np",
        )

        if self.is_i2v:
            if not image_raw:
                raise ValueError("Это image-to-video эндпоинт: нужно поле image (base64)")
            call["image"] = _decode_image(image_raw).resize((width, height), Image.LANCZOS)
            if params.get("last_image"):  # FLF2V: задаём и последний кадр — стык становится точным
                call["last_image"] = _decode_image(params["last_image"]).resize((width, height), Image.LANCZOS)
        elif not prompt.strip():
            raise ValueError("Пустой промпт")

        log.info("Генерация: %sx%s, %s кадров, %s шагов, shift=%s",
                 width, height, num_frames, call["num_inference_steps"], self._flow_shift)
        frames, used_mode, plan = self._run(call, width, height, num_frames)

        with tempfile.TemporaryDirectory() as tmp:
            out = os.path.join(tmp, "out.mp4")
            export_to_video(frames, out, fps=int(params["fps"]))
            video = Path(out).read_bytes()

        info = {
            "width": width, "height": height, "num_frames": num_frames,
            "fps": int(params["fps"]), "steps": call["num_inference_steps"],
            "guidance_scale": call["guidance_scale"], "flow_shift": self._flow_shift,
            "seed": seed, "seconds": round(time.time() - started, 1),
            "bytes": len(video),
            "mode": used_mode, "vram_gib": round(self.vram, 1), **plan,
        }
        log.info("Готово: %s", info)
        return {
            "video_base64": base64.b64encode(video).decode(),
            "content_type": "video/mp4",
            "info": info,
        }