| 12 | |
| 13 | |
| 14 | class 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 | |
| 54 | class WanVAEWrapper(torch.nn.Module): |
no outgoing calls
no test coverage detected