Initialize the Diffusion loss module.
(self, args, device)
| 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) |
nothing calls this directly
no test coverage detected