(
self,
model_paths=None, model_id_with_origin_paths=None,
tokenizer_path=None,
trainable_models=None,
lora_base_model=None, lora_target_modules="", lora_rank=32, lora_checkpoint=None,
preset_lora_path=None, preset_lora_model=None,
use_gradient_checkpointing=True,
use_gradient_checkpointing_offload=False,
extra_inputs=None,
enable_text=True,
modules=("dit", "text", "vae", "image", "action"),
fp8_models=None,
offload_models=None,
ckpt_path=None,
device="cpu",
task="sft",
max_timestep_boundary=1.0,
min_timestep_boundary=0.0,
num_history_frames=1,
args=None,
)
| 12 | |
| 13 | class WanTrainingModule(DiffusionTrainingModule): |
| 14 | def __init__( |
| 15 | self, |
| 16 | model_paths=None, model_id_with_origin_paths=None, |
| 17 | tokenizer_path=None, |
| 18 | trainable_models=None, |
| 19 | lora_base_model=None, lora_target_modules="", lora_rank=32, lora_checkpoint=None, |
| 20 | preset_lora_path=None, preset_lora_model=None, |
| 21 | use_gradient_checkpointing=True, |
| 22 | use_gradient_checkpointing_offload=False, |
| 23 | extra_inputs=None, |
| 24 | enable_text=True, |
| 25 | modules=("dit", "text", "vae", "image", "action"), |
| 26 | fp8_models=None, |
| 27 | offload_models=None, |
| 28 | ckpt_path=None, |
| 29 | device="cpu", |
| 30 | task="sft", |
| 31 | max_timestep_boundary=1.0, |
| 32 | min_timestep_boundary=0.0, |
| 33 | num_history_frames=1, |
| 34 | args=None, |
| 35 | ): |
| 36 | super().__init__() |
| 37 | # Warning |
| 38 | if not use_gradient_checkpointing: |
| 39 | warnings.warn("Gradient checkpointing is detected as disabled. To prevent out-of-memory errors, the training framework will forcibly enable gradient checkpointing.") |
| 40 | use_gradient_checkpointing = True |
| 41 | |
| 42 | # Load models |
| 43 | model_configs = self.parse_model_configs(model_paths, model_id_with_origin_paths, fp8_models=fp8_models, offload_models=offload_models, device=device) |
| 44 | tokenizer_config = ModelConfig(model_id="Wan-AI/Wan2.1-T2V-1.3B", origin_file_pattern="google/umt5-xxl/") if enable_text and tokenizer_path is None else (ModelConfig(tokenizer_path) if enable_text and tokenizer_path else None) |
| 45 | self.pipe = build_wan_video_action_pipeline(torch_dtype=torch.bfloat16, device=device, model_configs=model_configs, tokenizer_config=tokenizer_config, args=args) |
| 46 | self.pipe = self.split_pipeline_units(task, self.pipe, trainable_models, lora_base_model) |
| 47 | |
| 48 | # Training mode |
| 49 | self.switch_pipe_to_training_mode( |
| 50 | self.pipe, trainable_models, |
| 51 | lora_base_model, lora_target_modules, lora_rank, lora_checkpoint, |
| 52 | preset_lora_path, preset_lora_model, |
| 53 | task=task, |
| 54 | ) |
| 55 | |
| 56 | if not enable_text: |
| 57 | self.freeze_text_modules() |
| 58 | |
| 59 | # Store other configs |
| 60 | self.use_gradient_checkpointing = use_gradient_checkpointing |
| 61 | self.use_gradient_checkpointing_offload = use_gradient_checkpointing_offload |
| 62 | self.extra_inputs = extra_inputs.split(",") if extra_inputs is not None else [] |
| 63 | self.fp8_models = fp8_models |
| 64 | self.task = task |
| 65 | self.task_to_loss = { |
| 66 | "sft:data_process": lambda pipe, *args: args, |
| 67 | "direct_distill:data_process": lambda pipe, *args: args, |
| 68 | "sft": lambda pipe, inputs_shared, inputs_posi, inputs_nega: FlowMatchSFTLoss(pipe, **inputs_shared, **inputs_posi), |
| 69 | "sft:train": lambda pipe, inputs_shared, inputs_posi, inputs_nega: FlowMatchSFTLoss(pipe, **inputs_shared, **inputs_posi), |
| 70 | "direct_distill": lambda pipe, inputs_shared, inputs_posi, inputs_nega: DirectDistillLoss(pipe, **inputs_shared, **inputs_posi), |
| 71 | "direct_distill:train": lambda pipe, inputs_shared, inputs_posi, inputs_nega: DirectDistillLoss(pipe, **inputs_shared, **inputs_posi), |
nothing calls this directly
no test coverage detected