[refactor]: build resolved configs in the modular trainer

- moduleloader builds the component-loading config and the inference
  config as resolved configs (_build_training_resolved_config and
  build_inference_resolved_config) instead of TrainingArgs objects mutated
  after construction; the transformer override, weights path, and
  teacher/critic flag are config values or load_module arguments.
  keep_checkpoint_component_config keeps the trainer's PipelineConfig view
  of the checkpoint-filled VAE and text-encoder configs.
- The validation callback builds its pipeline with
  from_pretrained(path, resolved_config=...) and forwards with a resolved
  config built per run; it no longer replaces or mutates the trainer's
  pipeline_config, and generator_overrides takes nested GeneratorConfig
  fields instead of flat names.
- Against the 3f6893a0 snapshots of the 67 trainer YAML files, every
  PipelineConfig and value matches except ltx2_vae_tiling (the model's
  vae_tiling) and the unread pretrained_model_name_or_path.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_012M91wnVFmPEvJ9h39r5BH7
This commit is contained in:
Davids048
2026-10-03 06:19:13 +00:00
co-authored by Claude Opus 5.5
parent 4117696a4f
commit 00a8e5bcb6
15 changed files with 497 additions and 207 deletions
@@ -131,8 +131,9 @@ export default function CreateJobModal({
const editingJobId = editingJob?.id ?? null;
const editingJobModelId = editingJob?.model_id ?? null;
// Layerwise offload and FSDP compete for the DiT weights and FastVideoArgs
// silently picks a winner (fastvideo_args.py:859); resolve it visibly here.
// Layerwise offload and FSDP compete for the DiT weights and the device offload
// policy (resolve_device_offload_conflicts in fastvideo/api/device_policy.py)
// silently picks a winner; resolve it visibly here.
// dit_cpu_offload is deliberately not interlocked -- it is a modifier, not a
// competing strategy.
const handleDitLayerwiseOffloadChange = React.useCallback((next: boolean) => {
@@ -19,7 +19,7 @@
# The export is passed to all three models.*.init_from overrides: student,
# teacher, and critic all load the same weights. Teacher/critic are
# automatically masked back to dense attention and full-precision weights
# by the _loading_teacher_critic_model gate in
# by the loading_teacher_critic_model gate in
# fastvideo/models/loader/component_loader.py (family-agnostic, no
# Kandinsky5-specific handling needed) -- NOT by pointing them at a
# different export.
@@ -11,7 +11,7 @@
#
# - Teacher: frozen, full-precision Kandinsky5 (loaded from the stage-1 QAT
# checkpoint's base weights, or the original pretrained checkpoint --
# the _loading_teacher_critic_model gate in component_loader.py masks
# the loading_teacher_critic_model gate in component_loader.py masks
# quant_config and FASTVIDEO_ATTENTION_BACKEND for teacher/critic
# regardless, so they always run full precision / dense attention).
# - Student: trainable, initialized from the stage-1 QAT finetune checkpoint.
@@ -90,7 +90,7 @@ conversion and uses it directly.)
`--role student` exports the trained transformer weights (the only role
that matters here: student/teacher/critic in stage 2 all load from the same
checkpoint, and `_loading_teacher_critic_model` handles the full-precision
checkpoint, and `loading_teacher_critic_model` handles the full-precision
masking for teacher/critic at load time, not at export time). `--verify`
strictly reloads the exported transformer immediately, so a key-mapping bug
fails here instead of deep inside the stage-2 launch.
@@ -98,7 +98,7 @@ fails here instead of deep inside the stage-2 launch.
Only the student is quantized (Attn-QAT); the teacher and critic stay full
precision / dense attention. This is enforced in the loader
(`fastvideo/models/loader/component_loader.py`, via the
`_loading_teacher_critic_model` flag), which masks `quant_config` and clears
`loading_teacher_critic_model` flag), which masks `quant_config` and clears
`FASTVIDEO_ATTENTION_BACKEND` for teacher/critic -- the same global env var
reaches only the student, with no per-model flags or Kandinsky5-specific
handling. Validation runs the distilled student through
@@ -136,26 +136,31 @@ still says the base T2V pipeline and the registry
(`fastvideo/registry.py`) has no way to auto-detect that a given directory
is actually a DMD (four-step re-noise sampler) export -- it will resolve to
`Kandinsky5T2VPipeline` and run the full-length sampler on DMD-distilled
weights. Pass `override_pipeline_cls_name` and a `Kandinsky5DMDConfig`
explicitly to select the right pipeline/sampler and get
`dmd_denoising_steps` set (`Kandinsky5T2VConfig`'s default of `None` makes
weights. Set `pipeline.components.override_pipeline_cls_name` and pass a
`Kandinsky5DMDConfig` as `pipeline.experimental.pipeline_config` to select
the right pipeline/sampler and get `pipeline.dmd_denoising_steps` set
(`Kandinsky5T2VConfig`'s default of `None` makes
`Kandinsky5DmdDenoisingStage` raise immediately):
```python
from fastvideo import VideoGenerator
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5DMDConfig
from fastvideo.layers.quantization import get_quantization_config
gen = VideoGenerator.from_pretrained(
"path/to/kandinsky5_dmd_checkpoint", num_gpus=1,
override_pipeline_cls_name="Kandinsky5DMDPipeline",
pipeline_config=Kandinsky5DMDConfig(),
transformer_quant=get_quantization_config("nvfp4_qat")(),
use_fsdp_inference=False,
)
gen = VideoGenerator.from_config({
"model_path": "path/to/kandinsky5_dmd_checkpoint",
"engine": {
"num_gpus": 1,
"use_fsdp_inference": False,
"quantization": {"transformer_quant": "nvfp4_qat"},
},
"pipeline": {
"components": {"override_pipeline_cls_name": "Kandinsky5DMDPipeline"},
"experimental": {"pipeline_config": Kandinsky5DMDConfig()},
},
})
# Kandinsky5DmdDenoisingStage ignores request.sampling.num_inference_steps
# and guidance_scale -- it's a fixed 4-step, no-CFG sampler driven entirely
# by pipeline_config.dmd_denoising_steps (set above). Adjust
# by pipeline.dmd_denoising_steps, which Kandinsky5DMDConfig sets. Adjust
# Kandinsky5DMDConfig(dmd_denoising_steps=[...]) instead if the checkpoint
# was distilled/validated with a different schedule than the default
# [1000, 750, 500, 250].
@@ -176,8 +181,9 @@ to a supported dense backend while the FP4 linear layers still run.
and that gradients flow through a forward+backward pass.
- `fastvideo/tests/api/test_kandinsky5_dmd_pipeline_resolution.py` --
confirms an unmodified export resolves to `Kandinsky5T2VPipeline` (the bug)
and that `override_pipeline_cls_name="Kandinsky5DMDPipeline"` fixes it (the
documented workaround above), plus `Kandinsky5DMDConfig`'s
and that `pipeline.components.override_pipeline_cls_name:
Kandinsky5DMDPipeline` fixes it (the documented workaround above), plus
`Kandinsky5DMDConfig`'s
`dmd_denoising_steps` default.
- `fastvideo/tests/nightly/test_e2e_kandinsky5_dmd_t2v_overfit.py` -- a few
steps of both training stages on a single synthetic sample, exercising the
@@ -80,7 +80,9 @@ callbacks:
num_videos_per_prompt: 1
use_validation_media_conditioning: false
offload_training_state: true
text_encoder_cpu_offload: true
vae_cpu_offload: true
engine:
offload:
text_encoder: true
vae: true
pipeline: {}
@@ -21,6 +21,7 @@ import pytest
import torch
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.train.callbacks.callback import CallbackDict
from fastvideo.train.callbacks.ema import EMACallback
from fastvideo.train.callbacks.validation import (
@@ -31,6 +32,7 @@ from fastvideo.train.callbacks.validation import (
_ValidationMetricStats,
_ValidationStepResult,
)
from fastvideo.train.utils.training_config import DistributedConfig, TrainingConfig
# ---------------------------------------------------------------------------
# Helpers
@@ -138,17 +140,17 @@ class TestConstructor:
with pytest.raises(ValueError, match="num_videos_per_prompt must be positive"):
_make_callback(num_videos_per_prompt=0)
def test_pipeline_kwargs_collected(self) -> None:
def test_generator_overrides_collected(self) -> None:
cb = ValidationCallback(
pipeline_target=_PIPE_TARGET,
dataset_file="x.json",
extra_arg=123,
another="value",
engine={"offload": {"vae": True}},
pipeline={"flow_shift": 5.0},
)
# Unknown kwargs are stashed for the pipeline factory.
assert cb.pipeline_kwargs == {
"extra_arg": 123,
"another": "value",
# Unknown keys are stashed as GeneratorConfig fields of the validation pipeline.
assert cb.generator_overrides == {
"engine": {"offload": {"vae": True}},
"pipeline": {"flow_shift": 5.0},
}
def test_metrics_true_uses_default_vbench_subset(self) -> None:
@@ -159,7 +161,7 @@ class TestConstructor:
)
assert cb.metrics_config.enabled is True
assert cb.metrics_config.names == DEFAULT_VALIDATION_VBENCH_METRICS
assert "metrics" not in cb.pipeline_kwargs
assert "metrics" not in cb.generator_overrides
def test_metrics_mapping_is_coerced(self) -> None:
cb = ValidationCallback(
@@ -263,10 +265,7 @@ class TestOnValidationBegin:
class TestH3ValidationContract:
"""Verify synchronized MiniMax H3 validation through Weights & Biases logging."""
def test_prepare_validation_batch_forwards_video_count(
self,
monkeypatch: pytest.MonkeyPatch,
) -> None:
def test_prepare_validation_batch_forwards_video_count(self) -> None:
"""Verify that each ForwardBatch receives the configured output count."""
cb = _make_callback(num_videos_per_prompt=3)
cb.training_config = SimpleNamespace(
@@ -280,11 +279,6 @@ class TestH3ValidationContract:
model_path="unused",
vsa_sparsity=0.0,
)
monkeypatch.setattr(
"fastvideo.train.callbacks.validation.make_inference_args",
lambda *args, **kwargs: SimpleNamespace(),
)
batch = cb._prepare_validation_batch(
SamplingParam(),
{"prompt": "Generate synchronized media."},
@@ -295,7 +289,6 @@ class TestH3ValidationContract:
def test_prepare_validation_batch_ignores_media_for_text_only_generation(
self,
monkeypatch: pytest.MonkeyPatch,
tmp_path,
) -> None:
"""Verify text-to-video-with-audio uses captions without source media."""
@@ -313,11 +306,6 @@ class TestH3ValidationContract:
model_path="unused",
vsa_sparsity=0.0,
)
monkeypatch.setattr(
"fastvideo.train.callbacks.validation.make_inference_args",
lambda *args, **kwargs: SimpleNamespace(),
)
batch = cb._prepare_validation_batch(
SamplingParam(),
{
@@ -341,7 +329,7 @@ class TestH3ValidationContract:
resolved_config = SimpleNamespace(pipeline_config=SimpleNamespace())
def forward(self, batch, inference_args):
def forward(self, batch, resolved_config):
"""Produce media whose values identify the source prompt."""
prompt_index = int(batch.prompt.rsplit("-", 1)[1])
return SimpleNamespace(
@@ -363,13 +351,7 @@ class TestH3ValidationContract:
"fastvideo.train.callbacks.validation.ValidationDataset",
lambda filename: validation_records,
)
monkeypatch.setattr(
"fastvideo.train.callbacks.validation.make_inference_args",
lambda *args, **kwargs: SimpleNamespace(
pipeline_config=SimpleNamespace(dmd_denoising_steps=None),
dit_cpu_offload=True,
),
)
monkeypatch.setattr(cb, "_validation_forward_config", lambda pipeline, transformer: SimpleNamespace())
monkeypatch.setattr(cb, "_get_pipeline", lambda *, transformer: pipeline)
monkeypatch.setattr(cb, "_get_sampling_param", SamplingParam)
monkeypatch.setattr(
@@ -401,7 +383,7 @@ class TestH3ValidationContract:
resolved_config = SimpleNamespace(pipeline_config=SimpleNamespace())
def forward(self, batch, inference_args):
def forward(self, batch, resolved_config):
return SimpleNamespace(output=output, extra={})
cb = _make_callback()
@@ -415,13 +397,7 @@ class TestH3ValidationContract:
"fastvideo.train.callbacks.validation.ValidationDataset",
lambda filename: [{"caption": "prompt-0"}],
)
monkeypatch.setattr(
"fastvideo.train.callbacks.validation.make_inference_args",
lambda *args, **kwargs: SimpleNamespace(
pipeline_config=SimpleNamespace(dmd_denoising_steps=None),
dit_cpu_offload=True,
),
)
monkeypatch.setattr(cb, "_validation_forward_config", lambda pipeline, transformer: SimpleNamespace())
monkeypatch.setattr(cb, "_get_pipeline", lambda *, transformer: pipeline)
monkeypatch.setattr(cb, "_get_sampling_param", SamplingParam)
monkeypatch.setattr(
@@ -1059,8 +1035,8 @@ class TestKeepLoadedEncoderWidths:
assert validation.text_encoder_configs[0].arch_config.text_len == 1000
def test_keeps_config_objects_unshared(self) -> None:
# The second call site writes into ``tc.pipeline_config`` itself, so
# the merge must not alias the loaded encoder objects into it.
# The validation copy becomes the model definition of the forward
# config, so the merge must not alias the loaded encoder objects into it.
validation = SimpleNamespace(text_encoder_configs=(self._encoder(512), ))
original = validation.text_encoder_configs
loaded = SimpleNamespace(text_encoder_configs=(self._encoder(1472), ))
@@ -1110,3 +1086,77 @@ class TestKeepLoadedEncoderWidths:
) -> None:
# Pipelines without text encoders must not raise here.
ValidationCallback._keep_loaded_encoder_widths(validation, loaded)
class TestValidationResolvedConfigs:
"""The validation pipeline and its forwards run on resolved inference configs built from the training config."""
# A registered model path, so resolution finds its PipelineConfig class without a download.
_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
def _training_config(self) -> TrainingConfig:
return TrainingConfig(
distributed=DistributedConfig(num_gpus=2, sp_size=2, hsdp_shard_dim=2),
pipeline_config=PipelineConfig.from_kwargs({
"model_path": self._MODEL_PATH,
"flow_shift": 3.0
}),
model_path=self._MODEL_PATH,
vsa_sparsity=0.5,
)
def test_pipeline_is_built_from_resolved_inference_config(self, monkeypatch: pytest.MonkeyPatch) -> None:
"""Verify the parallel layout, DiT placement, flow shift, and YAML overrides of the validation pipeline."""
captured: dict = {}
class FakePipeline:
@classmethod
def from_pretrained(cls, model_path, *, resolved_config, loaded_modules):
captured.update(model_path=model_path, resolved_config=resolved_config, loaded_modules=loaded_modules)
return cls()
monkeypatch.setattr("fastvideo.train.callbacks.validation.resolve_target", lambda target: FakePipeline)
cb = ValidationCallback(
pipeline_target=_PIPE_TARGET,
dataset_file="x.json",
engine={"offload": {
"vae": True
}},
)
cb.training_config = self._training_config()
cb.method = SimpleNamespace()
transformer = torch.nn.Identity()
pipeline = cb._get_pipeline(transformer=transformer)
resolved_config = captured["resolved_config"]
assert isinstance(pipeline, FakePipeline)
assert captured["model_path"] == self._MODEL_PATH
assert captured["loaded_modules"] == {"transformer": transformer}
assert resolved_config.inference_mode
assert resolved_config.engine.num_gpus == 2
assert resolved_config.engine.parallelism.sp_size == 2
assert resolved_config.engine.offload.dit is False
assert resolved_config.engine.offload.dit_layerwise is False
assert resolved_config.engine.offload.vae is True
assert resolved_config.pipeline.flow_shift == 3.0
def test_forward_config_uses_validation_copy_and_sampling_timesteps(self) -> None:
"""Verify the forward config's model definition and DMD steps leave the training pipeline config unchanged."""
cb = _make_callback(sampling_timesteps=[1000, 757, 522])
training_config = self._training_config()
cb.training_config = training_config
loaded_config = PipelineConfig.from_kwargs({"model_path": self._MODEL_PATH})
loaded_config.text_encoder_configs[0].arch_config.hidden_size = 1234
pipeline = SimpleNamespace(resolved_config=SimpleNamespace(pipeline_config=loaded_config))
forward_config = cb._validation_forward_config(pipeline, torch.nn.Identity())
assert forward_config.inference_mode
assert forward_config.engine.offload.dit is True
assert forward_config.engine.attention.vsa_sparsity == 0.5
assert forward_config.pipeline.dmd_denoising_steps == (1000, 757, 522)
assert forward_config.pipeline_config.text_encoder_configs[0].arch_config.hidden_size == 1234
assert training_config.pipeline_config.dmd_denoising_steps is None
assert training_config.pipeline_config.text_encoder_configs[0].arch_config.hidden_size != 1234
@@ -247,8 +247,7 @@ def test_h3_experiment_config_uses_modular_validation_callback() -> None:
assert validation["num_videos_per_prompt"] == 1
assert validation["use_validation_media_conditioning"] is False
assert validation["offload_training_state"] is True
assert validation["text_encoder_cpu_offload"] is True
assert validation["vae_cpu_offload"] is True
assert validation["engine"]["offload"] == {"text_encoder": True, "vae": True}
assert tracker["trackers"] == ["wandb"]
assert tracker["project_name"] == "fastvideo_minimax_h3"
@@ -1,11 +1,12 @@
# SPDX-License-Identifier: Apache-2.0
"""Verify construction-scoped attention backends in the training loader."""
"""Verify construction-scoped attention backends and resolved configs in the training loader."""
from __future__ import annotations
import torch
import pytest
from fastvideo.api.schema import ExecutionMode
from fastvideo.attention.selector import _active_component_attention_backend_scope
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.platforms import AttentionBackendEnum
@@ -15,6 +16,9 @@ from fastvideo.train.utils.training_config import (
TrainingConfig,
)
# A registered model path, so resolution finds its PipelineConfig class without a download.
_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
def test_load_transformer_scopes_attention_backend(monkeypatch, tmp_path) -> None:
"""Apply one backend while accepting trailing modular-manifest metadata."""
@@ -46,7 +50,7 @@ def test_load_transformer_scopes_attention_backend(monkeypatch, tmp_path) -> Non
)
result = moduleloader.load_module_from_path(
model_path="fake/model",
model_path=_MODEL_PATH,
module_type="transformer",
training_config=training_config,
attention_backend="ATTN_QAT_TRAIN",
@@ -86,9 +90,100 @@ def test_load_transformer_restores_backend_when_loading_fails(
with pytest.raises(RuntimeError, match="load failed"):
moduleloader.load_module_from_path(
model_path="fake/model",
model_path=_MODEL_PATH,
module_type="transformer",
training_config=training_config,
attention_backend="ATTN_QAT_TRAIN",
)
assert _active_component_attention_backend_scope() is None
def _wan_training_config() -> TrainingConfig:
"""Training config with the Wan model definition and a two-GPU sequence-parallel layout."""
return TrainingConfig(
distributed=DistributedConfig(num_gpus=2, sp_size=2, hsdp_shard_dim=2),
pipeline_config=PipelineConfig.from_kwargs({"model_path": _MODEL_PATH}),
model_path=_MODEL_PATH,
vsa_sparsity=0.5,
)
def test_load_module_passes_distillation_config_and_teacher_flag(monkeypatch, tmp_path) -> None:
"""The loader receives the transformer overrides as typed values and the teacher/critic flag as a keyword."""
training_config = _wan_training_config()
captured: dict = {}
monkeypatch.setattr(moduleloader, "maybe_download_model", lambda path: str(tmp_path))
monkeypatch.setattr(
moduleloader,
"verify_model_config_and_directory",
lambda path: {"transformer": ("diffusers", "WanTransformer3DModel")},
)
def _fake_load_module(**kwargs):
captured.update(kwargs)
return torch.nn.Linear(1, 1)
monkeypatch.setattr(moduleloader.PipelineComponentLoader, "load_module", _fake_load_module)
moduleloader.load_module_from_path(
model_path=_MODEL_PATH,
module_type="transformer",
training_config=training_config,
disable_custom_init_weights=True,
override_transformer_cls_name="CausalWanTransformer3DModel",
transformer_override_safetensor="/weights/student.safetensors",
)
resolved_config = captured["resolved_config"]
assert captured["loading_teacher_critic_model"] is True
assert resolved_config.mode is ExecutionMode.DISTILLATION
assert resolved_config.training_mode
assert resolved_config.model_path == _MODEL_PATH
assert resolved_config.pipeline.components.override_transformer_cls_name == "CausalWanTransformer3DModel"
assert resolved_config.pipeline.components.transformer_weights == "/weights/student.safetensors"
assert resolved_config.engine.num_gpus == 2
assert resolved_config.engine.parallelism.sp_size == 2
assert resolved_config.engine.parallelism.hsdp_shard_dim == 2
assert resolved_config.engine.precision.dit == "fp32"
assert not any((resolved_config.engine.offload.dit, resolved_config.engine.offload.dit_layerwise,
resolved_config.engine.offload.text_encoder, resolved_config.engine.offload.vae))
assert type(resolved_config.pipeline_config) is type(training_config.pipeline_config)
assert resolved_config.pipeline_config is not training_config.pipeline_config
def test_vae_load_keeps_checkpoint_vae_config(monkeypatch, tmp_path) -> None:
"""The VAE config that the loader filled from the checkpoint becomes the training pipeline config's VAE config."""
training_config = _wan_training_config()
monkeypatch.setattr(moduleloader, "maybe_download_model", lambda path: str(tmp_path))
monkeypatch.setattr(
moduleloader,
"verify_model_config_and_directory",
lambda path: {"vae": ("diffusers", "AutoencoderKLWan")},
)
def _fill_vae_config(**kwargs):
kwargs["resolved_config"].pipeline_config.vae_config.arch_config.z_dim = 48
return torch.nn.Linear(1, 1)
monkeypatch.setattr(moduleloader.PipelineComponentLoader, "load_module", _fill_vae_config)
moduleloader.load_module_from_path(
model_path=_MODEL_PATH,
module_type="vae",
training_config=training_config,
)
assert training_config.pipeline_config.vae_config.arch_config.z_dim == 48
def test_inference_resolved_config_offloads_only_the_dit() -> None:
"""The inference config keeps encoders and VAE on the device and carries the training VSA sparsity."""
resolved_config = moduleloader.build_inference_resolved_config(_wan_training_config(), model_path=_MODEL_PATH)
assert resolved_config.mode is ExecutionMode.INFERENCE
assert resolved_config.inference_mode
assert resolved_config.engine.offload.dit is True
assert resolved_config.engine.offload.text_encoder is False
assert resolved_config.engine.offload.vae is False
assert resolved_config.engine.attention.vsa_sparsity == 0.5
assert resolved_config.engine.precision.dit == "fp32"
+116 -74
View File
@@ -13,6 +13,7 @@ import gc
import json
import os
import time
from collections.abc import Mapping
from copy import deepcopy
from dataclasses import dataclass, field
from typing import Any, TYPE_CHECKING
@@ -23,7 +24,11 @@ import torchvision
from einops import rearrange
from torch.utils.data import DataLoader
from fastvideo.api.inference_resolution import resolve_inference_config
from fastvideo.api.overrides import apply_overrides
from fastvideo.api.resolution import ResolvedGeneratorConfig
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.dataset.validation_dataset import (
ValidationDataset, )
from fastvideo.distributed import (
@@ -35,7 +40,7 @@ from fastvideo.pipelines import ForwardBatch
from fastvideo.train.callbacks.callback import Callback
from fastvideo.train.utils.instantiate import resolve_target
from fastvideo.train.utils.moduleloader import (
make_inference_args, )
build_inference_resolved_config, )
from fastvideo.train.utils.validation_media import write_validation_mp4
from fastvideo.training.trackers import DummyTracker
from fastvideo.utils import pixels_to_uint8, shallow_asdict
@@ -117,6 +122,18 @@ SYNTHETIC_OPTICAL_FLOW_LOG_KEYS = (
)
def _dotted_leaves(mapping: Mapping[str, Any], prefix: str = "") -> dict[str, Any]:
"""Return ``{dotted.path: value}`` for every value of a nested mapping that is not a non-empty mapping."""
leaves: dict[str, Any] = {}
for key, value in mapping.items():
path = f"{prefix}{key}"
if isinstance(value, Mapping) and value:
leaves.update(_dotted_leaves(value, f"{path}."))
else:
leaves[path] = value
return leaves
class ValidationCallback(Callback):
"""Generic validation callback driven entirely by YAML
config.
@@ -145,13 +162,15 @@ class ValidationCallback(Callback):
offload_training_state: bool = False,
unload_pipeline_after_validation: bool = False,
attn_qat_infer: bool = False,
**pipeline_kwargs: Any,
**generator_overrides: 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.
captions from a dataset that also contains source-media paths. The
remaining keys are nested ``GeneratorConfig`` fields, such as
``engine.offload.vae``, that override the validation pipeline's config.
"""
self.pipeline_target = str(pipeline_target)
self.dataset_file = str(dataset_file)
@@ -169,12 +188,12 @@ class ValidationCallback(Callback):
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)
metrics_config = generator_overrides.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)
self.generator_overrides = dict(generator_overrides)
# Set after on_train_start.
self._pipeline: Any | None = None
@@ -1191,13 +1210,27 @@ class ValidationCallback(Callback):
model_path = getattr(pipeline_config, "_fastvideo_train_model_path", None)
return str(model_path or tc.model_path)
def _validation_pipeline_config(self, transformer: torch.nn.Module) -> Any:
def _validation_pipeline_config(
self,
transformer: torch.nn.Module,
loaded_config: Any,
) -> Any:
"""Copy the training pipeline config for the validation forwards.
The copy takes the runtime attention window of ``transformer`` and the
encoder widths of ``loaded_config``, the validation pipeline's loaded
model definition.
"""
tc = self.training_config
pipeline_config = deepcopy(tc.pipeline_config)
pipeline_config = (deepcopy(tc.pipeline_config) if tc.pipeline_config is not None else PipelineConfig())
self._sync_runtime_dit_arch_config(
pipeline_config,
transformer,
)
self._keep_loaded_encoder_widths(
pipeline_config,
loaded_config,
)
return pipeline_config
@staticmethod
@@ -1207,12 +1240,11 @@ class ValidationCallback(Callback):
) -> None:
"""Carry the loader-populated encoder widths onto ``validation_config``.
Validation reaches the stages through two pipeline configs that never
pass through ``ModelConfig.update_model_arch``: the deep copy
``_validation_pipeline_config`` makes of the training-side config, and
``tc.pipeline_config`` itself, which ``make_inference_args`` hands to
``pipeline.forward`` by reference. Both still hold the encoder dataclass
defaults for whatever the checkpoint would have supplied.
Validation reaches the stages through a pipeline config that never
passes through ``ModelConfig.update_model_arch``: the deep copy
``_validation_pipeline_config`` makes of the training-side config. It
still holds the encoder dataclass defaults for whatever the checkpoint
would have supplied.
Stages read ``hidden_size`` only when they have to synthesise an
embedding instead of measuring one: HunyuanVideo 1.5 sizes its
@@ -1284,13 +1316,7 @@ class ValidationCallback(Callback):
if (self._pipeline is not None and self._pipeline_key == key):
return self._pipeline
tc = self.training_config
PipelineCls = resolve_target(self.pipeline_target)
flow_shift = getattr(
tc.pipeline_config,
"flow_shift",
None,
)
loaded_modules: dict[str, Any] = {"transformer": transformer}
# Distillation methods build the flow-match scheduler their few-step DMD
@@ -1299,45 +1325,84 @@ class ValidationCallback(Callback):
if method_scheduler is not None:
loaded_modules["scheduler"] = method_scheduler
kwargs: dict[str, Any] = {
"inference_mode": True,
"loaded_modules": loaded_modules,
"tp_size": tc.distributed.tp_size,
"sp_size": tc.distributed.sp_size,
"num_gpus": tc.distributed.num_gpus,
"pin_cpu_memory": (tc.distributed.pin_cpu_memory),
"dit_cpu_offload": False,
"dit_layerwise_offload": False,
}
if flow_shift is not None:
kwargs["flow_shift"] = float(flow_shift)
kwargs.update(self.pipeline_kwargs)
# The pipeline class comes from a YAML target, so static analysis cannot
# infer the dynamically resolved ``from_pretrained`` class method.
self._pipeline = PipelineCls.from_pretrained( # type: ignore[attr-defined]
self._pipeline_model_path(),
**kwargs,
resolved_config=resolve_inference_config(self._pipeline_generator_config()),
loaded_modules=loaded_modules,
)
if tc.pipeline_config is not None:
loaded_config = self._pipeline.resolved_config.pipeline_config
validation_config = self._validation_pipeline_config(transformer)
self._keep_loaded_encoder_widths(
validation_config,
loaded_config,
)
self._pipeline.resolved_config.pipeline_config = validation_config
arch_config = self._pipeline.resolved_config.pipeline_config.dit_config.arch_config
logger.info(
"Validation pipeline runtime config: local_attn_size=%s sink_size=%s boundary_ratio=%s",
getattr(arch_config, "local_attn_size", None),
getattr(arch_config, "sink_size", None),
getattr(self._pipeline.resolved_config.pipeline_config.dit_config, "boundary_ratio", None),
)
self._pipeline_key = key
return self._pipeline
def _pipeline_generator_config(self) -> dict[str, Any]:
"""Build the ``GeneratorConfig`` mapping of the validation pipeline.
It has the training parallel layout, keeps the DiT on the device, takes
``flow_shift`` from the training pipeline config, and applies
``generator_overrides`` last.
"""
tc = self.training_config
raw: dict[str, Any] = {
"model_path": self._pipeline_model_path(),
"engine": {
"num_gpus": tc.distributed.num_gpus,
"parallelism": {
"tp_size": tc.distributed.tp_size,
"sp_size": tc.distributed.sp_size,
},
"offload": {
"dit": False,
"dit_layerwise": False,
"pin_cpu_memory": tc.distributed.pin_cpu_memory,
},
},
}
flow_shift = getattr(
tc.pipeline_config,
"flow_shift",
None,
)
if flow_shift is not None:
raw["pipeline"] = {"flow_shift": float(flow_shift)}
return apply_overrides(raw, _dotted_leaves(self.generator_overrides))
def _validation_forward_config(
self,
pipeline: Any,
transformer: torch.nn.Module,
) -> ResolvedGeneratorConfig:
"""Build the resolved inference config that the validation forwards pass to the pipeline stages.
Its model definition is ``_validation_pipeline_config``. When that
definition has no DMD denoising steps, ``sampling_timesteps`` become
``pipeline.dmd_denoising_steps``, which the causal and DMD denoising
stages read.
"""
tc = self.training_config
forward_config = build_inference_resolved_config(
tc,
model_path=tc.model_path,
pipeline_config=self._validation_pipeline_config(
transformer,
pipeline.resolved_config.pipeline_config,
),
)
if (self.sampling_timesteps is not None and forward_config.pipeline.dmd_denoising_steps is None):
forward_config = forward_config.with_override(
"validation_callback:sampling_timesteps",
{"pipeline.dmd_denoising_steps": [int(s) for s in self.sampling_timesteps]},
)
dit_config = forward_config.pipeline_config.dit_config
logger.info(
"Validation forward config: local_attn_size=%s sink_size=%s boundary_ratio=%s",
getattr(dit_config.arch_config, "local_attn_size", None),
getattr(dit_config.arch_config, "sink_size", None),
getattr(dit_config, "boundary_ratio", None),
)
return forward_config
# ----------------------------------------------------------
# Batch preparation
# ----------------------------------------------------------
@@ -1393,11 +1458,6 @@ class ValidationCallback(Callback):
dtype=torch.long,
) if self.sampling_timesteps is not None else None)
inference_args = make_inference_args(
tc,
model_path=tc.model_path,
)
batch = ForwardBatch(
**shallow_asdict(sampling_param),
generator=self.validation_random_generator,
@@ -1415,7 +1475,6 @@ class ValidationCallback(Callback):
# mask instead of the current one, mismatching prompt_embeds.
batch.prompt_attention_mask = []
batch.negative_attention_mask = []
batch._inference_args = inference_args # type: ignore[attr-defined]
# Conditionally set I2V fields.
if ("image" in validation_batch and validation_batch["image"] is not None):
@@ -1479,7 +1538,6 @@ class ValidationCallback(Callback):
ranks in a group execute each forward pass; the group leader retains
decoded media.
"""
tc = self.training_config
pipeline = self._get_pipeline(transformer=transformer, )
sampling_param = self._get_sampling_param()
@@ -1490,23 +1548,7 @@ class ValidationCallback(Callback):
num_workers=0,
)
inference_args = make_inference_args(
tc,
model_path=tc.model_path,
)
self._sync_runtime_dit_arch_config(
inference_args.pipeline_config,
transformer,
)
self._keep_loaded_encoder_widths(
inference_args.pipeline_config,
pipeline.resolved_config.pipeline_config,
)
# Propagate sampling_timesteps to pipeline_config so
# causal/DMD denoising stages can read them.
if (self.sampling_timesteps is not None and inference_args.pipeline_config.dmd_denoising_steps is None):
inference_args.pipeline_config.dmd_denoising_steps = ([int(s) for s in self.sampling_timesteps])
forward_config = self._validation_forward_config(pipeline, transformer)
videos: list[list[np.ndarray]] = []
audio_waveforms: list[torch.Tensor | np.ndarray | None] = []
@@ -1532,7 +1574,7 @@ class ValidationCallback(Callback):
with torch.no_grad():
output_batch = pipeline.forward(
batch,
inference_args,
forward_config,
)
samples = output_batch.output.cpu()
@@ -56,8 +56,9 @@ from fastvideo.train.models.base import ModelBase
from fastvideo.train.utils.module_state import (
apply_trainable, )
from fastvideo.train.utils.moduleloader import (
build_inference_resolved_config,
keep_checkpoint_component_config,
load_module_from_path,
make_inference_args,
)
if TYPE_CHECKING:
@@ -444,8 +445,8 @@ class Kandinsky5Model(ModelBase):
pipeline_config = tc.pipeline_config
assert pipeline_config is not None
model_path = maybe_download_model(tc.model_path)
inference_args = make_inference_args(tc, model_path=model_path)
inference_args.text_encoder_cpu_offload = False
# The inference config keeps the text encoders on the device.
inference_config = build_inference_resolved_config(tc, model_path=model_path)
sampling_param = SamplingParam.from_pretrained(tc.model_path)
negative_prompt = sampling_param.negative_prompt
@@ -457,7 +458,7 @@ class Kandinsky5Model(ModelBase):
# --- Qwen / Reason1 ---
qwen_enc = loader.load(
os.path.join(model_path, "text_encoder"),
inference_args,
inference_config,
).to(device).eval()
qwen_tok = AutoTokenizer.from_pretrained(os.path.join(model_path, "tokenizer"))
qwen_tok_kwargs = dict(qwen_cfg.tokenizer_kwargs)
@@ -477,8 +478,9 @@ class Kandinsky5Model(ModelBase):
# --- CLIP ---
clip_enc = loader.load(
os.path.join(model_path, "text_encoder_2"),
inference_args,
inference_config,
).to(device).eval()
keep_checkpoint_component_config(tc, inference_config, "text_encoder_configs")
clip_tok = AutoTokenizer.from_pretrained(os.path.join(model_path, "tokenizer_2"))
clip_tok_kwargs = dict(clip_cfg.tokenizer_kwargs)
clip_text = preprocess_text(negative_prompt)
+1 -1
View File
@@ -23,7 +23,7 @@ def build_from_config(cfg: RunConfig, ) -> tuple[TrainingConfig, TrainingMethod,
1. Instantiate each model in ``cfg.models`` via ``_target_``.
2. Resolve the method class from ``cfg.method["_target_"]``
and construct it with ``(cfg=cfg, role_models=...)``.
3. Return ``(training_args, method, dataloader, start_step)``.
3. Return ``(training_config, method, dataloader, start_step)``.
"""
from fastvideo.train.models.base import ModelBase
+1 -1
View File
@@ -189,7 +189,7 @@ class CheckpointConfig:
class CheckpointManager:
"""Role-based checkpoint manager for training runtime.
- Checkpoint policy lives in YAML (via TrainingArgs fields).
- Checkpoint policy lives in YAML (the ``training.checkpoint`` section).
- Resume path is typically provided via CLI (``--resume-from-checkpoint``).
"""
+141 -51
View File
@@ -8,12 +8,22 @@ from typing import Any, TYPE_CHECKING
import torch
from fastvideo.api.inference_resolution import (
generator_resolution_steps,
resolve_config,
resolve_inference_config,
)
from fastvideo.api.resolution import ResolutionStep, ResolvedGeneratorConfig
from fastvideo.api.schema import ExecutionMode, GeneratorConfig
from fastvideo.api.training_schema import (
resolve_training_offload_conflicts,
validate_training_parallel_sizes,
)
from fastvideo.attention.selector import (
_component_attention_backend_scope,
coerce_attn_backend,
)
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.fastvideo_args import ExecutionMode, TrainingArgs
from fastvideo.models.loader.component_loader import (
PipelineComponentLoader, )
from fastvideo.utils import (
@@ -27,55 +37,137 @@ if TYPE_CHECKING:
TrainingConfig, )
# ------------------------------------------------------------------
# TrainingArgs builders (only place that creates FastVideoArgs)
# Resolved configs (the only place that builds them from a TrainingConfig)
# ------------------------------------------------------------------
def _make_training_args(
def _generator_config_mapping(
tc: TrainingConfig,
*,
model_path: str,
) -> TrainingArgs:
"""Build a TrainingArgs for PipelineComponentLoader."""
pipeline_config = tc.pipeline_config or PipelineConfig()
# Propagate dit_precision from TrainingConfig to PipelineConfig
# so that TransformerLoader.load() picks up the correct
# default_dtype (e.g. fp32 master weights for training).
if tc.dit_precision and tc.dit_precision != pipeline_config.dit_precision:
pipeline_config.dit_precision = tc.dit_precision
return TrainingArgs(
model_path=model_path,
mode=ExecutionMode.DISTILLATION,
inference_mode=False,
pipeline_config=pipeline_config,
num_gpus=tc.distributed.num_gpus,
tp_size=tc.distributed.tp_size,
sp_size=tc.distributed.sp_size,
hsdp_replicate_dim=tc.distributed.hsdp_replicate_dim,
hsdp_shard_dim=tc.distributed.hsdp_shard_dim,
pin_cpu_memory=tc.distributed.pin_cpu_memory,
dit_cpu_offload=False,
dit_layerwise_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=False,
image_encoder_cpu_offload=False,
use_fsdp_inference=False,
enable_torch_compile=False,
mode: ExecutionMode,
pipeline_config: PipelineConfig,
) -> dict[str, Any]:
"""``GeneratorConfig`` mapping with the training parallel layout and every offload and compile switch off.
``pipeline_config`` is the model definition that resolution copies and
materializes. ``tc.dit_precision`` sets the DiT precision, so
``TransformerLoader`` builds the master weights in it (fp32 for training).
"""
return {
"model_path": model_path,
"mode": mode,
"engine": {
"num_gpus": tc.distributed.num_gpus,
"parallelism": {
"tp_size": tc.distributed.tp_size,
"sp_size": tc.distributed.sp_size,
"hsdp_replicate_dim": tc.distributed.hsdp_replicate_dim,
"hsdp_shard_dim": tc.distributed.hsdp_shard_dim,
},
"offload": {
"dit": False,
"dit_layerwise": False,
"text_encoder": False,
"image_encoder": False,
"vae": False,
"pin_cpu_memory": tc.distributed.pin_cpu_memory,
},
"use_fsdp_inference": False,
"compile": {
"enabled": False
},
"precision": {
"dit": tc.dit_precision
},
},
"pipeline": {
"experimental": {
"pipeline_config": pipeline_config
}
},
}
def _training_resolution_steps(config: GeneratorConfig, defaults: Any) -> tuple[ResolutionStep, ...]:
"""Generator resolution steps plus the training offload-conflict and parallel-size checks."""
return generator_resolution_steps(
config,
defaults,
before_placeholders=(resolve_training_offload_conflicts, validate_training_parallel_sizes),
)
def make_inference_args(
def _build_training_resolved_config(
tc: TrainingConfig,
*,
model_path: str,
) -> TrainingArgs:
"""Build a TrainingArgs for inference (validation / pipelines)."""
args = _make_training_args(tc, model_path=model_path)
args.inference_mode = True
args.mode = ExecutionMode.INFERENCE
args.dit_cpu_offload = True
args.VSA_sparsity = tc.vsa_sparsity
return args
override_transformer_cls_name: str | None = None,
transformer_weights: str | None = None,
) -> ResolvedGeneratorConfig:
"""Build the distillation-mode resolved config that ``PipelineComponentLoader`` reads.
``override_transformer_cls_name`` and ``transformer_weights`` set
``pipeline.components.override_transformer_cls_name`` and
``pipeline.components.transformer_weights`` when they are given.
"""
raw = _generator_config_mapping(
tc,
model_path=model_path,
mode=ExecutionMode.DISTILLATION,
pipeline_config=tc.pipeline_config if tc.pipeline_config is not None else PipelineConfig(),
)
components: dict[str, str] = {}
if override_transformer_cls_name is not None:
components["override_transformer_cls_name"] = str(override_transformer_cls_name)
if transformer_weights:
components["transformer_weights"] = str(transformer_weights)
if components:
raw["pipeline"]["components"] = components
return resolve_config(raw, GeneratorConfig, _training_resolution_steps)
def build_inference_resolved_config(
tc: TrainingConfig,
*,
model_path: str,
pipeline_config: PipelineConfig | None = None,
) -> ResolvedGeneratorConfig:
"""Build the inference-mode resolved config for validation forwards and standalone encoder loads.
It has the training parallel layout, the DiT offloaded to the CPU, every
other offload off, and ``tc.vsa_sparsity`` as the VSA sparsity.
``pipeline_config`` replaces ``tc.pipeline_config`` as the model
definition.
"""
if pipeline_config is None:
pipeline_config = tc.pipeline_config if tc.pipeline_config is not None else PipelineConfig()
raw = _generator_config_mapping(
tc,
model_path=model_path,
mode=ExecutionMode.INFERENCE,
pipeline_config=pipeline_config,
)
raw["engine"]["offload"]["dit"] = True
raw["engine"]["attention"] = {"vsa_sparsity": tc.vsa_sparsity}
return resolve_inference_config(raw)
def keep_checkpoint_component_config(
tc: TrainingConfig,
resolved_config: ResolvedGeneratorConfig,
component_config_name: str,
) -> None:
"""Point ``tc.pipeline_config.<component_config_name>`` at the component config that a loader filled.
A resolved config holds a copy of ``tc.pipeline_config``, and the VAE and
encoder loaders fill that copy's arch configs from the checkpoint's
``config.json``. Training code reads component shapes, such as the VAE
latent channels and compression ratios, from ``tc.pipeline_config``.
"""
if tc.pipeline_config is not None:
setattr(tc.pipeline_config, component_config_name,
getattr(resolved_config.pipeline_config, component_config_name))
# ------------------------------------------------------------------
@@ -95,15 +187,20 @@ def load_module_from_path(
) -> torch.nn.Module:
"""Load one pipeline component with its role-scoped attention policy.
Accepts a ``TrainingConfig`` and internally builds the
``TrainingArgs`` needed by ``PipelineComponentLoader``.
Accepts a ``TrainingConfig`` and internally builds the distillation-mode
resolved config that ``PipelineComponentLoader`` reads.
Diffusers component entries retain provider and architecture as their
first two fields and can append modular loading metadata. Attention layers
bind their backend during construction, so the requested backend remains
scoped to this load call.
"""
resolved_config: Any = _make_training_args(training_config, model_path=model_path)
resolved_config = _build_training_resolved_config(
training_config,
model_path=model_path,
override_transformer_cls_name=override_transformer_cls_name,
transformer_weights=transformer_override_safetensor,
)
local_model_path = maybe_download_model(model_path)
config = verify_model_config_and_directory(local_model_path)
@@ -122,14 +219,6 @@ def load_module_from_path(
transformers_or_diffusers, _architecture = module_info[:2]
component_path = os.path.join(local_model_path, module_type)
# fastvideo_args is freshly built above and never escapes this function,
# so overrides are plain assignments — nothing to save or restore.
if override_transformer_cls_name is not None:
resolved_config.override_transformer_cls_name = str(override_transformer_cls_name)
if transformer_override_safetensor:
resolved_config.init_weights_from_safetensors = str(transformer_override_safetensor)
if attention_backend is not None and module_type != "transformer":
raise ValueError("attention_backend can only be set when loading "
f"a transformer, got module_type={module_type!r}")
@@ -140,8 +229,6 @@ def load_module_from_path(
attention_context = (nullcontext() if resolved_attention_backend is None else _component_attention_backend_scope(
resolved_attention_backend, component=module_type))
if disable_custom_init_weights:
resolved_config._loading_teacher_critic_model = True
# Attention implementations are bound while transformer layers are
# constructed. Scope the override to this one role so student,
# teacher, and critic can use independent backends in one process.
@@ -151,7 +238,10 @@ def load_module_from_path(
component_model_path=component_path,
transformers_or_diffusers=(transformers_or_diffusers),
resolved_config=resolved_config,
loading_teacher_critic_model=disable_custom_init_weights,
)
if module_type == "vae":
keep_checkpoint_component_config(training_config, resolved_config, "vae_config")
if not isinstance(module, torch.nn.Module):
raise TypeError(f"Loaded {module_type!r} is not a "
+9 -6
View File
@@ -19,7 +19,10 @@ from transformers import AutoTokenizer
from fastvideo.forward_context import set_forward_context
from fastvideo.models.loader.component_loader import TextEncoderLoader
from fastvideo.train.utils.moduleloader import make_inference_args
from fastvideo.train.utils.moduleloader import (
build_inference_resolved_config,
keep_checkpoint_component_config,
)
from fastvideo.utils import maybe_download_model
if TYPE_CHECKING:
@@ -65,16 +68,16 @@ def encode_negative_prompt(
tokenizer_subdir = f"tokenizer{suffix}"
model_path = maybe_download_model(tc.model_path)
inference_args = make_inference_args(tc, model_path=model_path)
# Keep the encoder on-device; CPU offload would init an FSDP device
# mesh and reintroduce the collective at load time.
inference_args.text_encoder_cpu_offload = False
# The inference config keeps the encoder on-device; CPU offload would init
# an FSDP device mesh and reintroduce the collective at load time.
inference_config = build_inference_resolved_config(tc, model_path=model_path)
loader = TextEncoderLoader()
text_encoder = loader.load(
os.path.join(model_path, encoder_subdir),
inference_args,
inference_config,
).to(device).eval()
keep_checkpoint_component_config(tc, inference_config, "text_encoder_configs")
tokenizer = AutoTokenizer.from_pretrained(os.path.join(model_path, tokenizer_subdir))
tok_kwargs = dict(encoder_config.tokenizer_kwargs)
+1 -1
View File
@@ -1,5 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
"""Typed training config — replaces TrainingArgs."""
"""Typed training config of the modular trainer."""
from __future__ import annotations