Bases: FineTuneMethod
Fine-tune a model against a velocity target supplied by its adapter.
Unlike :class:FineTuneMethod, this method does not assume a five- dimensional video latent layout. The model plugin owns sampling the flow path and stores its target in TrainingBatch.training_target.
Source code in fastvideo/train/methods/fine_tuning/flow_matching.py
| def __init__(
self,
*,
cfg: Any,
role_models: dict[str, ModelBase],
) -> None:
super().__init__(cfg=cfg, role_models=role_models)
flow_matching_forward = getattr(self.student, "flow_matching_forward", None)
if not callable(flow_matching_forward):
raise TypeError("FlowMatchingFineTuneMethod requires the student model to "
"implement flow_matching_forward()")
self._flow_matching_forward: TensorTrainFn = flow_matching_forward
self._train_fn: TensorTrainFn = self._eager_train_fn
if bool(self.training_config.model.compile_train_fn):
if not is_ddp_strategy(self.training_config):
raise ValueError("training.model.compile_train_fn is currently supported "
"only with training.distributed.strategy=ddp")
compile_kwargs = dict(self.training_config.model.torch_compile_kwargs)
compile_kwargs.setdefault("fullgraph", True)
compile_training_forward = getattr(self.student, "compile_training_forward", None)
if not callable(compile_training_forward):
raise TypeError("compile_train_fn requires the student model to implement "
"compile_training_forward()")
logger.info(
"Enabling fullgraph transformer compilation with kwargs=%s",
compile_kwargs,
)
compile_training_forward(compile_kwargs)
|