MCPcopy Create free account
hub / github.com/AMAP-ML/Eevee / TrainingModule

Class TrainingModule

models/training.py:8–218  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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

Callers 1

train.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected