| 6 | |
| 7 | |
| 8 | class TrainingModule(torch.nn.Module): |
| 9 | def __init__( |
| 10 | self, |
| 11 | vae_model_path = None, # |
| 12 | text_encoder_model_path = None, # |
| 13 | dit_model_path = None, # |
| 14 | tokenizer_path = None, # |
| 15 | |
| 16 | lora_base_model = None, # "vace" |
| 17 | lora_target_modules = "q,k,v,o,ffn.0,ffn.2", # "q,k,v,o,ffn.0,ffn.2" |
| 18 | lora_rank = 32, # 32 |
| 19 | |
| 20 | use_gradient_checkpointing = True, # True |
| 21 | use_gradient_checkpointing_offload = True, # True |
| 22 | extra_inputs = None, # "vace_video,vace_reference_image" |
| 23 | max_timestep_boundary = 1.0, # 1.0 |
| 24 | min_timestep_boundary = 0.0, # 0.0 |
| 25 | ): |
| 26 | super().__init__() |
| 27 | |
| 28 | self.pipe = EeveePipeline.from_pretrained( |
| 29 | torch_dtype = torch.bfloat16, |
| 30 | device = "cpu", |
| 31 | vae_model_path = vae_model_path, |
| 32 | text_encoder_model_path = text_encoder_model_path, |
| 33 | dit_model_path = dit_model_path, |
| 34 | tokenizer_path = tokenizer_path |
| 35 | ) |
| 36 | self.switch_pipe_to_training_mode( |
| 37 | self.pipe, |
| 38 | lora_base_model, |
| 39 | lora_target_modules, |
| 40 | lora_rank |
| 41 | ) |
| 42 | |
| 43 | self.use_gradient_checkpointing = use_gradient_checkpointing |
| 44 | self.use_gradient_checkpointing_offload = use_gradient_checkpointing_offload |
| 45 | self.extra_inputs = extra_inputs.split(",") if extra_inputs is not None else [] |
| 46 | self.max_timestep_boundary = max_timestep_boundary |
| 47 | self.min_timestep_boundary = min_timestep_boundary |
| 48 | |
| 49 | |
| 50 | def forward(self, data): |
| 51 | inputs = self.forward_preprocess(data) |
| 52 | models = {name: getattr(self.pipe, name) for name in self.pipe.in_iteration_models} # todo |
| 53 | loss = self.pipe.training_loss(**models, **inputs) |
| 54 | return loss |
| 55 | |
| 56 | |
| 57 | def switch_pipe_to_training_mode( |
| 58 | self, |
| 59 | pipe, |
| 60 | lora_base_model, # "vace" |
| 61 | lora_target_modules, # "q,k,v,o,ffn.0,ffn.2" |
| 62 | lora_rank, # 32 |
| 63 | ): |
| 64 | # Scheduler |
| 65 | pipe.scheduler.set_timesteps(1000, training=True) |