Skip to content

Commit b0d254b

Browse files
committed
[refactor]: remove monolithic main blocks from legacy training pipelines
1 parent e64fe04 commit b0d254b

12 files changed

Lines changed: 396 additions & 640 deletions

fastvideo/training/cosmos2_5_training_pipeline.py

Lines changed: 0 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -133,22 +133,3 @@ def _build_input_kwargs(self, training_batch: TrainingBatch) -> TrainingBatch:
133133
# Entry point (mirrors wan_training_pipeline.py)
134134
# ---------------------------------------------------------------------------
135135

136-
137-
def main(args) -> None:
138-
logger.info("Starting Cosmos 2.5 training pipeline...")
139-
pipeline = Cosmos25TrainingPipeline.from_pretrained(args.pretrained_model_name_or_path, args=args)
140-
args = pipeline.training_args
141-
pipeline.train()
142-
logger.info("Training pipeline done")
143-
144-
145-
if __name__ == "__main__":
146-
from fastvideo.fastvideo_args import TrainingArgs
147-
from fastvideo.utils import FlexibleArgumentParser
148-
149-
parser = FlexibleArgumentParser()
150-
parser = TrainingArgs.add_cli_args(parser)
151-
parser = FastVideoArgs.add_cli_args(parser)
152-
args = parser.parse_args()
153-
args.dit_cpu_offload = False
154-
main(args)

fastvideo/training/ltx2_training_pipeline.py

Lines changed: 0 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -474,23 +474,3 @@ def _clip_grad_norm(self, training_batch: TrainingBatch) -> TrainingBatch:
474474
training_batch.grad_norm = grad_norm
475475
return training_batch
476476

477-
478-
def main(args) -> None:
479-
logger.info("Starting LTX-2 training pipeline...")
480-
pipeline = LTX2TrainingPipeline.from_pretrained(args.pretrained_model_name_or_path, args=args)
481-
args = pipeline.training_args
482-
pipeline.train()
483-
logger.info("Training pipeline done")
484-
485-
486-
if __name__ == "__main__":
487-
argv = sys.argv
488-
from fastvideo.fastvideo_args import TrainingArgs
489-
from fastvideo.utils import FlexibleArgumentParser
490-
491-
parser = FlexibleArgumentParser()
492-
parser = TrainingArgs.add_cli_args(parser)
493-
parser = FastVideoArgs.add_cli_args(parser)
494-
args = parser.parse_args()
495-
args.dit_cpu_offload = False
496-
main(args)

fastvideo/training/matrixgame2_ar_diffusion_pipeline.py

