mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Added residual
This commit is contained in:
@@ -220,10 +220,11 @@ def conditioned_forward_pass(
|
||||
num_frames: int,
|
||||
height: int,
|
||||
width: int,
|
||||
noisy_latents_residual:torch.Tensor,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
rope_interpolation_scale = [1 / 25, 32, 32]
|
||||
|
||||
|
||||
denoised_latents = transformer(
|
||||
hidden_states=noisy_latents,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
@@ -234,6 +235,7 @@ def conditioned_forward_pass(
|
||||
width=width,
|
||||
rope_interpolation_scale=rope_interpolation_scale,
|
||||
return_dict=False,
|
||||
residual_x=noisy_latents_residual,
|
||||
)[0]
|
||||
|
||||
return {"latents": denoised_latents}
|
||||
|
||||
@@ -827,6 +827,7 @@ class Trainer:
|
||||
patch_size_t = 1)
|
||||
# pose information latent
|
||||
latent_conditions.update({"noisy_latents": latent_final["latents"]})
|
||||
latent_conditions.update({"noisy_latents_residual":noisy_latents_residual["latents"]})
|
||||
else:
|
||||
# Default to flow-matching noise addition
|
||||
noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise
|
||||
@@ -848,7 +849,6 @@ class Trainer:
|
||||
transformer=self.transformer,
|
||||
scheduler=self.scheduler,
|
||||
timesteps=timesteps,
|
||||
residual_x=noisy_latents_residual,
|
||||
**latent_conditions,
|
||||
**text_conditions,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user