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

Method __init__

model/diffusion.py:9–32  ·  view source on GitHub ↗

Initialize the Diffusion loss module.

(self, args, device)

Source from the content-addressed store, hash-verified

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)

Callers

nothing calls this directly

Calls 1

Tested by

no test coverage detected