Skip to content

Commit 105e034

Browse files
committed
[misc]: Use runner module in all tests and add unit test for runner dynamic instantiation
1 parent 232e677 commit 105e034

5 files changed

Lines changed: 50 additions & 37 deletions

File tree

fastvideo/tests/training/Vanilla/mfu_calculation.py

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,7 @@
1717

1818
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
1919
from fastvideo.utils import FlexibleArgumentParser
20-
from fastvideo.training.wan_training_pipeline import WanTrainingPipeline
20+
from fastvideo.training.runner import main
2121

2222
MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
2323
DATA_PATH = "data/crush-smol_processed_t2v/training_dataset/worker_1/worker_0/"
@@ -123,10 +123,11 @@ def run_worker():
123123
"1"
124124
])
125125
# Call the main training function
126-
pipeline = WanTrainingPipeline.from_pretrained(args.pretrained_model_name_or_path, args=args)
127-
args = pipeline.training_args
128-
pipeline.train()
129-
logger.info("Training pipeline done")
126+
args.pipeline_class = "WanTrainingPipeline"
127+
args.pipeline_module = "fastvideo.training.wan_training_pipeline"
128+
129+
# Call the main training function
130+
main(args)
130131

131132

132133
def test_distributed_training():

fastvideo/tests/training/Vanilla/test_training_loss.py

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@
1414

1515
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
1616
from fastvideo.utils import FlexibleArgumentParser
17-
from fastvideo.training.wan_training_pipeline import WanTrainingPipeline
17+
from fastvideo.training.runner import main
1818

1919
wandb_name = "test_training_loss"
2020
a40_reference_wandb_summary_file = "fastvideo/tests/training/Vanilla/a40_reference_wandb_summary.json"
@@ -50,10 +50,11 @@ def run_worker():
5050
"--not_apply_cfg_solver", "--dit_precision", "fp32", "--max_grad_norm", "1.0"
5151
])
5252
# Call the main training function
53-
pipeline = WanTrainingPipeline.from_pretrained(args.pretrained_model_name_or_path, args=args)
54-
args = pipeline.training_args
55-
pipeline.train()
56-
logger.info("Training pipeline done")
53+
args.pipeline_class = "WanTrainingPipeline"
54+
args.pipeline_module = "fastvideo.training.wan_training_pipeline"
55+
56+
# Call the main training function
57+
main(args)
5758

5859

5960
def test_distributed_training():

fastvideo/tests/training/distill/test_distill_dmd.py

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@
1414

1515
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
1616
from fastvideo.utils import FlexibleArgumentParser
17-
from fastvideo.training.wan_distillation_pipeline import WanDistillationPipeline
17+
from fastvideo.training.runner import main
1818

1919
wandb_name = "test_distill_dmd"
2020

@@ -49,10 +49,11 @@ def run_worker():
4949
"--enable_gradient_checkpointing_type", "full"
5050
])
5151
# Call the main training function
52-
pipeline = WanDistillationPipeline.from_pretrained(args.pretrained_model_name_or_path, args=args)
53-
args = pipeline.training_args
54-
pipeline.train()
55-
logger.info("Training pipeline done")
52+
args.pipeline_class = "WanDistillationPipeline"
53+
args.pipeline_module = "fastvideo.training.wan_distillation_pipeline"
54+
55+
# Call the main training function
56+
main(args)
5657

5758

5859
def test_distributed_training():

fastvideo/tests/training/self-forcing/test_self_forcing.py

Lines changed: 7 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -11,9 +11,10 @@
1111
from fastvideo.utils import logger
1212
# Import the training pipeline
1313
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
14-
from fastvideo.training.wan_self_forcing_distillation_pipeline import WanSelfForcingDistillationPipeline
14+
1515
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
1616
from fastvideo.utils import FlexibleArgumentParser
17+
from fastvideo.training.runner import main
1718

1819
wandb_name = "test_self_forcing_distill"
1920

@@ -148,10 +149,11 @@ def run_worker():
148149
])
149150

150151
# Call the main training function
151-
pipeline = WanSelfForcingDistillationPipeline.from_pretrained(args.pretrained_model_name_or_path, args=args)
152-
args = pipeline.training_args
153-
pipeline.train()
154-
logger.info("Self-forcing distillation training pipeline done")
152+
args.pipeline_class = "WanSelfForcingDistillationPipeline"
153+
args.pipeline_module = "fastvideo.training.wan_self_forcing_distillation_pipeline"
154+
155+
# Call the main training function
156+
main(args)
155157

156158

157159
def test_distributed_training():
Lines changed: 25 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -1,18 +1,26 @@
1-
import subprocess
2-
import sys
1+
from unittest.mock import patch, MagicMock
2+
from fastvideo.training.runner import main
3+
import argparse
34

4-
def test_runner_help():
5-
# Run the runner with --help to ensure it parses arguments correctly
6-
# and doesn't have any syntax errors or missing imports.
7-
result = subprocess.run(
8-
[sys.executable, "-m", "fastvideo.training.runner", "--help"],
9-
capture_output=True,
10-
text=True
11-
)
12-
13-
# Check if the process completed successfully
14-
assert result.returncode == 0, f"Runner failed with exit code {result.returncode}.\nStderr: {result.stderr}"
15-
16-
# Check if help output is printed
17-
assert "--pipeline_class" in result.stdout
18-
assert "--pipeline_module" in result.stdout
5+
def test_runner_invokes_correct_class():
6+
with patch('fastvideo.training.runner.importlib.import_module') as mock_import_module:
7+
mock_module = MagicMock()
8+
mock_import_module.return_value = mock_module
9+
10+
mock_pipeline_class = MagicMock()
11+
setattr(mock_module, "WanTrainingPipeline", mock_pipeline_class)
12+
13+
mock_pipeline_instance = MagicMock()
14+
mock_pipeline_class.from_pretrained.return_value = mock_pipeline_instance
15+
16+
args = argparse.Namespace(
17+
pipeline_module="fastvideo.training.wan_training_pipeline",
18+
pipeline_class="WanTrainingPipeline",
19+
pretrained_model_name_or_path="test_model"
20+
)
21+
22+
main(args)
23+
24+
mock_import_module.assert_called_once_with("fastvideo.training.wan_training_pipeline")
25+
mock_pipeline_class.from_pretrained.assert_called_once_with("test_model", args=args)
26+
mock_pipeline_instance.train.assert_called_once()

0 commit comments

Comments
 (0)