"""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)