MCPcopy Create free account
hub / github.com/JaydenLyh/Reward-Forcing / WanTextEncoder

Class WanTextEncoder

utils/wan_wrapper.py:14–51  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

12
13
14class WanTextEncoder(torch.nn.Module):
15 def __init__(self, checkpoint_path) -> None:
16 super().__init__()
17
18 self.text_encoder = umt5_xxl(
19 encoder_only=True,
20 return_tokenizer=False,
21 dtype=torch.float32,
22 device=torch.device('cpu')
23 ).eval().requires_grad_(False)
24
25 self.text_encoder.load_state_dict(
26 torch.load(f"{checkpoint_path}/Wan2.1-T2V-1.3B/models_t5_umt5-xxl-enc-bf16.pth",
27 map_location='cpu', weights_only=False)
28 )
29
30 self.tokenizer = HuggingfaceTokenizer(
31 name=f"{checkpoint_path}/Wan2.1-T2V-1.3B/google/umt5-xxl/", seq_len=512, clean='whitespace')
32
33 @property
34 def device(self):
35 # Assume we are always on GPU
36 return torch.cuda.current_device()
37
38 def forward(self, text_prompts: List[str]) -> dict:
39 ids, mask = self.tokenizer(
40 text_prompts, return_mask=True, add_special_tokens=True)
41 ids = ids.to(self.device)
42 mask = mask.to(self.device)
43 seq_lens = mask.gt(0).sum(dim=1).long()
44 context = self.text_encoder(ids, mask)
45
46 for u, v in zip(context, seq_lens):
47 u[v:] = 0.0 # set padding to 0.0
48
49 return {
50 "prompt_embeds": context
51 }
52
53
54class WanVAEWrapper(torch.nn.Module):

Callers 9

_initialize_modelsMethod · 0.90
_initialize_modelsMethod · 0.90
_initialize_modelsMethod · 0.90
_initialize_modelsMethod · 0.90
init_modelFunction · 0.90
__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected