MCPcopy Create free account
hub / github.com/OpenImagingLab/FlashVSR / StepVideoPrompter

Class StepVideoPrompter

diffsynth/prompters/stepvideo_prompter.py:8–56  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

6
7
8class 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

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected