File size: 7,501 Bytes
9a25493
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
"""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)