FloodDiffusion2-Live / space /text_encoder.py
caiyiyi1998's picture
Initial commit
9a25493
Raw History Blame Contribute Delete
7.5 kB
"""Original Wan UMT5-XXL encoder and CPU feature files for ZeroGPU.
Import spaces first. Call load_text_encoder at app-module scope, just like the
motion loader: CPU BF16 checkpoint -> virtual CUDA. Call encode only inside a
real GPU allocation. Save its CPU result to a session NPZ and read that file in
the generation worker; Python dictionaries modified in a fork are not shared.
The public Wan network, tokenizer cleaning, 512-token padding/truncation, and
T5EncoderModel.__call__ are unchanged. Meta construction plus strict assignment
avoids the original constructor's temporary 22.7 GB FP32 parameter allocation.
"""
import importlib
from pathlib import Path
import sys
import uuid
import numpy as np
import torch
def _cpu_feature(prompt, value):
if not isinstance(prompt, str):
raise TypeError("Prompt keys must be strings")
if isinstance(value, torch.Tensor):
if value.device.type != "cpu":
raise ValueError("Feature files require CPU tensors")
value = value.detach().float().numpy()
value = np.array(value, dtype=np.float32, copy=True)
if value.ndim != 2 or value.shape[1] != 4096 or not 0 < len(value) <= 512:
raise ValueError(f"Invalid UMT5 feature shape for {prompt!r}: {value.shape}")
if not np.isfinite(value).all():
raise ValueError(f"Non-finite UMT5 features for {prompt!r}")
return value
def save_prompt_features(path, bank):
"""Atomically save CPU features to a pickle-free NPZ shared by GPU workers."""
if not isinstance(bank, dict) or not bank:
raise ValueError("Expected a nonempty prompt-to-feature dictionary")
payload = {f"feature_{i}": _cpu_feature(prompt, value)
for i, (prompt, value) in enumerate(bank.items())}
payload["prompts"] = np.asarray(list(bank), dtype=np.str_)
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
temporary = path.with_name(f".{path.name}.{uuid.uuid4().hex}.tmp")
try:
with temporary.open("wb") as handle:
np.savez(handle, **payload)
temporary.replace(path)
finally:
temporary.unlink(missing_ok=True)
return str(path)
def load_prompt_features(path):
"""Read a fresh CPU bank in the generation fork; do not rely on globals."""
with np.load(path, allow_pickle=False) as archive:
names = archive["prompts"]
if names.ndim != 1 or names.dtype.kind not in ("U", "S"):
raise ValueError("Feature archive needs a one-dimensional string prompt list")
prompts = names.astype(str).tolist()
if not prompts or len(prompts) != len(set(prompts)):
raise ValueError("Feature archive prompt names must be nonempty and unique")
return {prompt: _cpu_feature(prompt, archive[f"feature_{i}"])
for i, prompt in enumerate(prompts)}
class TextEncoder:
"""Thin wrapper around the unchanged public T5EncoderModel inference call."""
def __init__(self, encoder, metadata):
self.encoder = encoder
self.metadata = metadata
@torch.inference_mode()
def encode(self, prompts):
"""Return {exact_prompt: CPU float32 ndarray(tokens, 4096)}.
Encode one unique prompt at a time to keep activation memory bounded.
The original tokenizer still pads each input to 512 tokens and trims
output according to its mask, including the EOS token.
"""
if isinstance(prompts, str):
prompts = [prompts]
prompts = list(prompts)
if not prompts or any(not isinstance(prompt, str) for prompt in prompts):
raise ValueError("Provide at least one string prompt")
device = next(self.encoder.model.parameters()).device
if device.type != "cuda":
raise RuntimeError("encode must run inside the ZeroGPU CUDA allocation")
bank = {}
for prompt in dict.fromkeys(prompts):
feature = self.encoder([prompt], device)[0]
bank[prompt] = _cpu_feature(prompt, feature.float().cpu())
return bank
def load_text_encoder(repo_root, encoder_path, tokenizer_path):
"""Load the exact public Wan UMT5-XXL weights without duplicate CPU copies.
encoder_path is the published models_t5_umt5-xxl-enc-bf16.pth; tokenizer_path
is the corresponding local google/umt5-xxl directory from the same deps.
No alternative Transformers T5 model or online tokenizer is loaded.
"""
root = Path(repo_root).resolve()
encoder_path = Path(encoder_path).resolve()
tokenizer_path = Path(tokenizer_path).resolve()
if not (root / "models/tools/t5.py").is_file():
raise FileNotFoundError(f"Public Wan T5 source is missing under {root}")
if not encoder_path.is_file():
raise FileNotFoundError(encoder_path)
if not tokenizer_path.is_dir():
raise FileNotFoundError(tokenizer_path)
if str(root) not in sys.path:
sys.path.insert(0, str(root))
source = importlib.import_module("models.tools.t5")
if Path(source.__file__).resolve() != root / "models/tools/t5.py":
raise RuntimeError("A different models package is already imported; use a clean process")
# The original constructor first creates all FP32 parameters before BF16
# conversion. Its model factory on meta allocates no parameter storage.
model = source.umt5_xxl(encoder_only=True, return_tokenizer=False,
dtype=torch.bfloat16, device="meta")
state = torch.load(encoder_path, map_location="cpu", weights_only=True, mmap=True)
if not isinstance(state, dict) or not state:
raise ValueError("Expected the original flat Wan UMT5 state dictionary")
if any(not isinstance(value, torch.Tensor) or value.dtype != torch.bfloat16
for value in state.values()):
raise ValueError("Expected the original all-BF16 Wan UMT5 checkpoint")
model.load_state_dict(state, strict=True, assign=True)
del state
if any(value.is_meta for value in list(model.parameters()) + list(model.buffers())):
raise RuntimeError("UMT5 checkpoint left uninitialized meta tensors")
model.eval().requires_grad_(False)
# Keep the exact public wrapper's __call__ and tokenizer rather than a
# rewritten forward path; only its allocation-heavy constructor is bypassed.
encoder = source.T5EncoderModel.__new__(source.T5EncoderModel)
encoder.text_len = 512
encoder.dtype = torch.bfloat16
encoder.device = torch.device("cuda")
encoder.checkpoint_path = str(encoder_path)
encoder.tokenizer_path = str(tokenizer_path)
encoder.t5_size = "xxl"
encoder.model = model
encoder.tokenizer = source.HuggingfaceTokenizer(
name=str(tokenizer_path), seq_len=512, clean="whitespace", local_files_only=True)
parameter_count = sum(parameter.numel() for parameter in model.parameters())
parameter_bytes = sum(parameter.numel() * parameter.element_size()
for parameter in model.parameters())
metadata = {"source": str(root), "encoder_path": str(encoder_path),
"tokenizer_path": str(tokenizer_path), "architecture": "public Wan UMT5-XXL encoder",
"dtype": "bfloat16", "text_len": 512, "strict_load": True,
"parameter_count": parameter_count, "parameter_bytes": parameter_bytes,
"load_method": "meta factory + CPU mmap weights + strict assign; module-scope CUDA"}
model.to("cuda")
return TextEncoder(encoder, metadata)