small fixes

This commit is contained in:
kijai
2024-12-14 22:28:02 +02:00
parent 4c5507ffc6
commit ecae5f270e
3 changed files with 11 additions and 5 deletions
@@ -451,7 +451,12 @@ class HunyuanVideoPipeline(DiffusionPipeline):
if prompt_mask is not None:
prompt_mask = torch.cat([prompt_mask, prompt_mask])
if prompt_embeds_2 is not None:
prompt_embeds_2 = torch.cat([prompt_embeds_2, prompt_embeds_2])
prompt_embeds_2 = torch.cat([prompt_embeds_2, prompt_embeds_2])
prompt_embeds = prompt_embeds.to(device = device, dtype = self.base_dtype)
prompt_mask = prompt_mask.to(device)
if prompt_embeds_2 is not None:
prompt_embeds_2 = prompt_embeds_2.to(device = device, dtype = self.base_dtype)
# 4. Prepare timesteps
extra_set_timesteps_kwargs = self.prepare_extra_func_kwargs(
@@ -564,7 +569,6 @@ class HunyuanVideoPipeline(DiffusionPipeline):
else:
guidance_expand = None
#print("latent_model_input", latent_model_input.shape, "guidance_expand", guidance_expand)
# predict the noise residual
with torch.autocast(
device_type="cuda", dtype=self.base_dtype, enabled=True
+3 -3
View File
@@ -665,7 +665,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
out = {}
img = x
txt = text_states.to(x.device)
txt = text_states
_, _, ot, oh, ow = x.shape
tt, th, tw = (
ot // self.patch_size[0],
@@ -678,7 +678,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
# text modulation
if text_states_2 is not None:
vec = vec + self.vector_in(text_states_2.to(x.device))
vec = vec + self.vector_in(text_states_2)
# guidance modulation
if guidance is not None:
@@ -695,7 +695,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
if self.text_projection == "linear":
txt = self.txt_in(txt)
elif self.text_projection == "single_refiner":
txt = self.txt_in(txt, t, text_mask.to(x.device) if self.use_attention_mask else None)
txt = self.txt_in(txt, t, text_mask if self.use_attention_mask else None)
else:
raise NotImplementedError(
f"Unsupported text_projection: {self.text_projection}"
+2
View File
@@ -848,6 +848,8 @@ class HyVideoTextEncode:
tokens = clip_l.tokenize(negative_prompt, return_word_ids=True)
negative_prompt_embeds_2 = clip_l.encode_from_tokens(tokens, return_pooled=True, return_dict=False)[1]
negative_prompt_embeds_2 = negative_prompt_embeds_2.to(device=device)
else:
negative_prompt_embeds_2 = None
attention_mask_2, negative_attention_mask_2 = None, None
if force_offload: