MCPcopy Create free account
hub / github.com/JaydenLyh/Reward-Forcing / CausalDiffusion

Class CausalDiffusion

model/diffusion.py:8–125  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

6
7
8class CausalDiffusion(BaseModel):
9 def __init__(self, args, device):
10 """
11 Initialize the Diffusion loss module.
12 """
13 super().__init__(args, device)
14 self.num_frame_per_block = getattr(args, "num_frame_per_block", 1)
15 if self.num_frame_per_block > 1:
16 self.generator.model.num_frame_per_block = self.num_frame_per_block
17 self.independent_first_frame = getattr(args, "independent_first_frame", False)
18 if self.independent_first_frame:
19 self.generator.model.independent_first_frame = True
20
21 if args.gradient_checkpointing:
22 self.generator.enable_gradient_checkpointing()
23
24 # Step 2: Initialize all hyperparameters
25 self.num_train_timestep = args.num_train_timestep
26 self.min_step = int(0.02 * self.num_train_timestep)
27 self.max_step = int(0.98 * self.num_train_timestep)
28 self.guidance_scale = args.guidance_scale
29 self.timestep_shift = getattr(args, "timestep_shift", 1.0)
30 self.teacher_forcing = getattr(args, "teacher_forcing", False)
31 # Noise augmentation in teacher forcing, we add small noise to clean context latents
32 self.noise_augmentation_max_timestep = getattr(args, "noise_augmentation_max_timestep", 0)
33
34 def _initialize_models(self, args):
35 self.generator = WanDiffusionWrapper(**getattr(args, "model_kwargs", {}), is_causal=True)
36 self.generator.model.requires_grad_(True)
37
38 self.text_encoder = WanTextEncoder()
39 self.text_encoder.requires_grad_(False)
40
41 self.vae = WanVAEWrapper()
42 self.vae.requires_grad_(False)
43
44 def generator_loss(
45 self,
46 image_or_video_shape,
47 conditional_dict: dict,
48 unconditional_dict: dict,
49 clean_latent: torch.Tensor,
50 initial_latent: torch.Tensor = None
51 ) -> Tuple[torch.Tensor, dict]:
52 """
53 Generate image/videos from noise and compute the DMD loss.
54 The noisy input to the generator is backward simulated.
55 This removes the need of any datasets during distillation.
56 See Sec 4.5 of the DMD2 paper (https://arxiv.org/abs/2405.14867) for details.
57 Input:
58 - image_or_video_shape: a list containing the shape of the image or video [B, F, C, H, W].
59 - conditional_dict: a dictionary containing the conditional information (e.g. text embeddings, image embeddings).
60 - unconditional_dict: a dictionary containing the unconditional information (e.g. null/negative text embeddings, null/negative image embeddings).
61 - clean_latent: a tensor containing the clean latents [B, F, C, H, W]. Need to be passed when no backward simulation is used.
62 Output:
63 - loss: a scalar tensor representing the generator loss.
64 - generator_log_dict: a dictionary containing the intermediate tensors for logging.
65 """

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected