Instructions to use digleone/wan22-t2v-a14b-lightning-endpoint with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use digleone/wan22-t2v-a14b-lightning-endpoint with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("digleone/wan22-t2v-a14b-lightning-endpoint", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
| """ | |
| 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, | |
| } | |