Simple conditioning sorta merges but checking inputs now.

This commit is contained in:
CrossProduct
2025-02-06 17:25:13 +00:00
parent d65958b181
commit 93e5069dc0
3 changed files with 9 additions and 82 deletions
@@ -48,58 +48,6 @@ class LTXVideoConditionedTransformer3DModel(LTXVideoTransformer3DModel):
attention_out_bias=attention_out_bias
)
# adapter.down_proj.weight
self.adapter = ConditionedResidualAdapterBottleneck(
input_dim=adapter_in_dim,
output_dim=128,
bottleneck_dim=64,
adapter_dropout=0.1,
adapter_init_scale=1e-3
)
@classmethod # rewrite this to ensure that if it doesn't have adapter weights it will load this way but if it has adapter weights it will not.
def from_pretrained(cls, pretrained_model_name_or_path,save_directory, **kwargs):
# First create an empty model with the desired architecture
print("Pretrain Loading")
model = cls()
model_dict = model.state_dict()
# # Then load the pretrained weights into it
pretrained = LTXVideoTransformer3DModel.from_pretrained(pretrained_model_name_or_path, **kwargs)
# Copy over the pretrained weights for the shared components
pretrained_dict = pretrained.state_dict()
# Filter out adapter weights from the pretrained dict
filtered_dict = {}
for k, v in pretrained_dict.items():
if k in model_dict:
filtered_dict[k] = v
else:
print(k)
# Update model with pretrained weights
model_dict.update(filtered_dict)
adapter_weights_path = os.path.join(save_directory, "adapter_weights.pth")
if os.path.exists(adapter_weights_path):
adapter_weights = torch.load(adapter_weights_path)
model.adapter.load_state_dict(adapter_weights)
print("Adapter weights loaded successfully.")
else:
print("No adapter weights file found. Using default adapter weights.")
model.load_state_dict(model_dict)
return model
def save_pretrained(self,save_directory):
model_state_dict = self.state_dict()
torch.save(model_state_dict, save_directory)
adapter_state_dict = self.adapter.state_dict()
save_directory = os.path.dirname(save_directory)
path = os.path.join(save_directory, "adapter_weights.pth")
torch.save(adapter_state_dict, path)
print("Saving Model with Adapter")
def forward(
self,
@@ -125,9 +73,6 @@ class LTXVideoConditionedTransformer3DModel(LTXVideoTransformer3DModel):
batch_size = hidden_states.size(0)
# 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)
@@ -213,31 +213,23 @@ class LTXConditionedPipeline(LTXPipeline):
# residual x latent
# add pose information to both channels
noisy_latents = noise_latent + pose_latent
pose_img_ref_latents = img_ref_latent + pose_latent
# expand channel information # B x 2C latent will be projected to adapter to scale it back to 128d using adapter.
condition_latent = torch.cat([pose_img_ref_latents,noisy_latents], dim=1)
noisy_latents = noise_latent + pose_latent + img_ref_latent
noise_latent_tokens = post_conditioned_latent_patchify(latents=noise_latent, num_frames=num_frames,
height=height,
width=width,
patch_size = 1,
patch_size_t = 1)["latents"]
# # expand channel information # B x 2C latent will be projected to adapter to scale it back to 128d using adapter.
# condition_latent = torch.cat([pose_img_ref_latents,noisy_latents], dim=1)
condition_tokens = post_conditioned_latent_patchify(latents=condition_latent,
condition_tokens = post_conditioned_latent_patchify(latents=noisy_latents,
num_frames=num_frames,
height=height,
width=width,
patch_size = 1,
patch_size_t = 1)["latents"]
# target video as a input residual
# noisy_residual_tokens = post_conditioned_latent_patchify(latents=noisy_latents,
# num_frames=num_frames,
# height=height,
# width=width,
# patch_size = 1,
# patch_size_t = 1)["latents"]
# Need to change the latents to
# 5. Prepare timesteps
+4 -14
View File
@@ -813,28 +813,18 @@ 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_img_ref_latents = img_refs_latents["latents"] + pose_video_latents["latents"]
pose_noisy_latents = noisy_latents + pose_video_latents["latents"] + img_refs_latents["latents"]
# expand channel information # B x 2C latent will be projected to adapter to scale it back to 128d using adapter.
condition_latents = torch.cat([pose_img_ref_latents,pose_noisy_latents], dim=1)
# project this turn into patches
# concat tensor
condition_tokens = post_conditioned_latent_patchify(latents=condition_latents,
condition_tokens = post_conditioned_latent_patchify(latents=pose_noisy_latents,
num_frames=latent_conditions["num_frames"],
height=latent_conditions["height"],
width=latent_conditions["width"],
patch_size = 1,
patch_size_t = 1)
# TODO REMOVE target video as a input residual might not need.
noisy_residual_tokens = post_conditioned_latent_patchify(latents=noisy_latents,
num_frames=latent_conditions["num_frames"],
height=latent_conditions["height"],
width=latent_conditions["width"],
patch_size = 1,
patch_size_t = 1)
# Target Latent Patchified concat tensor
# pose template noisey input [cat] img_ref + pose video
# That dict says latents but actually tokens.