Skip to content

flow_matching

Shape-agnostic conditional flow-matching fine-tuning.

Classes

fastvideo.train.methods.fine_tuning.flow_matching.FlowMatchingFineTuneMethod

FlowMatchingFineTuneMethod(*, cfg: Any, role_models: dict[str, ModelBase])

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)

Functions: