| 6 | |
| 7 | |
| 8 | class StepVideoPrompter(BasePrompter): |
| 9 | |
| 10 | def __init__( |
| 11 | self, |
| 12 | tokenizer_1_path=None, |
| 13 | ): |
| 14 | if tokenizer_1_path is None: |
| 15 | base_path = os.path.dirname(os.path.dirname(__file__)) |
| 16 | tokenizer_1_path = os.path.join( |
| 17 | base_path, "tokenizer_configs/hunyuan_dit/tokenizer") |
| 18 | super().__init__() |
| 19 | self.tokenizer_1 = BertTokenizer.from_pretrained(tokenizer_1_path) |
| 20 | |
| 21 | def fetch_models(self, text_encoder_1: HunyuanDiTCLIPTextEncoder = None, text_encoder_2: STEP1TextEncoder = None): |
| 22 | self.text_encoder_1 = text_encoder_1 |
| 23 | self.text_encoder_2 = text_encoder_2 |
| 24 | |
| 25 | def encode_prompt_using_clip(self, prompt, max_length, device): |
| 26 | text_inputs = self.tokenizer_1( |
| 27 | prompt, |
| 28 | padding="max_length", |
| 29 | max_length=max_length, |
| 30 | truncation=True, |
| 31 | return_attention_mask=True, |
| 32 | return_tensors="pt", |
| 33 | ) |
| 34 | prompt_embeds = self.text_encoder_1( |
| 35 | text_inputs.input_ids.to(device), |
| 36 | attention_mask=text_inputs.attention_mask.to(device), |
| 37 | ) |
| 38 | return prompt_embeds |
| 39 | |
| 40 | def encode_prompt_using_llm(self, prompt, max_length, device): |
| 41 | y, y_mask = self.text_encoder_2(prompt, max_length=max_length, device=device) |
| 42 | return y, y_mask |
| 43 | |
| 44 | def encode_prompt(self, |
| 45 | prompt, |
| 46 | positive=True, |
| 47 | device="cuda"): |
| 48 | |
| 49 | prompt = self.process_prompt(prompt, positive=positive) |
| 50 | |
| 51 | clip_embeds = self.encode_prompt_using_clip(prompt, max_length=77, device=device) |
| 52 | llm_embeds, llm_mask = self.encode_prompt_using_llm(prompt, max_length=320, device=device) |
| 53 | |
| 54 | llm_mask = torch.nn.functional.pad(llm_mask, (clip_embeds.shape[1], 0), value=1) |
| 55 | |
| 56 | return clip_embeds, llm_embeds, llm_mask |