[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:
co-authored by
Claude Opus 5.5
parent
4117696a4f
commit
00a8e5bcb6
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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``).
|
||||
"""
|
||||
|
||||
|
||||
@@ -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 "
|
||||
|
||||
@@ -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,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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user