mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Simple conditioning sorta merges but checking inputs now.
This commit is contained in:
@@ -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
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user