mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Conditioned Adapter make residual optional.
This commit is contained in:
@@ -219,7 +219,7 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
# 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_latent_tokens = post_conditioned_latent_patchify(latents=noisy_latents, num_frames=num_frames,
|
||||
noise_latent_tokens = post_conditioned_latent_patchify(latents=noise_latent, num_frames=num_frames,
|
||||
height=height,
|
||||
width=width,
|
||||
patch_size = 1,
|
||||
@@ -232,12 +232,12 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
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"]
|
||||
# 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
|
||||
@@ -280,11 +280,9 @@ 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)
|
||||
|
||||
|
||||
latent_model_input = torch.cat([noisy_latent_tokens] * 2) if self.do_classifier_free_guidance else noisy_latent_tokens
|
||||
latent_model_input = torch.cat([condition_tokens] * 2) if self.do_classifier_free_guidance else condition_tokens
|
||||
latent_model_input = latent_model_input.to(prompt_embeds.dtype)
|
||||
condition_tokens = condition_tokens.to(prompt_embeds.dtype)
|
||||
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latent_model_input.shape[0])
|
||||
|
||||
@@ -300,7 +298,7 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
rope_interpolation_scale=rope_interpolation_scale,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
residual_x=latent_model_input,
|
||||
residual_x=None,
|
||||
)[0]
|
||||
# patchified
|
||||
noise_pred = noise_pred.float()
|
||||
@@ -311,7 +309,7 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
# Noisy latents needs to be
|
||||
noisy_residual_tokens = self.scheduler.step(noise_pred, t, noisy_residual_tokens, return_dict=False)[0]
|
||||
noise_latent_tokens = self.scheduler.step(noise_pred, t, noise_latent_tokens, return_dict=False)[0]
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
@@ -319,7 +317,7 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
noisy_residual_tokens = callback_outputs.pop("latents", noisy_residual_tokens)
|
||||
noise_latent_tokens = callback_outputs.pop("latents", noise_latent_tokens)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
|
||||
# call the callback, if provided
|
||||
@@ -330,25 +328,25 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
xm.mark_step()
|
||||
|
||||
if output_type == "latent":
|
||||
video = noisy_residual_tokens
|
||||
video = noise_latent_tokens
|
||||
else:
|
||||
noisy_residual_tokens = self._unpack_latents(
|
||||
noisy_residual_tokens,
|
||||
noise_latent_tokens = self._unpack_latents(
|
||||
noise_latent_tokens,
|
||||
latent_num_frames,
|
||||
latent_height,
|
||||
latent_width,
|
||||
self.transformer_spatial_patch_size,
|
||||
self.transformer_temporal_patch_size,
|
||||
)
|
||||
noisy_residual_tokens = self._denormalize_latents(
|
||||
noisy_residual_tokens, self.vae.latents_mean, self.vae.latents_std, self.vae.config.scaling_factor
|
||||
noise_latent_tokens = self._denormalize_latents(
|
||||
noise_latent_tokens, self.vae.latents_mean, self.vae.latents_std, self.vae.config.scaling_factor
|
||||
)
|
||||
noisy_residual_tokens = noisy_residual_tokens.to(prompt_embeds.dtype)
|
||||
noise_latent_tokens = noise_latent_tokens.to(prompt_embeds.dtype)
|
||||
|
||||
if not self.vae.config.timestep_conditioning:
|
||||
timestep = None
|
||||
else:
|
||||
noise = torch.randn(noisy_residual_tokens.shape, generator=generator, device=device, dtype=noisy_residual_tokens.dtype)
|
||||
noise = torch.randn(noise_latent_tokens.shape, generator=generator, device=device, dtype=noise_latent_tokens.dtype)
|
||||
if not isinstance(decode_timestep, list):
|
||||
decode_timestep = [decode_timestep] * batch_size
|
||||
if decode_noise_scale is None:
|
||||
@@ -356,13 +354,13 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
elif not isinstance(decode_noise_scale, list):
|
||||
decode_noise_scale = [decode_noise_scale] * batch_size
|
||||
|
||||
timestep = torch.tensor(decode_timestep, device=device, dtype=noisy_residual_tokens.dtype)
|
||||
decode_noise_scale = torch.tensor(decode_noise_scale, device=device, dtype=noisy_residual_tokens.dtype)[
|
||||
timestep = torch.tensor(decode_timestep, device=device, dtype=noise_latent_tokens.dtype)
|
||||
decode_noise_scale = torch.tensor(decode_noise_scale, device=device, dtype=noise_latent_tokens.dtype)[
|
||||
:, None, None, None, None
|
||||
]
|
||||
noisy_residual_tokens = (1 - decode_noise_scale) * noisy_residual_tokens + decode_noise_scale * noise
|
||||
noise_latent_tokens = (1 - decode_noise_scale) * noise_latent_tokens + decode_noise_scale * noise
|
||||
|
||||
video = self.vae.decode(noisy_residual_tokens, timestep, return_dict=False)[0]
|
||||
video = self.vae.decode(noise_latent_tokens, timestep, return_dict=False)[0]
|
||||
video = self.video_processor.postprocess_video(video, output_type=output_type)
|
||||
|
||||
# Offload all models
|
||||
|
||||
@@ -52,7 +52,10 @@ class ConditionedResidualAdapterBottleneck(nn.Module):
|
||||
# Up projection to new dimension
|
||||
hidden_states = self.up_proj(hidden_states)
|
||||
|
||||
hidden_states = hidden_states + residual_x
|
||||
if residual_x is None:
|
||||
return hidden_states
|
||||
else:
|
||||
hidden_states = hidden_states + residual_x
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
Reference in New Issue
Block a user