ValidationCallback(*, pipeline_target: str, dataset_file: str, every_steps: int = 100, run_at_start: bool = True, sampling_steps: list[int] | None = None, guidance_scale: float | None = None, num_frames: int | None = None, num_videos_per_prompt: int = 1, use_validation_media_conditioning: bool = True, output_dir: str | None = None, sampling_timesteps: list[int] | None = None, overlay_actions: bool = False, keyboard_value_scale: float = 1.0, offload_training_state: bool = False, unload_pipeline_after_validation: bool = False, attn_qat_infer: bool = False, **pipeline_kwargs: Any)
Bases: Callback
Generic validation callback driven entirely by YAML config.
Works with any pipeline that follows the PipelineCls.from_pretrained(...) + pipeline.forward() contract.
Configure validation cadence, generation parameters, and pipeline loading.
run_at_start controls the pre-training baseline event. use_validation_media_conditioning lets text-to-video recipes use captions from a dataset that also contains source-media paths.
Source code in fastvideo/train/callbacks/validation.py
| def __init__(
self,
*,
pipeline_target: str,
dataset_file: str,
every_steps: int = 100,
run_at_start: bool = True,
sampling_steps: list[int] | None = None,
guidance_scale: float | None = None,
num_frames: int | None = None,
num_videos_per_prompt: int = 1,
use_validation_media_conditioning: bool = True,
output_dir: str | None = None,
sampling_timesteps: list[int] | None = None,
overlay_actions: bool = False,
keyboard_value_scale: float = 1.0,
offload_training_state: bool = False,
unload_pipeline_after_validation: bool = False,
attn_qat_infer: bool = False,
**pipeline_kwargs: Any,
) -> None:
"""Configure validation cadence, generation parameters, and pipeline loading.
``run_at_start`` controls the pre-training baseline event.
``use_validation_media_conditioning`` lets text-to-video recipes use
captions from a dataset that also contains source-media paths.
"""
self.pipeline_target = str(pipeline_target)
self.dataset_file = str(dataset_file)
self.every_steps = int(every_steps)
self.run_at_start = self._coerce_bool(run_at_start)
self.sampling_steps = ([int(s) for s in sampling_steps] if sampling_steps else [40])
self.guidance_scale = (float(guidance_scale) if guidance_scale is not None else None)
self.num_frames = (int(num_frames) if num_frames is not None else None)
self.num_videos_per_prompt = int(num_videos_per_prompt)
if self.num_videos_per_prompt <= 0:
raise ValueError("callbacks.validation.num_videos_per_prompt must be positive")
self.use_validation_media_conditioning = self._coerce_bool(use_validation_media_conditioning)
self.output_dir = (str(output_dir) if output_dir is not None else None)
self.sampling_timesteps = ([int(s) for s in sampling_timesteps] if sampling_timesteps is not None else None)
self.overlay_actions = self._coerce_bool(overlay_actions)
# Validation-only action amplification for world model; training keeps raw action values.
self.keyboard_value_scale = float(keyboard_value_scale)
metrics_config = pipeline_kwargs.pop("metrics", None)
self.metrics_config = self._parse_metrics_config(metrics_config)
self.offload_training_state = self._coerce_bool(offload_training_state)
self.unload_pipeline_after_validation = self._coerce_bool(unload_pipeline_after_validation)
self.attn_qat_infer = self._coerce_bool(attn_qat_infer)
self.pipeline_kwargs = dict(pipeline_kwargs)
# Set after on_train_start.
self._pipeline: Any | None = None
self._pipeline_key: tuple[Any, ...] | None = None
self._sampling_param: SamplingParam | None = None
self._metric_evaluator: Any | None = None
self.tracker: Any = DummyTracker()
self.validation_random_generator: (torch.Generator | None) = None
self.seed: int = 0
|
Methods:
fastvideo.train.callbacks.validation.ValidationCallback.on_validation_begin
Run the optional step-zero baseline and each scheduled validation event.
Source code in fastvideo/train/callbacks/validation.py
| def on_validation_begin(
self,
method: TrainingMethod,
iteration: int = 0,
) -> None:
"""Run the optional step-zero baseline and each scheduled validation event."""
if self.every_steps <= 0:
return
# Step zero measures the checkpoint before the first optimizer update.
if iteration == 0 and not self.run_at_start:
return
if iteration % self.every_steps != 0:
return
self._run_validation(method, iteration)
|