Compare commits
1
Commits
main
...
feat/h3-dmd2
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
dc70667483 |
Executable
+29
@@ -0,0 +1,29 @@
|
||||
#!/bin/bash
|
||||
# DMD2 distillation for MiniMax H3 (joint video + audio) on the modular
|
||||
# trainer (fastvideo/train). Unlike the legacy SFWan scripts, the modern
|
||||
# stack is YAML-driven: the {student, teacher, critic} trio, DMD2 method
|
||||
# knobs, and callbacks all live in the run config.
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export MASTER_PORT=${MASTER_PORT:-29513}
|
||||
# H3 training pins the dense TORCH_SDPA backend; do not export a sparse
|
||||
# attention backend here.
|
||||
|
||||
NUM_GPUS=${NUM_GPUS:-4}
|
||||
CONFIG=${CONFIG:-examples/train/configs/distribution_matching/minimax_h3/dmd2_t2va.yaml}
|
||||
DATA_DIR=${DATA_DIR:-data/crush-smol_h3_t2va_single_sample_preprocessed}
|
||||
OUTPUT_DIR=${OUTPUT_DIR:-outputs/minimax_h3_dmd2_3steps}
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--master_port "$MASTER_PORT" \
|
||||
--nproc_per_node "$NUM_GPUS" \
|
||||
-m fastvideo.train.entrypoint.train \
|
||||
--config "$CONFIG" \
|
||||
--training.distributed.num_gpus "$NUM_GPUS" \
|
||||
--training.distributed.sp_size "$NUM_GPUS" \
|
||||
--training.distributed.hsdp_shard_dim "$NUM_GPUS" \
|
||||
--training.data.data_path "$DATA_DIR" \
|
||||
--training.checkpoint.output_dir "$OUTPUT_DIR"
|
||||
@@ -0,0 +1,121 @@
|
||||
# DMD2 distillation: MiniMax H3 T2VA (teacher 50-step -> student 3-step).
|
||||
#
|
||||
# - Teacher: frozen pretrained MiniMax H3
|
||||
# - Student: trainable, initialized from the same pretrained weights
|
||||
# - Critic: trainable, initialized from the same pretrained weights
|
||||
# - Both modality streams (video + stereo audio) are distilled jointly through
|
||||
# the packed-latent adapter (MiniMaxH3DMDModel): one shared base timestep,
|
||||
# per-modality scheduler shifts (video 12.0, audio 3.0).
|
||||
# - Validation: standard MiniMaxH3Pipeline at 3 sampling steps (no dedicated
|
||||
# H3 DMD few-step pipeline yet).
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel
|
||||
init_from: MiniMaxAI/MiniMax-H3
|
||||
trainable: true
|
||||
enable_gradient_checkpointing_type: full
|
||||
attention_backend: TORCH_SDPA
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel
|
||||
init_from: MiniMaxAI/MiniMax-H3
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
attention_backend: TORCH_SDPA
|
||||
critic:
|
||||
_target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel
|
||||
init_from: MiniMaxAI/MiniMax-H3
|
||||
trainable: true
|
||||
disable_custom_init_weights: true
|
||||
enable_gradient_checkpointing_type: full
|
||||
attention_backend: TORCH_SDPA
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method
|
||||
# The preprocessed t2va parquet rows carry paired VAE latents, so perturb
|
||||
# those directly; switch to `simulate` for data-free rollout from noise.
|
||||
rollout_mode: data_latent
|
||||
generator_update_interval: 5
|
||||
real_score_guidance_scale: 3.5
|
||||
dmd_denoising_steps: [1000, 757, 522]
|
||||
min_timestep_ratio: 0.02
|
||||
max_timestep_ratio: 0.98
|
||||
# H3 has no negative-prompt encoder at training time; unconditional teacher
|
||||
# forwards zero the text embeddings instead.
|
||||
cfg_uncond:
|
||||
text: zero
|
||||
|
||||
# Critic optimizer (required — no fallback to training.optimizer)
|
||||
fake_score_learning_rate: 8.0e-6
|
||||
fake_score_betas: [0.0, 0.999]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 4
|
||||
sp_size: 4
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 4
|
||||
|
||||
data:
|
||||
data_path: data/crush-smol_h3_t2va_single_sample_preprocessed
|
||||
preprocessed_data_type: t2va
|
||||
dataloader_num_workers: 0
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 42
|
||||
num_latent_t: 37
|
||||
num_height: 768
|
||||
num_width: 1344
|
||||
num_frames: 124
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.0, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/minimax_h3_dmd2_3steps
|
||||
training_state_checkpointing_steps: 20
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_minimax_h3
|
||||
run_name: minimax_h3_dmd2_3steps
|
||||
|
||||
model:
|
||||
precondition_outputs: false
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: bf16
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline.MiniMaxH3Pipeline
|
||||
dataset_file: examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json
|
||||
every_steps: 50
|
||||
sampling_steps: [3]
|
||||
guidance_scale: 1.0
|
||||
num_frames: 124
|
||||
num_videos_per_prompt: 1
|
||||
use_validation_media_conditioning: false
|
||||
offload_training_state: true
|
||||
text_encoder_cpu_offload: true
|
||||
vae_cpu_offload: true
|
||||
ema:
|
||||
_target_: fastvideo.train.callbacks.ema.EMACallback
|
||||
decay: 0.98
|
||||
start_iter: 0
|
||||
|
||||
pipeline: {}
|
||||
@@ -0,0 +1,61 @@
|
||||
# Minimal H3 DMD2 trio and method configuration for CPU contract tests.
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel
|
||||
init_from: MiniMaxAI/MiniMax-H3
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel
|
||||
init_from: MiniMaxAI/MiniMax-H3
|
||||
trainable: false
|
||||
critic:
|
||||
_target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel
|
||||
init_from: MiniMaxAI/MiniMax-H3
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method
|
||||
rollout_mode: data_latent
|
||||
generator_update_interval: 1
|
||||
real_score_guidance_scale: 3.5
|
||||
dmd_denoising_steps: [1000, 757, 522]
|
||||
min_timestep_ratio: 0.02
|
||||
max_timestep_ratio: 0.98
|
||||
cfg_uncond:
|
||||
text: zero
|
||||
fake_score_learning_rate: 1.0e-3
|
||||
fake_score_betas: [0.0, 0.999]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 1
|
||||
data:
|
||||
data_path: /tmp/minimax_h3_t2va
|
||||
preprocessed_data_type: t2va
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 42
|
||||
# Tiny geometry: [1, 24, 2, 4, 4] video latents, [1, 2, 32, 8] audio.
|
||||
num_latent_t: 2
|
||||
num_height: 64
|
||||
num_width: 64
|
||||
num_frames: 5
|
||||
optimizer:
|
||||
learning_rate: 1.0e-3
|
||||
betas: [0.0, 0.999]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
loop:
|
||||
max_train_steps: 4
|
||||
model:
|
||||
precondition_outputs: false
|
||||
dit_precision: bf16
|
||||
|
||||
callbacks: {}
|
||||
pipeline: {}
|
||||
@@ -0,0 +1,335 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU contract tests for MiniMax H3 DMD2 distillation.
|
||||
|
||||
Covers the packed dual-modality adapter (MiniMaxH3DMDModel) and one full
|
||||
DMD2Method.single_train_step on a tiny CPU trio: student rollout, critic
|
||||
flow-matching loss, generator DMD loss, both backwards, both optimizers.
|
||||
"""
|
||||
|
||||
import math
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import yaml
|
||||
|
||||
from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method
|
||||
from fastvideo.train.models.minimax_h3 import MiniMaxH3DMDModel, MiniMaxH3Model
|
||||
from fastvideo.train.models.minimax_h3.minimax_h3 import shift_noise_amount
|
||||
from fastvideo.train.utils.config import load_run_config
|
||||
|
||||
_FIXTURE = Path(__file__).resolve().parent.parent / "fixtures" / "minimax_h3_dmd2_min.yaml"
|
||||
_REPO_ROOT = Path(__file__).resolve().parents[4]
|
||||
_EXPERIMENT_CONFIG = _REPO_ROOT / "examples/train/configs/distribution_matching/minimax_h3/dmd2_t2va.yaml"
|
||||
|
||||
# Fixture geometry: video latents [1, 24, 2, 4, 4] and audio latents
|
||||
# [1, 2, 32, 8]; the packed adapter stores video-major [1, T, C, H, W].
|
||||
_VIDEO_SHAPE = (1, 2, 24, 4, 4)
|
||||
_AUDIO_SHAPE = (1, 2, 32, 8)
|
||||
_PACKED_NUMEL = math.prod(_VIDEO_SHAPE) + math.prod(_AUDIO_SHAPE)
|
||||
|
||||
|
||||
class _TinyJointTransformer(torch.nn.Module):
|
||||
"""Scale packed H3 rows with one trainable parameter."""
|
||||
|
||||
patch_size = (1, 2, 2)
|
||||
|
||||
def __init__(self, scale: float = 1.0) -> None:
|
||||
super().__init__()
|
||||
self.scale = torch.nn.Parameter(torch.tensor(scale))
|
||||
self.last_encoder_hidden_states: torch.Tensor | None = None
|
||||
|
||||
def forward(self, **kwargs):
|
||||
self.last_encoder_hidden_states = kwargs["encoder_hidden_states"]
|
||||
return (
|
||||
kwargs["hidden_states"] * self.scale,
|
||||
kwargs["audio_hidden_states"] * self.scale,
|
||||
)
|
||||
|
||||
|
||||
def _make_model(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
training_config,
|
||||
*,
|
||||
trainable: bool = True,
|
||||
scale: float = 1.0,
|
||||
) -> MiniMaxH3DMDModel:
|
||||
monkeypatch.setattr(MiniMaxH3Model, "device", property(lambda _self: torch.device("cpu")))
|
||||
model = MiniMaxH3DMDModel.__new__(MiniMaxH3DMDModel)
|
||||
model._trainable = trainable
|
||||
model.transformer = _TinyJointTransformer(scale)
|
||||
model.training_config = training_config
|
||||
model.sp_group = None
|
||||
return model
|
||||
|
||||
|
||||
def _tiny_training_config():
|
||||
return SimpleNamespace(
|
||||
data=SimpleNamespace(
|
||||
num_latent_t=2,
|
||||
num_frames=5,
|
||||
num_height=64,
|
||||
num_width=64,
|
||||
),
|
||||
distributed=SimpleNamespace(sp_size=1),
|
||||
)
|
||||
|
||||
|
||||
def _raw_batch(seed: int = 1) -> dict[str, torch.Tensor]:
|
||||
generator = torch.Generator().manual_seed(seed)
|
||||
return {
|
||||
"vae_latent": torch.randn(1, 24, 2, 4, 4, generator=generator),
|
||||
"audio_latent": torch.randn(1, 2, 32, 8, generator=generator),
|
||||
"text_embedding": torch.randn(1, 4, 5120, generator=generator),
|
||||
"text_attention_mask": torch.tensor([[1, 1, 0, 0]], dtype=torch.float32),
|
||||
}
|
||||
|
||||
|
||||
def _build_method(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
*,
|
||||
rollout_mode: str,
|
||||
generator_update_interval: int = 1,
|
||||
) -> DMD2Method:
|
||||
config = load_run_config(str(_FIXTURE))
|
||||
config.method["rollout_mode"] = rollout_mode
|
||||
config.method["generator_update_interval"] = generator_update_interval
|
||||
# Distinct role scales keep the critic-vs-teacher DMD gradient non-zero.
|
||||
student = _make_model(monkeypatch, config.training, scale=1.0)
|
||||
teacher = _make_model(monkeypatch, config.training, trainable=False, scale=0.5)
|
||||
critic = _make_model(monkeypatch, config.training, scale=0.25)
|
||||
student.init_preprocessors = lambda training_config: None
|
||||
method = DMD2Method(
|
||||
cfg=config,
|
||||
role_models={
|
||||
"student": student,
|
||||
"teacher": teacher,
|
||||
"critic": critic,
|
||||
},
|
||||
)
|
||||
method.cuda_generator = torch.Generator(device="cpu").manual_seed(0)
|
||||
return method
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Core gate: one full DMD2 train step on CPU
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize("rollout_mode", ["data_latent", "simulate"])
|
||||
def test_dmd2_single_train_step_updates_student_and_critic(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
rollout_mode: str,
|
||||
) -> None:
|
||||
"""Run rollout, critic loss, generator loss, both backwards and steps."""
|
||||
method = _build_method(monkeypatch, rollout_mode=rollout_mode)
|
||||
student = method.student
|
||||
teacher = method.teacher
|
||||
critic = method.critic
|
||||
|
||||
loss_map, outputs, metrics = method.single_train_step(_raw_batch(), iteration=0)
|
||||
|
||||
assert metrics["update_student"] == 1.0
|
||||
for key in ("total_loss", "generator_loss", "fake_score_loss"):
|
||||
assert torch.isfinite(loss_map[key]), key
|
||||
assert loss_map["generator_loss"].item() > 0.0
|
||||
assert loss_map["fake_score_loss"].item() > 0.0
|
||||
|
||||
method.backward(loss_map, outputs)
|
||||
assert student.transformer.scale.grad is not None
|
||||
assert torch.isfinite(student.transformer.scale.grad)
|
||||
assert critic.transformer.scale.grad is not None
|
||||
assert torch.isfinite(critic.transformer.scale.grad)
|
||||
assert teacher.transformer.scale.grad is None
|
||||
|
||||
student_before = student.transformer.scale.detach().clone()
|
||||
critic_before = critic.transformer.scale.detach().clone()
|
||||
method.optimizers_schedulers_step(0)
|
||||
assert student.transformer.scale.detach() != student_before
|
||||
assert critic.transformer.scale.detach() != critic_before
|
||||
|
||||
|
||||
def test_dmd2_generator_update_interval_gates_student(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Off-interval iterations train the critic only."""
|
||||
method = _build_method(
|
||||
monkeypatch,
|
||||
rollout_mode="data_latent",
|
||||
generator_update_interval=5,
|
||||
)
|
||||
|
||||
loss_map, _outputs, metrics = method.single_train_step(_raw_batch(), iteration=1)
|
||||
|
||||
assert metrics["update_student"] == 0.0
|
||||
assert loss_map["generator_loss"].item() == 0.0
|
||||
assert loss_map["fake_score_loss"].item() > 0.0
|
||||
assert method.get_optimizers(1) == [method._critic_optimizer]
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Packed dual-modality adapter units
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_packed_adapter_roundtrip_and_prepare_batch(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Verify pack/unpack inversion and packed clean latents in the batch."""
|
||||
model = _make_model(monkeypatch, _tiny_training_config())
|
||||
video = torch.randn(_VIDEO_SHAPE)
|
||||
audio = torch.randn(_AUDIO_SHAPE)
|
||||
|
||||
packed = model.pack_latents(video, audio)
|
||||
assert packed.shape == (1, _PACKED_NUMEL)
|
||||
video_out, audio_out = model.unpack_latents(packed)
|
||||
torch.testing.assert_close(video_out, video)
|
||||
torch.testing.assert_close(audio_out, audio)
|
||||
|
||||
raw_batch = _raw_batch()
|
||||
batch = model.prepare_batch(
|
||||
raw_batch,
|
||||
generator=torch.Generator().manual_seed(7),
|
||||
)
|
||||
assert batch.latents.shape == (1, _PACKED_NUMEL)
|
||||
video_clean, audio_clean = model.unpack_latents(batch.latents)
|
||||
torch.testing.assert_close(
|
||||
video_clean,
|
||||
raw_batch["vae_latent"].permute(0, 2, 1, 3, 4).to(torch.bfloat16),
|
||||
)
|
||||
torch.testing.assert_close(audio_clean, batch.audio_latents)
|
||||
|
||||
|
||||
def test_packed_add_noise_applies_modality_shifts(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""One shared base timestep must map to two shifted noise amounts."""
|
||||
model = _make_model(monkeypatch, _tiny_training_config())
|
||||
clean = torch.ones(1, _PACKED_NUMEL)
|
||||
noise = torch.zeros(1, _PACKED_NUMEL)
|
||||
|
||||
torch.testing.assert_close(
|
||||
model.add_noise(clean, noise, torch.tensor([0])),
|
||||
clean,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
model.add_noise(clean, noise, torch.tensor([1000])),
|
||||
noise,
|
||||
)
|
||||
|
||||
mixed = model.add_noise(clean, noise, torch.tensor([500]))
|
||||
video_mixed, audio_mixed = model.unpack_latents(mixed)
|
||||
base = torch.tensor([0.5])
|
||||
torch.testing.assert_close(
|
||||
video_mixed,
|
||||
torch.full(_VIDEO_SHAPE, float(1.0 - shift_noise_amount(base, 12.0))),
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
audio_mixed,
|
||||
torch.full(_AUDIO_SHAPE, float(1.0 - shift_noise_amount(base, 3.0))),
|
||||
)
|
||||
|
||||
|
||||
def test_packed_predict_noise_plumbs_timesteps_and_tolerates_vsa(monkeypatch: pytest.MonkeyPatch, ) -> None:
|
||||
"""Explicit method timesteps must rewrite both modality clean-times."""
|
||||
model = _make_model(monkeypatch, _tiny_training_config())
|
||||
batch = model.prepare_batch(_raw_batch(), generator=torch.Generator().manual_seed(7))
|
||||
noisy = torch.randn(1, _PACKED_NUMEL).to(torch.bfloat16)
|
||||
timestep = torch.tensor([757], dtype=torch.long)
|
||||
|
||||
# attn_kind="vsa" must silently mean dense (both metadata views are None).
|
||||
prediction = model.predict_noise(
|
||||
noisy,
|
||||
timestep,
|
||||
batch,
|
||||
conditional=True,
|
||||
attn_kind="vsa",
|
||||
)
|
||||
|
||||
base = torch.tensor([0.757])
|
||||
torch.testing.assert_close(
|
||||
batch.timesteps,
|
||||
1.0 - shift_noise_amount(base, 12.0),
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
batch.audio_timesteps,
|
||||
1.0 - shift_noise_amount(base, 3.0),
|
||||
)
|
||||
# The unit-scale transformer echoes packed rows, and the H3 wrapper
|
||||
# negates them into noise-minus-clean form.
|
||||
torch.testing.assert_close(prediction, -noisy)
|
||||
|
||||
x0 = model.predict_x0(noisy, timestep, batch, conditional=True)
|
||||
noisy_video, noisy_audio = model.unpack_latents(noisy)
|
||||
sigma_video = shift_noise_amount(base, 12.0).to(torch.bfloat16)
|
||||
sigma_audio = shift_noise_amount(base, 3.0).to(torch.bfloat16)
|
||||
expected = model.pack_latents(
|
||||
noisy_video + sigma_video * noisy_video,
|
||||
noisy_audio + sigma_audio * noisy_audio,
|
||||
)
|
||||
torch.testing.assert_close(x0, expected)
|
||||
|
||||
|
||||
def test_uncond_forward_zeroes_text_and_guards_policies(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Teacher-CFG unconditional forwards zero text; other policies fail fast."""
|
||||
model = _make_model(monkeypatch, _tiny_training_config())
|
||||
batch = model.prepare_batch(_raw_batch(), generator=torch.Generator().manual_seed(7))
|
||||
noisy = torch.randn(1, _PACKED_NUMEL).to(torch.bfloat16)
|
||||
timestep = torch.tensor([500], dtype=torch.long)
|
||||
|
||||
model.predict_noise(
|
||||
noisy,
|
||||
timestep,
|
||||
batch,
|
||||
conditional=False,
|
||||
cfg_uncond={"text": "zero"},
|
||||
)
|
||||
assert torch.all(model.transformer.last_encoder_hidden_states == 0)
|
||||
|
||||
model.predict_noise(
|
||||
noisy,
|
||||
timestep,
|
||||
batch,
|
||||
conditional=True,
|
||||
cfg_uncond={"text": "zero"},
|
||||
)
|
||||
assert torch.any(model.transformer.last_encoder_hidden_states != 0)
|
||||
|
||||
with pytest.raises(ValueError, match="cfg_uncond"):
|
||||
model.predict_noise(noisy, timestep, batch, conditional=False)
|
||||
with pytest.raises(ValueError, match="negative-prompt"):
|
||||
model.set_requires_negative_conditioning(True)
|
||||
model.set_requires_negative_conditioning(False)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Config contracts
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_h3_dmd2_fixture_resolves_trio_contract() -> None:
|
||||
"""The fixture must wire the H3 DMD trio through the modular builder path."""
|
||||
config = load_run_config(str(_FIXTURE))
|
||||
|
||||
for role in ("student", "teacher", "critic"):
|
||||
assert config.models[role]["_target_"] == ("fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel")
|
||||
assert config.models["teacher"]["trainable"] is False
|
||||
assert config.method["_target_"] == ("fastvideo.train.methods.distribution_matching.dmd2.DMD2Method")
|
||||
assert config.training.data.preprocessed_data_type == "t2va"
|
||||
|
||||
|
||||
def test_h3_dmd2_experiment_config_mirrors_wan_recipe() -> None:
|
||||
"""The example config keeps the studio DMD2 defaults and the H3 contract."""
|
||||
config = yaml.safe_load(_EXPERIMENT_CONFIG.read_text())
|
||||
method = config["method"]
|
||||
|
||||
for role in ("student", "teacher", "critic"):
|
||||
assert config["models"][role]["_target_"] == ("fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel")
|
||||
assert config["models"]["teacher"]["trainable"] is False
|
||||
assert config["models"]["critic"]["trainable"] is True
|
||||
assert method["_target_"] == ("fastvideo.train.methods.distribution_matching.dmd2.DMD2Method")
|
||||
assert method["rollout_mode"] == "data_latent"
|
||||
assert method["generator_update_interval"] == 5
|
||||
assert method["dmd_denoising_steps"] == [1000, 757, 522]
|
||||
assert method["cfg_uncond"] == {"text": "zero"}
|
||||
assert method["fake_score_learning_rate"] == 8.0e-6
|
||||
assert method["fake_score_betas"] == [0.0, 0.999]
|
||||
assert method["fake_score_lr_scheduler"] == "constant"
|
||||
assert config["training"]["data"]["preprocessed_data_type"] == "t2va"
|
||||
assert config["training"]["data"]["train_batch_size"] == 1
|
||||
assert config["training"]["data"]["training_cfg_rate"] == 0.0
|
||||
@@ -3,3 +3,5 @@
|
||||
|
||||
from fastvideo.train.models.minimax_h3.minimax_h3 import (
|
||||
MiniMaxH3Model as MiniMaxH3Model, )
|
||||
from fastvideo.train.models.minimax_h3.minimax_h3_dmd import (
|
||||
MiniMaxH3DMDModel as MiniMaxH3DMDModel, )
|
||||
|
||||
@@ -320,10 +320,12 @@ class MiniMaxH3Model(ModelBase):
|
||||
) -> NoisePrediction:
|
||||
"""Pack modality timesteps and convert H3 outputs to noise-minus-clean."""
|
||||
del timestep
|
||||
if not conditional or cfg_uncond is not None:
|
||||
raise ValueError("MiniMaxH3Model predicts one conditional T2VA sample")
|
||||
if attn_kind != "dense":
|
||||
raise ValueError("MiniMaxH3Model supports dense attention for training")
|
||||
# Both attention-metadata views are None under TORCH_SDPA, so "vsa"
|
||||
# silently means dense here (mirrors WanModel under dense backends).
|
||||
# Real VSA-H3 metadata would slot in at this branch once the VSA-H3
|
||||
# attention backend lands (PR #1695).
|
||||
if attn_kind not in ("dense", "vsa"):
|
||||
raise ValueError(f"Unknown attn_kind: {attn_kind!r}")
|
||||
layout = batch.minimax_h3_layout
|
||||
if not isinstance(layout, MiniMaxH3PackedLayout):
|
||||
raise RuntimeError("prepare_batch() must set TrainingBatch.minimax_h3_layout")
|
||||
@@ -332,6 +334,16 @@ class MiniMaxH3Model(ModelBase):
|
||||
if batch.timesteps is None or batch.audio_timesteps is None:
|
||||
raise RuntimeError("prepare_batch() must set video and audio timesteps")
|
||||
|
||||
encoder_hidden_states = batch.encoder_hidden_states
|
||||
if not conditional:
|
||||
# H3 has no negative-prompt encoder at training time, so the only
|
||||
# supported unconditional branch (teacher CFG in distillation)
|
||||
# zeroes the text embeddings.
|
||||
if (cfg_uncond or {}).get("text") != "zero":
|
||||
raise ValueError("MiniMaxH3Model unconditional forwards require "
|
||||
"method.cfg_uncond={'text': 'zero'}")
|
||||
encoder_hidden_states = torch.zeros_like(encoder_hidden_states)
|
||||
|
||||
dtype = torch.bfloat16
|
||||
device = self.device
|
||||
video_bcthw = noisy_latents.permute(0, 2, 1, 3, 4).to(dtype)
|
||||
@@ -359,7 +371,7 @@ class MiniMaxH3Model(ModelBase):
|
||||
video_velocity, audio_velocity = self.transformer(
|
||||
hidden_states=video_rows[None],
|
||||
audio_hidden_states=audio_rows[None],
|
||||
encoder_hidden_states=batch.encoder_hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=unique_timesteps,
|
||||
timestep_indices=timestep_indices,
|
||||
token_tags=layout.token_tags.to(device),
|
||||
|
||||
@@ -0,0 +1,226 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""MiniMax H3 distribution-matching adapter (packed dual-modality latents)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import Any, Literal
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.pipelines import TrainingBatch
|
||||
from fastvideo.pipelines.basic.minimax_h3.packing import (
|
||||
MINIMAX_H3_AUDIO_CHANNELS,
|
||||
audio_latent_num_frames,
|
||||
)
|
||||
from fastvideo.train.models.minimax_h3.minimax_h3 import (
|
||||
_AUDIO_LATENT_CHANNELS,
|
||||
_AUDIO_SCHEDULER_SHIFT,
|
||||
_VIDEO_LATENT_CHANNELS,
|
||||
_VIDEO_SCHEDULER_SHIFT,
|
||||
MiniMaxH3Model,
|
||||
shift_noise_amount,
|
||||
)
|
||||
|
||||
# DMD2 samples integer score timesteps on the legacy [0, 1000] scale. H3 maps
|
||||
# them to its shared base noise amount before applying the modality shifts.
|
||||
_DMD_TIMESTEP_SCALE = 1000
|
||||
|
||||
|
||||
class MiniMaxH3DMDModel(MiniMaxH3Model):
|
||||
"""Present H3's dual (video, audio) streams to DMD2 as one packed tensor.
|
||||
|
||||
``DMD2Method``'s rollout and loss math assume one latent tensor per
|
||||
sample. This adapter flattens both modality latents into one ``[1, N]``
|
||||
tensor (video's ``[1, T, 24, H, W]`` elements first, stereo audio's
|
||||
``[1, 2, 32, Ta]`` elements after) so ``dmd2.py`` stays model-agnostic.
|
||||
Integer method timesteps become one shared base noise amount that is
|
||||
shifted per modality (video 12.0, audio 3.0), exactly as H3's paired
|
||||
schedulers synchronize the two streams during fine-tuning and inference.
|
||||
|
||||
ponytail: packed means weight modalities by element count (video
|
||||
dominates); switch to per-modality mean losses if audio quality lags.
|
||||
"""
|
||||
|
||||
@property
|
||||
def num_train_timesteps(self) -> int:
|
||||
return _DMD_TIMESTEP_SCALE
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Packed dual-modality helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _modality_shapes(self) -> tuple[tuple[int, ...], tuple[int, ...]]:
|
||||
"""Return the ``[1, T, C, H, W]`` video and ``[1, 2, 32, Ta]`` audio shapes."""
|
||||
data = self.training_config.data
|
||||
video_shape = (
|
||||
1,
|
||||
int(data.num_latent_t),
|
||||
_VIDEO_LATENT_CHANNELS,
|
||||
int(data.num_height) // 16,
|
||||
int(data.num_width) // 16,
|
||||
)
|
||||
audio_shape = (
|
||||
1,
|
||||
MINIMAX_H3_AUDIO_CHANNELS,
|
||||
_AUDIO_LATENT_CHANNELS,
|
||||
audio_latent_num_frames(int(data.num_frames)),
|
||||
)
|
||||
return video_shape, audio_shape
|
||||
|
||||
def pack_latents(
|
||||
self,
|
||||
video_latents: torch.Tensor,
|
||||
audio_latents: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Flatten both modality latents into one ``[1, N]`` tensor."""
|
||||
return torch.cat(
|
||||
(video_latents.reshape(1, -1), audio_latents.reshape(1, -1)),
|
||||
dim=1,
|
||||
)
|
||||
|
||||
def unpack_latents(self, packed: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Split one packed ``[1, N]`` tensor back into (video, audio) latents."""
|
||||
video_shape, audio_shape = self._modality_shapes()
|
||||
split = math.prod(video_shape)
|
||||
if packed.shape != (1, split + math.prod(audio_shape)):
|
||||
raise ValueError("Packed latents must have shape "
|
||||
f"[1, {split + math.prod(audio_shape)}], got {tuple(packed.shape)}")
|
||||
return (
|
||||
packed[:, :split].reshape(video_shape),
|
||||
packed[:, split:].reshape(audio_shape),
|
||||
)
|
||||
|
||||
def _noise_amounts(self, timestep: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Map one integer method timestep to both modality noise amounts."""
|
||||
base = (timestep.reshape(-1)[:1].to(torch.float32) / _DMD_TIMESTEP_SCALE)
|
||||
base = base.clamp(0.0, 1.0)
|
||||
return (
|
||||
shift_noise_amount(base, _VIDEO_SCHEDULER_SHIFT),
|
||||
shift_noise_amount(base, _AUDIO_SCHEDULER_SHIFT),
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# ModelBase overrides (packed convention)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def set_requires_negative_conditioning(self, requires: bool) -> None:
|
||||
"""Fail fast: H3 cannot encode negative prompts at training time."""
|
||||
if requires:
|
||||
raise ValueError("MiniMaxH3DMDModel has no negative-prompt encoder; set "
|
||||
"method.cfg_uncond={'text': 'zero'} for unconditional forwards")
|
||||
|
||||
def prepare_batch(
|
||||
self,
|
||||
raw_batch: dict[str, Any],
|
||||
*,
|
||||
generator: torch.Generator,
|
||||
latents_source: Literal["data", "zeros"] = "data",
|
||||
) -> TrainingBatch:
|
||||
"""Prepare the T2VA batch, then expose clean latents in packed form."""
|
||||
batch = super().prepare_batch(
|
||||
raw_batch,
|
||||
generator=generator,
|
||||
latents_source=latents_source,
|
||||
)
|
||||
# DMD2 draws its own noise and timesteps per forward; only the packed
|
||||
# clean latents matter here. The fine-tuning noisy fields are
|
||||
# refreshed by predict_noise on every call.
|
||||
batch.latents = self.pack_latents(batch.latents, batch.audio_latents)
|
||||
return batch
|
||||
|
||||
def add_noise(
|
||||
self,
|
||||
clean_latents: torch.Tensor,
|
||||
noise: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Noise packed latents at one shared timestep, shifted per modality."""
|
||||
sigma_video, sigma_audio = self._noise_amounts(timestep)
|
||||
clean_video, clean_audio = self.unpack_latents(clean_latents)
|
||||
noise_video, noise_audio = self.unpack_latents(noise)
|
||||
return self.pack_latents(
|
||||
self._mix(clean_video, noise_video, sigma_video),
|
||||
self._mix(clean_audio, noise_audio, sigma_audio),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _mix(
|
||||
clean: torch.Tensor,
|
||||
noise: torch.Tensor,
|
||||
sigma: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
sigma = sigma.to(device=clean.device, dtype=clean.dtype)
|
||||
return (1.0 - sigma) * clean + sigma * noise
|
||||
|
||||
def predict_noise(
|
||||
self,
|
||||
noisy_latents: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
batch: TrainingBatch,
|
||||
*,
|
||||
conditional: bool,
|
||||
cfg_uncond: dict[str, Any] | None = None,
|
||||
attn_kind: Literal["dense", "vsa"] = "dense",
|
||||
) -> torch.Tensor:
|
||||
"""Run one packed joint forward at an explicit method timestep.
|
||||
|
||||
Both modality clean-time fields on ``batch`` are rewritten from
|
||||
``timestep`` so the packed-row timestep plan and the backward
|
||||
forward-context stay coherent with this call.
|
||||
"""
|
||||
sigma_video, sigma_audio = self._noise_amounts(timestep)
|
||||
noisy_video, noisy_audio = self.unpack_latents(noisy_latents)
|
||||
batch.timesteps = (1.0 - sigma_video).to(noisy_latents.device)
|
||||
batch.audio_timesteps = (1.0 - sigma_audio).to(noisy_latents.device)
|
||||
batch.audio_noisy_model_input = noisy_audio
|
||||
video_pred, audio_pred = super().predict_noise(
|
||||
noisy_video,
|
||||
timestep,
|
||||
batch,
|
||||
conditional=conditional,
|
||||
cfg_uncond=cfg_uncond,
|
||||
attn_kind=attn_kind,
|
||||
)
|
||||
return self.pack_latents(video_pred, audio_pred)
|
||||
|
||||
def predict_x0(
|
||||
self,
|
||||
noisy_latents: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
batch: TrainingBatch,
|
||||
*,
|
||||
conditional: bool,
|
||||
cfg_uncond: dict[str, Any] | None = None,
|
||||
attn_kind: Literal["dense", "vsa"] = "dense",
|
||||
) -> torch.Tensor:
|
||||
"""Convert packed noise-minus-clean predictions to packed clean latents."""
|
||||
pred_noise = self.predict_noise(
|
||||
noisy_latents,
|
||||
timestep,
|
||||
batch,
|
||||
conditional=conditional,
|
||||
cfg_uncond=cfg_uncond,
|
||||
attn_kind=attn_kind,
|
||||
)
|
||||
sigma_video, sigma_audio = self._noise_amounts(timestep)
|
||||
noisy_video, noisy_audio = self.unpack_latents(noisy_latents)
|
||||
pred_video, pred_audio = self.unpack_latents(pred_noise)
|
||||
return self.pack_latents(
|
||||
self._to_x0(noisy_video, pred_video, sigma_video),
|
||||
self._to_x0(noisy_audio, pred_audio, sigma_audio),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _to_x0(
|
||||
noisy: torch.Tensor,
|
||||
pred_noise: torch.Tensor,
|
||||
sigma: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
# noisy = (1 - sigma) * clean + sigma * noise and pred approximates
|
||||
# noise - clean, so clean = noisy - sigma * pred.
|
||||
sigma = sigma.to(device=noisy.device, dtype=noisy.dtype)
|
||||
return noisy - sigma * pred_noise
|
||||
|
||||
|
||||
__all__ = ["MiniMaxH3DMDModel"]
|
||||
Reference in New Issue
Block a user