Spaces:
Running on Zero
Running on Zero
Download space/text_encoder.py from AlayaLab/FloodDiffusion2-Live: direct link, hf CLI and curl.
- Browser
- Download file 7.5 kB
-
https://huggingface.co/spaces/AlayaLab/FloodDiffusion2-Live/resolve/main/space/text_encoder.py
- Command line
-
hf download hf://spaces/AlayaLab/FloodDiffusion2-Live/space/text_encoder.py
-
curl -L -o text_encoder.py https://huggingface.co/spaces/AlayaLab/FloodDiffusion2-Live/resolve/main/space/text_encoder.py
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 | |
| 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) | |