The encoder hidden states and encode attention mask sizes do not match up in the validation pipeline.

This commit is contained in:
CrossProduct
2025-01-29 05:24:06 +00:00
parent 4ab3b25c13
commit 707e001f7b
3 changed files with 13 additions and 11 deletions
@@ -106,7 +106,7 @@ class LTXVideoConditionedTransformer3DModel(LTXVideoTransformer3DModel):
# inject the condition and the residual then project it into the pretrained proj_in
hidden_states = self.adapter(residual_x=residual_x,
conditioned_x=hidden_states)
# whats this value ?
hidden_states = self.proj_in(hidden_states)
temb, embedded_timestep = self.time_embed(
@@ -134,7 +134,7 @@ class LTXVideoConditionedTransformer3DModel(LTXVideoTransformer3DModel):
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
hidden_states = torch.utils.checkpoint.checkpoint(
hidden_states = torch.utils.checkpoint.checkpoint( # crashes here
create_custom_forward(block),
hidden_states,
encoder_hidden_states,
@@ -180,7 +180,7 @@ class LTXConditionedPipeline(LTXPipeline):
height,
width,
num_frames,
torch.float32,
torch.bfloat16,
device,
generator,
)
@@ -192,7 +192,7 @@ class LTXConditionedPipeline(LTXPipeline):
patch_size=self.transformer.config.patch_size,
patch_size_t=self.transformer.config.patch_size_t,
device=device,
dtype=torch.float32,
dtype=torch.bfloat16,
generator=generator,
)["latents"]
@@ -202,7 +202,7 @@ class LTXConditionedPipeline(LTXPipeline):
patch_size=self.transformer.config.patch_size,
patch_size_t=self.transformer.config.patch_size_t,
device=device,
dtype=torch.float32,
dtype=torch.bfloat16,
generator=generator,
)["latents"]
@@ -280,6 +280,8 @@ class LTXConditionedPipeline(LTXPipeline):
# change this
# latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents
# latent_model_input = latent_model_input.to(prompt_embeds.dtype)
# might be the bug ..
latent_model_input = noisy_latent_tokens
latent_model_input = latent_model_input.to(prompt_embeds.dtype)
+6 -6
View File
@@ -813,7 +813,7 @@ class Trainer:
noisy_latents = noisy_latents.to(target_video_latents.dtype)
# add pose information to both channels
pose_noisy_latents = noisy_latents + pose_video_latents["latents"]
pose_noisy_latents = noisy_latents # + pose_video_latents["latents"]
pose_img_ref_latents = img_refs_latents["latents"] + pose_video_latents["latents"]
# expand channel information # B x 2C latent will be projected to adapter to scale it back to 128d using adapter.
@@ -821,14 +821,14 @@ class Trainer:
# project this turn into patches
# concat tensor
pose_img_ref_latents = post_conditioned_latent_patchify(latents=pose_img_ref_latents,
pose_img_ref_tokens = post_conditioned_latent_patchify(latents=pose_img_ref_latents,
num_frames=latent_conditions["num_frames"],
height=latent_conditions["height"],
width=latent_conditions["width"],
patch_size = 1,
patch_size_t = 1)
# target video as a input residual
noisy_latents_residual = post_conditioned_latent_patchify(latents=noisy_latents,
noisy_residual_tokens = post_conditioned_latent_patchify(latents=noisy_latents,
num_frames=latent_conditions["num_frames"],
height=latent_conditions["height"],
width=latent_conditions["width"],
@@ -836,9 +836,9 @@ class Trainer:
patch_size_t = 1)
# Target Latent Patchified concat tensor
# pose template noisey input [cat] img_ref + pose video
latent_conditions.update({"noisy_latents": pose_img_ref_latents["latents"]})
latent_conditions.update({"noisy_latents": pose_img_ref_tokens["latents"]})
# input video noise at level residual information to adapter
latent_conditions.update({"noisy_latents_residual":noisy_latents_residual["latents"]})
latent_conditions.update({"noisy_latents_residual":noisy_residual_tokens["latents"]})
else:
# Default to flow-matching noise addition
noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise
@@ -1029,7 +1029,7 @@ class Trainer:
image = load_image(image)
if video is not None:
video = load_video(video)
# loading videos for inference ..
if self.pose_condition:
if pose_video is not None:
pose_video = self.preprocess_condition_video(pose_video,max_num_frames=self.dataset.max_num_frames)