""" 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, "loras": [{"file": "nsfw/NSFW-22-H-e8.safetensors", "strength": 0.8}, {"file": "nsfw/NSFW-22-L-e8.safetensors", "strength": 0.8}] }, "image": "" # только для i2v } Выход: {"video_base64": "...", "content_type": "video/mp4", "info": {...}} """ import base64 import io import json import logging import os import re import struct 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压缩残留, 丑陋的, 残缺的, 多余的手指, 画得不好的手部, " "画得不好的脸部, 畸形的, 毁容的, 形态畸形的肢体, 手指融合, 静止不动的画面, " "杂乱的背景, 三条腿, 背景人很多, 倒着走" ) # --------------------------------------------------------------------------- LoRA # # Файлы лежат в отдельном ПРИВАТНОМ датасете, а не в репозитории эндпоинта, и это # осознанно: их там 24 ГиБ, и переезд в репозиторий модели раздул бы холодный старт # с 68.8 ГБ до девяноста с лишним — платится он при каждом подъёме спящей реплики, # а нужны за раз две-три LoRA. Взамен эндпоинту нужен HF_TOKEN в переменных окружения: # без него приватный датасет не читается. LORA_REPO = os.environ.get("WAN_LORA_REPO", "digleone/wan2.2-Loras") LORA_REPO_TYPE = os.environ.get("WAN_LORA_REPO_TYPE", "dataset") # Имя секрета на эндпоинте задаётся руками в UI, и промахнуться в нём легко: у нас # он приехал как HF, а не HF_TOKEN. Перебираем ходовые вместо того, чтобы требовать # одно конкретное — цена ошибки тут пересборка и холодный старт на 68.8 ГБ, а не # сообщение в логе. НЕ передавать token=True как запасной вариант: huggingface_hub # ответит «Token is required, but no token found», и в 400-ке это выглядит так, # будто виноват запрос, а не незаполненная настройка. LORA_TOKEN_VARS = ("WAN_LORA_TOKEN", "HF_TOKEN", "HF", "HUGGING_FACE_HUB_TOKEN", "HUGGINGFACE_TOKEN") def _lora_token(): for var in LORA_TOKEN_VARS: val = (os.environ.get(var) or "").strip() if val: return val raise RuntimeError( f"LoRA лежат в приватном датасете {LORA_REPO}, а токена в окружении эндпоинта нет. " f"Добавь секрет с любым из имён: {', '.join(LORA_TOKEN_VARS)} " f"(страница эндпоинта -> Settings -> Environment variables, тип secret).") # У Wan2.2 A14B ровно 40 блоков трансформера. Проверка ловит то, что по имени файла # не видно: в датасете лежат LoRA под TI2V-5B (30 блоков) и одна на 4 блока. Они # сконвертируются без ошибки, а упадут уже в peft при подгонке ключей — с трейсбеком, # по которому причина не читается. LORA_BLOCKS_EXPECTED = 40 # Lightning в веса эндпоинта УЖЕ вмержен (см. DEPLOY_HF.md). Если подгрузить его же # сверху, дистилляция применится дважды — 8 шагов при cfg=1 превращаются в кашу, и # выглядит это как «LoRA сломала модель», хотя сломал её повтор ускорялки. LORA_DENY = ("lightning", "lightx2v", "cfg_step_distill", "4steps") def _lora_expert(file_name): """ high или low — только по имени файла: внутри тензоров этого нет. Пары обучаются порознь под двух экспертов MoE, и high обязан ехать в transformer, а low — в transformer_2. Перепутать местами хуже, чем не грузить вовсе: LoRA отработает на той половине шагов, для которой её не учили. """ low = file_name.lower() if re.search(r"(^|[^a-z])(high|high[_-]?noise)([^a-z]|$)|[-_]h[-_]?e?\d|-h-", low): return "high" if re.search(r"(^|[^a-z])(low|low[_-]?noise)([^a-z]|$)|[-_]l[-_]?e?\d|-l-", low): return "low" return None def _lora_blocks(path): """Число блоков из заголовка safetensors — читаем 8 байт длины и сам заголовок.""" with open(path, "rb") as fh: n = struct.unpack(" файл. При старте пусто и НИЧЕГО не грузится. # Ленивость тут не оптимизация, а страховка: битая LoRA обязана уронить один # запрос, а не эндпоинт. Упади она в __init__ — вебсервис не поднимется вовсе, # и снаружи это выглядит как «эндпоинт сломался», без единого намёка на причину. self._adapters = {} self._active = None log.info("Готов. i2v=%s, flow_shift по умолчанию=%s, LoRA из %s", self.is_i2v, self._flow_shift, LORA_REPO) # ----------------------------------------------------------------- LoRA def _lora_fetch(self, file_name): """Скачивает LoRA в кеш контейнера и проверяет, что она вообще от этой модели.""" from huggingface_hub import hf_hub_download low = file_name.lower() hit = next((d for d in LORA_DENY if d in low), None) if hit: raise ValueError( f"{file_name}: похоже на ускоряющую LoRA ({hit}), а Lightning в веса этого " f"эндпоинта уже вмержен — вторая копия применится поверх первой. " f"Если это ошибка распознавания, переименуй файл.") path = hf_hub_download( repo_id=LORA_REPO, repo_type=LORA_REPO_TYPE, filename=file_name, token=_lora_token()) blocks = _lora_blocks(path) if blocks and blocks != LORA_BLOCKS_EXPECTED: raise ValueError( f"{file_name}: {blocks} блоков, а у Wan2.2 A14B их {LORA_BLOCKS_EXPECTED}. " f"Это LoRA от другой модели (30 блоков — TI2V-5B), она сюда не встанет.") return path def _lora_load(self, file_name, expert): """Один файл -> один адаптер в нужном эксперте. Повторно не грузим.""" name = re.sub(r"[^0-9a-zA-Z]+", "_", f"{file_name}_{expert}").strip("_")[:60] if name in self._adapters: return name path = self._lora_fetch(file_name) log.info("Гружу LoRA %s -> %s (%s)", file_name, expert, name) # load_into_transformer_2 — единственный способ адресовать второго эксперта: # по умолчанию diffusers кладёт всё в transformer. T2V-LoRA на I2V-весах # доедет: _maybe_expand_t2v_lora_for_i2v дополнит недостающие проекции # картиночного внимания нулями, то есть на image-путь она просто не влияет. self.pipe.load_lora_weights(path, adapter_name=name, load_into_transformer_2=(expert == "low")) self._adapters[name] = file_name return name def _lora_apply(self, specs): """ Приводит пайплайн к запрошенному набору LoRA. Пусто — выключаем все. Силы задаются НА ЗАПРОС, а не при загрузке, и это главное решение здесь. Иначе перебор 0.6 / 0.8 / 1.0 означал бы правку handler'а, новую ревизию репозитория, пересборку и холодный старт на 68.8 ГБ — цикл, в котором подобрать силу невозможно в принципе. """ if not specs: if self._active: self.pipe.disable_lora() self._active = None return [] names, weights, used = [], [], [] for s in specs: file_name = (s.get("file") or s.get("name") or "").strip() if not file_name: raise ValueError("В loras нужен ключ file — имя файла в датасете") expert = (s.get("expert") or "").lower() or _lora_expert(file_name) if expert not in ("high", "low"): raise ValueError( f"{file_name}: не понял, к какому эксперту его цеплять. У Wan2.2 A14B " f"их два, и LoRA обучается под одного. Добавь в запрос " f'"expert": "high" или "low".') strength = float(s.get("strength", 1.0)) names.append(self._lora_load(file_name, expert)) weights.append(strength) used.append({"file": file_name, "expert": expert, "strength": strength}) self.pipe.set_adapters(names, weights) self.pipe.enable_lora() self._active = tuple(zip(names, weights)) return used 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")) # LoRA применяем ДО выбора режима весов: set_adapters трогает оба # трансформера, и делать это посреди переезда CPU↔GPU незачем. loras_used = self._lora_apply(params.get("loras") or []) 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), # В паспорт ролика: без имён и сил повторить кадр через полгода нечем, # ровно та же причина, по которой они пишутся в сайдкар у картинок. "loras": loras_used, **plan, } log.info("Готово: %s", info) return { "video_base64": base64.b64encode(video).decode(), "content_type": "video/mp4", "info": info, }