Compare commits

...
Author SHA1 Message Date
Will Lin efd886a188 fix 2025-06-10 15:14:07 -07:00
3 changed files with 7 additions and 7 deletions
+2 -4
View File
@@ -581,9 +581,6 @@ class TrainingArgs(FastVideoArgs):
# master_weight_type
master_weight_type: str = ""
# For fast checking in LoRA pipeline
training_mode: bool = True
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
# Get all fields from the dataclass
@@ -606,7 +603,8 @@ class TrainingArgs(FastVideoArgs):
# Use getattr with default value from the dataclass for potentially missing attributes
else:
default_value = getattr(cls, attr, None)
kwargs[attr] = getattr(args, attr, default_value)
if getattr(args, attr, default_value) is not None:
kwargs[attr] = getattr(args, attr, default_value)
return cls(**kwargs)
@@ -55,6 +55,7 @@ class ComposedPipelineBase(ABC):
use. The pipeline should be stateless and not hold any batch state.
"""
assert fastvideo_args.training_mode != fastvideo_args.inference_mode, "training_mode and inference_mode cannot be both True"
if fastvideo_args.training_mode:
assert isinstance(fastvideo_args, TrainingArgs)
self.training_args = fastvideo_args
@@ -55,10 +55,11 @@ class WanTrainingPipeline(TrainingPipeline):
args_copy.inference_mode = True
args_copy.vae_config.load_encoder = False
validation_pipeline = WanValidationPipeline.from_pretrained(
args_copy.log_validation = False
# no need to use .from_pretrained here as we already have TrainingArgs
validation_pipeline = WanValidationPipeline(
training_args.model_path,
args=None,
inference_mode=True,
fastvideo_args=args_copy,
loaded_modules={"transformer": self.get_module("transformer")})
self.validation_pipeline = validation_pipeline