MCPcopy Create free account
hub / github.com/boundless-large-model/boundless-world-model / __init__

Method __init__

scripts/train.py:14–75  ·  view source on GitHub ↗
(
        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,
    )

Source from the content-addressed store, hash-verified

12
13class 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),

Callers

nothing calls this directly

Calls 2

freeze_text_modulesMethod · 0.95

Tested by

no test coverage detected