Skip to content

Commit d9890ab

Browse files
committed
sampled video looks correct
1 parent 0fb70ad commit d9890ab

2 files changed

Lines changed: 20 additions & 21 deletions

File tree

fastvideo/train/methods/rl/embeddings.py

Lines changed: 6 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -34,13 +34,12 @@ def compute_text_embeddings(
3434
return_tensors="pt",
3535
)
3636
text_input_ids = text_inputs.input_ids.to(device)
37-
37+
attention_mask = text_inputs.attention_mask.to(device)
3838
with torch.no_grad():
3939
prompt_embeds = text_encoder(
40-
text_input_ids
41-
)[0]
42-
43-
prompt_embeds = prompt_embeds.to(
44-
dtype=text_encoder.dtype, device=device
45-
)
40+
text_input_ids, attention_mask=attention_mask
41+
).last_hidden_state
42+
43+
# make padding token 0
44+
prompt_embeds[attention_mask == 0] = 0
4645
return prompt_embeds

fastvideo/train/models/wan/wan_genrl.py

Lines changed: 14 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -104,17 +104,20 @@ def init_preprocessors(
104104

105105
# Load text encoder and tokenizer.
106106
model_path = str(training_config.model_path)
107-
self._load_text_encoder(model_path)
107+
self._load_text_encoder(
108+
model_path, training_config
109+
)
108110

109111
# Dummy dataloader for the trainer's outer loop.
110112
self.dataloader = _InfiniteDummyLoader()
111113
self.start_step = 0
112114

113-
def _load_text_encoder(self, model_path: str) -> None:
114-
from transformers import (
115-
AutoTokenizer,
116-
UMT5EncoderModel,
117-
)
115+
def _load_text_encoder(
116+
self,
117+
model_path: str,
118+
training_config: TrainingConfig,
119+
) -> None:
120+
from transformers import AutoTokenizer
118121

119122
logger.info(
120123
"Loading tokenizer from %s", model_path
@@ -124,17 +127,14 @@ def _load_text_encoder(self, model_path: str) -> None:
124127
)
125128

126129
logger.info(
127-
"Loading T5 text encoder from %s", model_path
130+
"Loading text encoder from %s", model_path
128131
)
129-
dtype = self._get_training_dtype()
130-
self.text_encoder = UMT5EncoderModel.from_pretrained(
131-
model_path,
132-
subfolder="text_encoder",
133-
torch_dtype=dtype,
132+
self.text_encoder = load_module_from_path(
133+
model_path=model_path,
134+
module_type="text_encoder",
135+
training_config=training_config,
134136
)
135-
self.text_encoder.to(self.device)
136137
self.text_encoder.requires_grad_(False)
137-
self.text_encoder.eval()
138138

139139
def on_train_start(self) -> None:
140140
"""Skip negative conditioning (handled by method)."""

0 commit comments

Comments
 (0)