Lines changed: 0 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -427,24 +427,3 @@ def _prepare_validation_batch(
427427

428428
return batch
429429

430-
431-
def main(args) -> None:
432-
logger.info("Starting Matrix-Game 2.0 AR diffusion training pipeline...")
433-
434-
pipeline = MatrixGame2ARDiffusionPipeline.from_pretrained(args.pretrained_model_name_or_path, args=args)
435-
args = pipeline.training_args
436-
pipeline.train()
437-
logger.info("Matrix-Game 2.0 AR diffusion training pipeline done")
438-
439-
440-
if __name__ == "__main__":
441-
argv = sys.argv
442-
from fastvideo.fastvideo_args import TrainingArgs
443-
from fastvideo.utils import FlexibleArgumentParser
444-
445-
parser = FlexibleArgumentParser()
446-
parser = TrainingArgs.add_cli_args(parser)
447-
parser = FastVideoArgs.add_cli_args(parser)
448-
args = parser.parse_args()
449-
args.dit_cpu_offload = False
450-
main(args)

fastvideo/training/matrixgame2_ode_causal_pipeline.py

Lines changed: 396 additions & 415 deletions
Large diffs are not rendered by default.

fastvideo/training/matrixgame2_self_forcing_distillation_pipeline.py

Lines changed: 0 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -875,23 +875,3 @@ def _prepare_validation_batch(self, sampling_param: SamplingParam, training_args
875875

876876
return batch
877877

878-
879-
def main(args) -> None:
880-
logger.info("Starting Matrix-Game 2.0 self-forcing distillation pipeline...")
881-
882-
pipeline = MatrixGame2SelfForcingDistillationPipeline.from_pretrained(args.pretrained_model_name_or_path, args=args)
883-
884-
args = pipeline.training_args
885-
pipeline.train()
886-
logger.info("Matrix-Game 2.0 self-forcing distillation pipeline completed")
887-
888-
889-
if __name__ == "__main__":
890-
argv = sys.argv
891-
from fastvideo.fastvideo_args import TrainingArgs
892-
from fastvideo.utils import FlexibleArgumentParser
893-
parser = FlexibleArgumentParser()
894-
parser = TrainingArgs.add_cli_args(parser)
895-
parser = FastVideoArgs.add_cli_args(parser)
896-
args = parser.parse_args()
897-
main(args)

fastvideo/training/matrixgame2_training_pipeline.py

Lines changed: 0 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -194,23 +194,3 @@ def _prepare_validation_batch(self, sampling_param: SamplingParam, training_args
194194

195195
return batch
196196

197-
198-
def main(args) -> None:
199-
logger.info("Starting training pipeline...")
200-
201-
pipeline = MatrixGame2TrainingPipeline.from_pretrained(args.pretrained_model_name_or_path, args=args)
202-
args = pipeline.training_args
203-
pipeline.train()
204-
logger.info("Training pipeline done")
205-
206-
207-
if __name__ == "__main__":
208-
argv = sys.argv
209-
from fastvideo.fastvideo_args import TrainingArgs
210-
from fastvideo.utils import FlexibleArgumentParser
211-
parser = FlexibleArgumentParser()
212-
parser = TrainingArgs.add_cli_args(parser)
213-
parser = FastVideoArgs.add_cli_args(parser)
214-
args = parser.parse_args()
215-
args.dit_cpu_offload = False
216-
main(args)

fastvideo/training/ode_causal_pipeline.py

Lines changed: 0 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -309,22 +309,3 @@ def visualize_intermediate_latents(self, training_batch: TrainingBatch, training
309309
if self.global_rank == 0 and tracker_loss_dict:
310310
self.tracker.log_artifacts(tracker_loss_dict, step)
311311

312-
313-
def main(args) -> None:
314-
logger.info("Starting ODE-init training pipeline...")
315-
pipeline = ODEInitTrainingPipeline.from_pretrained(args.pretrained_model_name_or_path, args=args)
316-
args = pipeline.training_args
317-
pipeline.train()
318-
logger.info("ODE-init training pipeline done")
319-
320-
321-
if __name__ == "__main__":
322-
argv = sys.argv
323-
from fastvideo.fastvideo_args import TrainingArgs
324-
from fastvideo.utils import FlexibleArgumentParser
325-
parser = FlexibleArgumentParser()
326-
parser = TrainingArgs.add_cli_args(parser)
327-
parser = FastVideoArgs.add_cli_args(parser)
328-
args = parser.parse_args()
329-
args.dit_cpu_offload = False
330-
main(args)

fastvideo/training/wan_distillation_pipeline.py

Lines changed: 0 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -52,25 +52,3 @@ def initialize_validation_pipeline(self, training_args: TrainingArgs):
5252

5353
self.validation_pipeline = validation_pipeline
5454

55-
56-
def main(args) -> None:
57-
logger.info("Starting Wan distillation pipeline...")
58-
59-
# Create pipeline with original args
60-
pipeline = WanDistillationPipeline.from_pretrained(args.pretrained_model_name_or_path, args=args)
61-
62-
args = pipeline.training_args
63-
# Start training
64-
pipeline.train()
65-
logger.info("Wan distillation pipeline completed")
66-
67-
68-
if __name__ == "__main__":
69-
argv = sys.argv
70-
from fastvideo.fastvideo_args import TrainingArgs
71-
from fastvideo.utils import FlexibleArgumentParser
72-
parser = FlexibleArgumentParser()
73-
parser = TrainingArgs.add_cli_args(parser)
74-
parser = FastVideoArgs.add_cli_args(parser)
75-
args = parser.parse_args()
76-
main(args)

fastvideo/training/wan_i2v_distillation_pipeline.py

Lines changed: 0 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -174,27 +174,3 @@ def _build_distill_input_kwargs(self, noise_input: torch.Tensor, timestep: torch
174174

175175
return training_batch
176176

177-
178-
def main(args) -> None:
179-
logger.info("Starting Wan distillation pipeline...")
180-
181-
# Create pipeline with original args
182-
pipeline = WanI2VDistillationPipeline.from_pretrained(args.pretrained_model_name_or_path, args=args)
183-
184-
args = pipeline.training_args
185-
186-
# Start training
187-
pipeline.train()
188-
logger.info("Wan distillation pipeline completed")
189-
190-
191-
if __name__ == "__main__":
192-
argv = sys.argv
193-
from fastvideo.fastvideo_args import TrainingArgs
194-
from fastvideo.utils import FlexibleArgumentParser
195-
parser = FlexibleArgumentParser()
196-
parser = TrainingArgs.add_cli_args(parser)
197-
parser = FastVideoArgs.add_cli_args(parser)
198-
args = parser.parse_args()
199-
args.dit_cpu_offload = False
200-
main(args)

fastvideo/training/wan_i2v_training_pipeline.py

Lines changed: 0 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -165,23 +165,3 @@ def _prepare_validation_batch(self, sampling_param: SamplingParam, training_args
165165

166166
return batch
167167

168-
169-
def main(args) -> None:
170-
logger.info("Starting training pipeline...")
171-
172-
pipeline = WanI2VTrainingPipeline.from_pretrained(args.pretrained_model_name_or_path, args=args)
173-
args = pipeline.training_args
174-
pipeline.train()
175-
logger.info("Training pipeline done")
176-
177-
178-
if __name__ == "__main__":
179-
argv = sys.argv
180-
from fastvideo.fastvideo_args import TrainingArgs
181-
from fastvideo.utils import FlexibleArgumentParser
182-
parser = FlexibleArgumentParser()
183-
parser = TrainingArgs.add_cli_args(parser)
184-
parser = FastVideoArgs.add_cli_args(parser)
185-
args = parser.parse_args()
186-
args.dit_cpu_offload = False
187-
main(args)

0 commit comments

Comments
 (0)