| 6 | |
| 7 | |
| 8 | class 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 | """ |