Compare commits

...
Author SHA1 Message Date
SolitaryThinker 4d8713f6e4 [bugfix]: bound WanCausalModel num_frames_per_block override to <= 3
The constructor override only checked >= 1, bypassing the <= 3 limit the
config path enforces in CausalWanTransformer3DModel. Apply the same upper
bound with a clear error message.
2026-07-05 15:01:25 -07:00
SolitaryThinker fda43fc610 [bugfix]: refuse teacher forcing with a local attention window
_prepare_teacher_forcing_mask silently ignored local_attn_size while the
block-wise causal mask honors it, so a configured attention window was
dropped on the teacher-forcing path. Raise NotImplementedError instead of
training with a mask that contradicts the config.
2026-07-05 15:01:09 -07:00
SolitaryThinker 5a8329ddf6 [bugfix]: set causal-CD timesteps before the teacher CFG forwards
training_batch.timesteps was assigned t_pf only after the two teacher
CFG passes, so their set_forward_context(current_timestep=...) carried
the stale random timesteps from prepare_batch — wrong VSA sparsity
gating when attn_kind == "vsa". Assign before any forward.
2026-07-05 15:00:47 -07:00
SolitaryThinker fcb5b465c5 [bugfix]: checkpoint causal-CD EMA consistency target on save/resume
The base TrainingMethod.checkpoint_state() only persists trainable roles,
so CausalConsistencyDistillationMethod's EMA target (frozen but mutated by
_update_ema every step) was never saved; a resume reloaded it from
init_from, snapping the consistency target back to the base checkpoint.

Persist roles.ema.transformer via the same full-state DCP wrapper
DiffusionNFT already uses for its frozen 'old' role, moving _FullModelState
into fastvideo/train/utils/checkpoint.py so both methods share it.
2026-07-05 15:00:27 -07:00
H1yori233 4f13fb0fee fix 2026-07-05 14:56:08 -07:00
H1yori233 d96ec99b4b fix 2026-07-05 14:56:08 -07:00
H1yori233 2555f25cce fix scheduler 2026-07-05 14:56:08 -07:00
H1yori233 fbe56fce8d update scheduler and ema 2026-07-05 14:56:08 -07:00
H1yori233 5064bcdc47 cleanup 2026-07-05 14:56:08 -07:00
H1yori233 2aa824b615 cleanup 2026-07-05 14:56:08 -07:00
H1yori233 f934efc58d add cf 2026-07-05 14:56:08 -07:00
21 changed files with 1237 additions and 57 deletions
@@ -0,0 +1,78 @@
# Causal Consistency Distillation: Wan 2.1 T2V 1.3B Causal
#
# ODE-data-free distillation. A frozen teacher takes a single CFG Euler step;
# the student matches an EMA copy of itself at the next timestep, all under
# clean-history teacher forcing.
#
# All three roles initialize from the SAME checkpoint (the teacher-forcing
# AR-diffusion model). Point init_from at that checkpoint for a real run.
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
teacher:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
ema:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
method:
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
discrete_cd_N: 48
guidance_scale: 3.0
ema_decay: 0.99
ema_start_step: 200
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/Wan-Syn_77x448x832_600k
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 18
num_height: 448
num_width: 832
num_frames: 69
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 3000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wan2.1_causal_cd
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
project_name: distillation_wan_r
run_name: wan2.1_causal_cd_shift5
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
pipeline:
flow_shift: 5
@@ -0,0 +1,73 @@
# DFSFT (Diffusion-Forcing SFT), frame-wise: Wan 2.1 T2V 1.3B Causal
#
# - Student: trainable causal Wan model with a block size of 1 frame
# - Training: each frame gets its own independent noise level (frame-wise
# diffusion forcing), versus the chunk-wise variant that shares one noise
# level across num_frames_per_block frames.
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
num_frames_per_block: 1
method:
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
chunk_size: 1
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/Wan-Syn_77x448x832_600k
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 18
num_height: 448
num_width: 832
num_frames: 69
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 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/wan2.1_causal_dfsft_framewise
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
project_name: distillation_wan_r
run_name: wan2.1_causal_dfsft_framewise_shift5_gauss_weight
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 50
sampling_steps: [40]
guidance_scale: 6.0
num_frames: 69
pipeline:
flow_shift: 5
@@ -0,0 +1,72 @@
# TFSFT (Teacher-Forcing SFT): Wan 2.1 T2V 1.3B Causal
#
# - Student: trainable causal Wan model
# - Training: inhomogeneous timesteps per chunk, but the causal transformer
# denoises the current block while attending to *clean* history (clean_x),
# not its own noisy rollout.
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
chunk_size: 3
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/Wan-Syn_77x448x832_600k
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 18
num_height: 448
num_width: 832
num_frames: 69
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 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/wan2.1_causal_tfsft
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
project_name: distillation_wan_r
run_name: wan2.1_causal_tfsft_shift5_gauss_weight
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 50
sampling_steps: [40]
guidance_scale: 6.0
num_frames: 69
pipeline:
flow_shift: 5
+108 -11
View File
@@ -437,6 +437,7 @@ class CausalWanTransformer3DModel(BaseDiT):
# Causal-specific
self.block_mask = None
self.teacher_forcing_block_mask = None
self.num_frame_per_block = config.arch_config.num_frames_per_block
assert self.num_frame_per_block <= 3
self.independent_first_frame = False
@@ -500,6 +501,70 @@ class CausalWanTransformer3DModel(BaseDiT):
return block_mask
@staticmethod
def _prepare_teacher_forcing_mask(
device: torch.device | str, num_frames: int = 21,
frame_seqlen: int = 1560, num_frame_per_block=1, local_attn_size=-1
) -> BlockMask:
"""Attention mask for the teacher-forcing ``[clean | noisy]`` sequence.
A noisy token attends to its own block plus the clean context of all
strictly previous blocks; clean tokens are block-wise causal.
"""
if local_attn_size != -1:
raise NotImplementedError(
f"Teacher forcing ignores local_attn_size={local_attn_size}: "
"unlike the block-wise causal mask, this mask always attends "
"to the full clean context. Windowed teacher forcing is not "
"implemented; use local_attn_size=-1 for teacher-forcing "
"training.")
total_length = num_frames * frame_seqlen * 2
padded_length = math.ceil(total_length / 128) * 128 - total_length
clean_ends = num_frames * frame_seqlen
context_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
noise_context_starts = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
noise_context_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
noise_noise_starts = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
noise_noise_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
attention_block_size = frame_seqlen * num_frame_per_block
frame_indices = torch.arange(
start=0, end=num_frames * frame_seqlen,
step=attention_block_size, device=device, dtype=torch.long
)
for start in frame_indices:
context_ends[start:start + attention_block_size] = start + attention_block_size
noisy_image_start_list = torch.arange(
num_frames * frame_seqlen, total_length,
step=attention_block_size, device=device, dtype=torch.long
)
noisy_image_end_list = noisy_image_start_list + attention_block_size
for block_index, (start, end) in enumerate(zip(noisy_image_start_list, noisy_image_end_list)):
noise_noise_starts[start:end] = start
noise_noise_ends[start:end] = end
noise_context_ends[start:end] = block_index * attention_block_size
def attention_mask(b, h, q_idx, kv_idx):
clean_mask = (q_idx < clean_ends) & (kv_idx < context_ends[q_idx])
c1 = (kv_idx < noise_noise_ends[q_idx]) & (kv_idx >= noise_noise_starts[q_idx])
c2 = (kv_idx < noise_context_ends[q_idx]) & (kv_idx >= noise_context_starts[q_idx])
noise_mask = (q_idx >= clean_ends) & (c1 | c2)
eye_mask = q_idx == kv_idx
return eye_mask | clean_mask | noise_mask
block_mask = create_block_mask(
attention_mask, B=None, H=None,
Q_LEN=total_length + padded_length, KV_LEN=total_length + padded_length,
_compile=False, device=device)
if not dist.is_initialized() or dist.get_rank() == 0:
print(f" cache a teacher-forcing mask with block size of {num_frame_per_block} frames")
print(block_mask)
return block_mask
def _forward_inference(
self,
hidden_states: torch.Tensor,
@@ -628,9 +693,12 @@ class CausalWanTransformer3DModel(BaseDiT):
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
start_frame: int = 0,
clean_x: torch.Tensor | None = None,
aug_t: torch.Tensor | None = None,
**kwargs) -> torch.Tensor:
orig_dtype = hidden_states.dtype
teacher_forcing = clean_x is not None
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
@@ -663,15 +731,26 @@ class CausalWanTransformer3DModel(BaseDiT):
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
# Construct blockwise causal attn mask
if self.block_mask is None:
self.block_mask = self._prepare_blockwise_causal_attn_mask(
device=hidden_states.device,
num_frames=num_frames,
frame_seqlen=post_patch_height * post_patch_width,
num_frame_per_block=self.num_frame_per_block,
local_attn_size=self.local_attn_size
)
if teacher_forcing:
if self.teacher_forcing_block_mask is None:
self.teacher_forcing_block_mask = self._prepare_teacher_forcing_mask(
device=hidden_states.device,
num_frames=num_frames,
frame_seqlen=post_patch_height * post_patch_width,
num_frame_per_block=self.num_frame_per_block,
local_attn_size=self.local_attn_size,
)
block_mask = self.teacher_forcing_block_mask
else:
if self.block_mask is None:
self.block_mask = self._prepare_blockwise_causal_attn_mask(
device=hidden_states.device,
num_frames=num_frames,
frame_seqlen=post_patch_height * post_patch_width,
num_frame_per_block=self.num_frame_per_block,
local_attn_size=self.local_attn_size
)
block_mask = self.block_mask
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
@@ -679,6 +758,7 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = hidden_states.flatten(2).transpose(1, 2)
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
encoder_hidden_states_text = encoder_hidden_states
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
@@ -694,18 +774,35 @@ class CausalWanTransformer3DModel(BaseDiT):
assert encoder_hidden_states.dtype == orig_dtype
if teacher_forcing:
# Tile RoPE/modulation so clean frame i and noisy frame i share a position.
clean_tokens = self.patch_embedding(clean_x).flatten(2).transpose(1, 2)
hidden_states = torch.cat([clean_tokens, hidden_states], dim=1)
if aug_t is None:
aug_t = torch.zeros_like(timestep)
_, timestep_proj_clean, _, _ = self.condition_embedder(
aug_t.flatten(), encoder_hidden_states_text, None)
timestep_proj_clean = timestep_proj_clean.unflatten(
1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
timestep_proj = torch.cat([timestep_proj_clean, timestep_proj], dim=1)
freqs_cis = (torch.cat([freqs_cos, freqs_cos], dim=0),
torch.cat([freqs_sin, freqs_sin], dim=0))
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.blocks:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis,
block_mask=self.block_mask)
block_mask=block_mask)
else:
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis,
block_mask=self.block_mask)
block_mask=block_mask)
if teacher_forcing:
hidden_states = hidden_states[:, hidden_states.shape[1] // 2:]
# 5. Output norm, projection & unpatchify
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
@@ -0,0 +1,56 @@
# SPDX-License-Identifier: Apache-2.0
# Minimum config to run a single training step of
# CausalConsistencyDistillationMethod on WanCausalModel for the
# per-method smoke test. Uses the real Wan 2.1 1.3B checkpoint for both
# the trainable student and the frozen teacher (AR Euler-step target),
# with tiny synthetic latents.
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
teacher:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
ema:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
method:
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
discrete_cd_N: 12
guidance_scale: 3.0
ema_decay: 0.95
training:
dit_precision: bf16
distributed:
num_gpus: 1
sp_size: 1
tp_size: 1
data:
seed: 42
train_batch_size: 1
training_cfg_rate: 0.0
num_latent_t: 6
num_height: 64
num_width: 64
num_frames: 21
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 1
gradient_accumulation_steps: 1
pipeline: {}
@@ -0,0 +1,46 @@
# Minimum config to run a single training step of
# DiffusionForcingSFTMethod on a frame-wise WanCausalModel
# (num_frames_per_block=1, chunk_size=1) for the per-method smoke
# test. Uses the real Wan 2.1 1.3B checkpoint with tiny synthetic
# latents.
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
num_frames_per_block: 1
method:
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
chunk_size: 1
training:
dit_precision: bf16
distributed:
num_gpus: 1
sp_size: 1
tp_size: 1
data:
seed: 42
train_batch_size: 1
training_cfg_rate: 0.0
num_latent_t: 6
num_height: 64
num_width: 64
num_frames: 21
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 1
gradient_accumulation_steps: 1
pipeline: {}
@@ -0,0 +1,45 @@
# SPDX-License-Identifier: Apache-2.0
# Minimum config to run a single training step of
# TeacherForcingSFTMethod on WanCausalModel for the per-method smoke
# test. Identical to wan_causal_t2v_dfsft_min.yaml except the method,
# which feeds clean history to the causal transformer (teacher forcing).
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
chunk_size: 3
training:
dit_precision: bf16
distributed:
num_gpus: 1
sp_size: 1
tp_size: 1
data:
seed: 42
train_batch_size: 1
training_cfg_rate: 0.0
num_latent_t: 6
num_height: 64
num_width: 64
num_frames: 21
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 1
gradient_accumulation_steps: 1
pipeline: {}
@@ -0,0 +1,147 @@
# SPDX-License-Identifier: Apache-2.0
"""Per-method GPU smoke test: ``WanCausalModel`` + ``CausalConsistencyDistillationMethod``.
Mirrors ``test_wan_causal_dfsft.py``. Causal consistency distillation
bootstraps a consistency MSE between the student's ``x0`` at ``t`` and an EMA
copy of the student at ``t_next``, where ``t_next`` is produced online by a
single CFG Euler step of a frozen teacher (all under clean-history teacher
forcing). This test exercises the full step: finite loss, nonzero student
gradients, frozen teacher, and a post-step EMA update.
"""
from __future__ import annotations
import os
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29519")
from pathlib import Path
import pytest
import torch
from fastvideo.train.methods.consistency_model.causal_cd import (
CausalConsistencyDistillationMethod, )
from fastvideo.train.models.wan import WanCausalModel
from fastvideo.train.utils.config import load_run_config
_FIXTURE = str(
Path(__file__).resolve().parent.parent / "fixtures"
/ "wan_causal_t2v_causal_cd_min.yaml")
def _build_synthetic_batch(
device: torch.device,
dtype: torch.dtype,
) -> dict[str, torch.Tensor]:
batch_size = 1
return {
"text_embedding":
torch.randn(batch_size, 16, 4096, device=device, dtype=dtype),
"text_attention_mask":
torch.ones(batch_size, 16, device=device),
"vae_latent":
torch.randn(batch_size, 16, 6, 8, 8, device=device, dtype=dtype),
}
@pytest.mark.usefixtures("distributed_setup")
def test_wan_causal_cd_single_train_step(
monkeypatch: pytest.MonkeyPatch) -> None:
if not torch.cuda.is_available():
pytest.skip("requires CUDA")
cfg = load_run_config(_FIXTURE)
device = torch.device("cuda:0")
dtype = torch.bfloat16
monkeypatch.setattr(
"fastvideo.train.utils.dataloader."
"build_parquet_t2v_train_dataloader",
lambda *args, **kwargs: None,
)
student = WanCausalModel(
init_from=cfg.models["student"]["init_from"],
training_config=cfg.training,
trainable=True,
)
student.transformer = student.transformer.to(device=device, dtype=dtype)
teacher = WanCausalModel(
init_from=cfg.models["teacher"]["init_from"],
training_config=cfg.training,
trainable=False,
)
teacher.transformer = teacher.transformer.to(device=device, dtype=dtype)
ema = WanCausalModel(
init_from=cfg.models["ema"]["init_from"],
training_config=cfg.training,
trainable=False,
)
ema.transformer = ema.transformer.to(device=device, dtype=dtype)
method = CausalConsistencyDistillationMethod(
cfg=cfg,
role_models={"student": student, "teacher": teacher, "ema": ema},
)
method.on_train_start()
batch = _build_synthetic_batch(device, dtype)
loss_map, outputs, _metrics = method.single_train_step(batch, iteration=0)
loss = loss_map["total_loss"]
assert torch.is_tensor(loss), "total_loss must be a torch.Tensor"
assert torch.isfinite(loss).item(), (
f"total_loss is not finite: {loss.item()}")
method.backward(loss_map, outputs, grad_accum_rounds=1)
blocks = student.transformer.blocks
assert blocks is not None and len(blocks) > 0
layer0 = blocks[0]
trainable = [p for p in layer0.parameters() if p.requires_grad]
assert len(trainable) > 0, "student layer 0 has no trainable parameters"
for i, p in enumerate(trainable):
assert p.grad is not None, f"student layer 0 param[{i}] has None grad"
assert torch.isfinite(p.grad).all().item(), (
f"student layer 0 param[{i}] grad contains NaN/Inf")
assert any(p.grad.detach().float().norm().item() > 0.0 for p in trainable), (
"all student layer-0 grads are exactly zero; consistency loss "
"did not reach the first transformer block")
# Teacher must stay frozen.
assert all(not p.requires_grad for p in teacher.transformer.parameters()), (
"teacher must be frozen for Causal-CD")
# The EMA model and student start from the same checkpoint, so the first
# parameter must match before any update. FSDP fully_shard params are
# DTensors; compare the local shards (torch.equal is unsupported on
# DTensor).
def _local(p: torch.Tensor) -> torch.Tensor:
return p.to_local() if hasattr(p, "to_local") else p
ema_param = next(ema.transformer.parameters())
student_param = next(student.transformer.parameters())
assert torch.equal(_local(ema_param), _local(student_param)), (
"EMA model should start identical to the student (same checkpoint)")
# The EMA update must move EMA toward the student. Apply a visibly large
# perturbation so the bf16 lerp is well above rounding noise (the real
# optimizer step at lr=2e-6 would be sub-ULP in bf16).
with torch.no_grad():
student_param.add_(1.0)
before = _local(ema_param).detach().float().clone()
method._update_ema()
after = _local(ema_param).detach().float()
assert not torch.equal(before, after), (
"EMA weights did not move after _update_ema")
# EMA = decay*ema + (1-decay)*student moves ~ (1-decay) of the gap.
expected = before + (1.0 - method._ema_decay) * (
_local(student_param).detach().float() - before)
assert torch.allclose(after, expected, atol=1e-2), (
"EMA update did not follow the expected lerp")
@@ -0,0 +1,109 @@
# SPDX-License-Identifier: Apache-2.0
"""Per-method GPU smoke test: frame-wise ``WanCausalModel`` + ``DiffusionForcingSFTMethod``.
Mirrors ``test_wan_causal_dfsft.py`` but with a block size of 1 frame
(``num_frames_per_block=1`` on the model, ``chunk_size=1`` on the method),
so each frame gets its own independent noise level. The test asserts the
override took effect and runs one train step: forward, finite loss, and
nonzero gradients reaching the first transformer block.
"""
from __future__ import annotations
import os
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29520")
from pathlib import Path
import pytest
import torch
from fastvideo.train.methods.fine_tuning.dfsft import (
DiffusionForcingSFTMethod, )
from fastvideo.train.models.wan import WanCausalModel
from fastvideo.train.utils.config import load_run_config
_FIXTURE = str(
Path(__file__).resolve().parent.parent / "fixtures"
/ "wan_causal_t2v_dfsft_framewise_min.yaml")
def _build_synthetic_batch(
device: torch.device,
dtype: torch.dtype,
) -> dict[str, torch.Tensor]:
batch_size = 1
return {
"text_embedding":
torch.randn(batch_size, 16, 4096, device=device, dtype=dtype),
"text_attention_mask":
torch.ones(batch_size, 16, device=device),
"vae_latent":
torch.randn(batch_size, 16, 6, 8, 8, device=device, dtype=dtype),
}
@pytest.mark.usefixtures("distributed_setup")
def test_wan_causal_dfsft_framewise_single_train_step(
monkeypatch: pytest.MonkeyPatch) -> None:
if not torch.cuda.is_available():
pytest.skip("requires CUDA")
cfg = load_run_config(_FIXTURE)
device = torch.device("cuda:0")
dtype = torch.bfloat16
monkeypatch.setattr(
"fastvideo.train.utils.dataloader."
"build_parquet_t2v_train_dataloader",
lambda *args, **kwargs: None,
)
student_cfg = cfg.models["student"]
model = WanCausalModel(
init_from=student_cfg["init_from"],
training_config=cfg.training,
trainable=True,
num_frames_per_block=student_cfg.get("num_frames_per_block"),
)
assert model.transformer.num_frame_per_block == 1, (
"frame-wise override did not reach the transformer")
model.transformer = model.transformer.to(device=device, dtype=dtype)
method = DiffusionForcingSFTMethod(
cfg=cfg,
role_models={"student": model},
)
method.on_train_start()
batch = _build_synthetic_batch(device, dtype)
loss_map, outputs, _metrics = method.single_train_step(batch, iteration=0)
loss = loss_map["total_loss"]
assert torch.is_tensor(loss), "total_loss must be a torch.Tensor"
assert torch.isfinite(loss).item(), (
f"total_loss is not finite: {loss.item()}")
method.backward(loss_map, outputs, grad_accum_rounds=1)
blocks = getattr(model.transformer, "blocks", None)
assert blocks is not None and len(blocks) > 0
layer0 = blocks[0]
trainable = [p for p in layer0.parameters() if p.requires_grad]
assert len(trainable) > 0, "layer 0 has no trainable parameters"
for i, p in enumerate(trainable):
assert p.grad is not None, f"layer 0 param[{i}] has None grad"
assert torch.isfinite(p.grad).all().item(), (
f"layer 0 param[{i}] grad contains NaN/Inf")
any_nonzero = any(
p.grad.detach().float().norm().item() > 0.0 for p in trainable)
assert any_nonzero, (
"all layer-0 grads are exactly zero; backward did not "
"reach the first transformer block")
@@ -0,0 +1,113 @@
# SPDX-License-Identifier: Apache-2.0
"""Per-method GPU smoke test: ``WanCausalModel`` + ``TeacherForcingSFTMethod``.
Mirrors ``test_wan_causal_dfsft.py``. Teacher forcing concatenates a clean
context copy of every frame inside the causal transformer (``clean_x``) and
denoises the current block while attending to *clean* history. This test
exercises the ``clean_x`` path end-to-end: forward, finite loss, and nonzero
gradients reaching the first transformer block.
"""
from __future__ import annotations
import os
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29518")
from pathlib import Path
import pytest
import torch
from fastvideo.train.methods.fine_tuning.tfsft import (
TeacherForcingSFTMethod, )
from fastvideo.train.models.wan import WanCausalModel
from fastvideo.train.utils.config import load_run_config
_FIXTURE = str(
Path(__file__).resolve().parent.parent / "fixtures"
/ "wan_causal_t2v_tfsft_min.yaml")
def _build_synthetic_batch(
device: torch.device,
dtype: torch.dtype,
) -> dict[str, torch.Tensor]:
batch_size = 1
return {
"text_embedding":
torch.randn(batch_size, 16, 4096, device=device, dtype=dtype),
"text_attention_mask":
torch.ones(batch_size, 16, device=device),
"vae_latent":
torch.randn(batch_size, 16, 6, 8, 8, device=device, dtype=dtype),
}
@pytest.mark.usefixtures("distributed_setup")
def test_wan_causal_tfsft_single_train_step(
monkeypatch: pytest.MonkeyPatch) -> None:
if not torch.cuda.is_available():
pytest.skip("requires CUDA")
cfg = load_run_config(_FIXTURE)
device = torch.device("cuda:0")
dtype = torch.bfloat16
monkeypatch.setattr(
"fastvideo.train.utils.dataloader."
"build_parquet_t2v_train_dataloader",
lambda *args, **kwargs: None,
)
model = WanCausalModel(
init_from=cfg.models["student"]["init_from"],
training_config=cfg.training,
trainable=True,
)
model.transformer = model.transformer.to(device=device, dtype=dtype)
method = TeacherForcingSFTMethod(
cfg=cfg,
role_models={"student": model},
)
method.on_train_start()
batch = _build_synthetic_batch(device, dtype)
loss_map, outputs, _metrics = method.single_train_step(batch, iteration=0)
loss = loss_map["total_loss"]
assert torch.is_tensor(loss), "total_loss must be a torch.Tensor"
assert torch.isfinite(loss).item(), (
f"total_loss is not finite: {loss.item()}")
method.backward(loss_map, outputs, grad_accum_rounds=1)
blocks = getattr(model.transformer, "blocks", None)
assert blocks is not None and len(blocks) > 0, (
"CausalWanTransformer is expected to expose ``.blocks``")
layer0 = blocks[0]
trainable = [p for p in layer0.parameters() if p.requires_grad]
assert len(trainable) > 0, "layer 0 has no trainable parameters"
for i, p in enumerate(trainable):
assert p.grad is not None, f"layer 0 param[{i}] has None grad"
assert torch.isfinite(p.grad).all().item(), (
f"layer 0 param[{i}] grad contains NaN/Inf")
any_nonzero = any(
p.grad.detach().float().norm().item() > 0.0 for p in trainable)
assert any_nonzero, (
"all layer-0 grads are exactly zero; backward did not "
"reach the first transformer block")
# Teacher forcing must build its own (concatenated) attention mask and
# must not have constructed the diffusion-forcing mask.
assert model.transformer.teacher_forcing_block_mask is not None, (
"teacher-forcing mask was not constructed")
assert model.transformer.block_mask is None, (
"diffusion-forcing mask should not be built on the TF path")
+8 -9
View File
@@ -456,11 +456,16 @@ class ValidationCallback(Callback):
None,
)
loaded_modules: dict[str, Any] = {"transformer": transformer}
# Distillation methods build the flow-match scheduler their few-step DMD
# sampler needs; inject it so the pipeline doesn't fall back to UniPC.
method_scheduler = getattr(self.method, "_sf_scheduler", None)
if method_scheduler is not None:
loaded_modules["scheduler"] = method_scheduler
kwargs: dict[str, Any] = {
"inference_mode": True,
"loaded_modules": {
"transformer": transformer,
},
"loaded_modules": loaded_modules,
"tp_size": tc.distributed.tp_size,
"sp_size": tc.distributed.sp_size,
"num_gpus": tc.distributed.num_gpus,
@@ -477,12 +482,6 @@ class ValidationCallback(Callback):
**kwargs,
)
scheduler = self._pipeline.get_module("scheduler")
if (scheduler is not None and type(scheduler).__name__ == "SelfForcingFlowMatchScheduler"):
scheduler.sigma_min = 0.0
scheduler.extra_one_step = True
scheduler.set_timesteps(num_inference_steps=1000, training=True)
self._pipeline_key = key
return self._pipeline
+9
View File
@@ -9,6 +9,8 @@ __all__ = [
"KDMethod",
"SelfForcingMethod",
"DiffusionForcingSFTMethod",
"TeacherForcingSFTMethod",
"CausalConsistencyDistillationMethod",
]
@@ -28,4 +30,11 @@ def __getattr__(name: str) -> object:
if name == "DiffusionForcingSFTMethod":
from fastvideo.train.methods.fine_tuning.dfsft import DiffusionForcingSFTMethod
return DiffusionForcingSFTMethod
if name == "TeacherForcingSFTMethod":
from fastvideo.train.methods.fine_tuning.tfsft import TeacherForcingSFTMethod
return TeacherForcingSFTMethod
if name == "CausalConsistencyDistillationMethod":
from fastvideo.train.methods.consistency_model.causal_cd import (
CausalConsistencyDistillationMethod, )
return CausalConsistencyDistillationMethod
raise AttributeError(name)
@@ -1,3 +1,22 @@
# SPDX-License-Identifier: Apache-2.0
__all__: list[str] = []
from __future__ import annotations
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from fastvideo.train.methods.consistency_model.causal_cd import (
CausalConsistencyDistillationMethod, )
__all__ = [
"CausalConsistencyDistillationMethod",
]
def __getattr__(name: str) -> object:
if name == "CausalConsistencyDistillationMethod":
from fastvideo.train.methods.consistency_model.causal_cd import (
CausalConsistencyDistillationMethod, )
return CausalConsistencyDistillationMethod
raise AttributeError(name)
@@ -0,0 +1,237 @@
# SPDX-License-Identifier: Apache-2.0
"""Causal consistency distillation method (algorithm layer)."""
from __future__ import annotations
from typing import Any
import torch
import torch.nn.functional as F
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (SelfForcingFlowMatchScheduler)
from fastvideo.train.methods.base import LogScalar, TrainingMethod
from fastvideo.train.models.base import ModelBase
from fastvideo.train.utils.checkpoint import _FullModelState
from fastvideo.train.utils.optimizer import build_optimizer_and_scheduler
class CausalConsistencyDistillationMethod(TrainingMethod):
def __init__(
self,
*,
cfg: Any,
role_models: dict[str, ModelBase],
) -> None:
super().__init__(cfg=cfg, role_models=role_models)
for role in ("student", "teacher", "ema"):
if role not in role_models:
raise ValueError(f"Causal-CD requires role {role!r} "
"(student trainable; teacher + ema frozen, "
"both initialized from the student's "
"checkpoint)")
if not self.student._trainable:
raise ValueError("Causal-CD requires student to be trainable")
self.teacher = role_models["teacher"]
self.ema_model = role_models["ema"]
self._attn_kind = self._infer_attn_kind()
self._guidance_scale = float(self.method_config.get("guidance_scale", 3.0))
self._discrete_cd_n = int(self.method_config.get("discrete_cd_N", 48))
if self._discrete_cd_n < 2:
raise ValueError("method.discrete_cd_N must be >= 2")
self._ema_decay = float(self.method_config.get("ema_decay", 0.99))
self._ema_start_step = int(self.method_config.get("ema_start_step", 200))
shift = getattr(self.training_config.pipeline_config, "flow_shift", None)
self._flow_shift = float(shift) if shift else 5.0
self.student.init_preprocessors(self.training_config)
self._sf_scheduler = SelfForcingFlowMatchScheduler(
num_inference_steps=self._discrete_cd_n,
num_train_timesteps=int(self.student.num_train_timesteps),
shift=self._flow_shift,
sigma_min=0.0,
sigma_max=1.0,
extra_one_step=True,
training=False,
)
self._init_optimizers_and_schedulers()
# ------------------------------------------------------------------
@property
def _optimizer_dict(self) -> dict[str, Any]:
return {"student": self._student_optimizer}
@property
def _lr_scheduler_dict(self) -> dict[str, Any]:
return {"student": self._student_lr_scheduler}
def get_optimizers(self, iteration: int) -> list[torch.optim.Optimizer]:
del iteration
return [self._student_optimizer]
def get_lr_schedulers(self, iteration: int) -> list[Any]:
del iteration
return [self._student_lr_scheduler]
def checkpoint_state(self) -> dict[str, Any]:
# The EMA role is frozen (so the base class skips it) but mutated by
# _update_ema every step; without persisting it a resume reloads the
# EMA from init_from and the consistency target snaps back to the
# base checkpoint. Mirrors DiffusionNFT's frozen "old" role.
states = super().checkpoint_state()
states["roles.ema.transformer"] = _FullModelState(self.ema_model.transformer)
return states
# ------------------------------------------------------------------
def single_train_step(
self,
batch: dict[str, Any],
iteration: int,
) -> tuple[dict[str, torch.Tensor], dict[str, Any], dict[str, LogScalar]]:
del iteration
training_batch = self.student.prepare_batch(
batch,
generator=self.cuda_generator,
latents_source="data",
)
clean_latents = training_batch.latents
if not torch.is_tensor(clean_latents) or clean_latents.ndim != 5:
raise ValueError("Causal-CD expects [B, T, C, H, W] latents")
batch_size, num_latents = int(clean_latents.shape[0]), int(clean_latents.shape[1])
device = clean_latents.device
sigmas = self._sf_scheduler.sigmas.to(device)
timesteps = self._sf_scheduler.timesteps.to(device)
idx = torch.randint(0, self._discrete_cd_n - 1, (1, ), generator=self.cuda_generator, device=device).squeeze(0)
t, t_next = timesteps[idx], timesteps[idx + 1]
sigma_t, sigma_t_next = sigmas[idx], sigmas[idx + 1]
t_pf = t * torch.ones(batch_size, num_latents, device=device)
t_next_pf = t_next * torch.ones(batch_size, num_latents, device=device)
noise = torch.randn(
clean_latents.shape,
generator=self.cuda_generator,
device=device,
dtype=clean_latents.dtype,
)
latent_t = (1.0 - sigma_t) * clean_latents + sigma_t * noise
# Set before any forward: predict_noise feeds batch.timesteps into
# set_forward_context (VSA sparsity gating), so the teacher CFG
# passes below must not see the stale timesteps from prepare_batch.
training_batch.timesteps = t_pf
with torch.no_grad():
v_cond = self._predict_flow(self.teacher,
latent_t,
t_pf,
training_batch,
conditional=True,
clean_x=clean_latents)
v_uncond = self._predict_flow(self.teacher,
latent_t,
t_pf,
training_batch,
conditional=False,
clean_x=clean_latents)
v_pred = v_uncond + self._guidance_scale * (v_cond - v_uncond)
dt = ((t - t_next) / float(self.student.num_train_timesteps))
latent_t_next = latent_t - dt * v_pred
flow_student = self._predict_flow(self.student,
latent_t,
t_pf,
training_batch,
conditional=True,
clean_x=clean_latents)
x0_t = latent_t - sigma_t * flow_student
with torch.no_grad():
flow_ema = self._predict_flow(self.ema_model,
latent_t_next,
t_next_pf,
training_batch,
conditional=True,
clean_x=clean_latents)
x0_t_next = latent_t_next - sigma_t_next * flow_ema
loss = F.mse_loss(x0_t.float(), x0_t_next.float())
loss_map = {"total_loss": loss, "causal_cd_loss": loss}
attn_metadata = (training_batch.attn_metadata_vsa if self._attn_kind == "vsa" else training_batch.attn_metadata)
outputs: dict[str, Any] = {"_fv_backward": (t_pf, attn_metadata)}
metrics: dict[str, LogScalar] = {}
return loss_map, outputs, metrics
# ------------------------------------------------------------------
def backward(
self,
loss_map: dict[str, torch.Tensor],
outputs: dict[str, Any],
*,
grad_accum_rounds: int = 1,
) -> None:
grad_accum_rounds = max(1, int(grad_accum_rounds))
ctx = outputs.get("_fv_backward")
if ctx is None:
super().backward(loss_map, outputs, grad_accum_rounds=grad_accum_rounds)
return
self.student.backward(loss_map["total_loss"], ctx, grad_accum_rounds=grad_accum_rounds)
def optimizers_schedulers_step(self, iteration: int) -> None:
super().optimizers_schedulers_step(iteration)
if iteration >= self._ema_start_step:
self._update_ema()
# ------------------------------------------------------------------
# Internal helpers
# ------------------------------------------------------------------
def _predict_flow(
self,
model: ModelBase,
latents: torch.Tensor,
timestep: torch.Tensor,
batch: Any,
*,
conditional: bool,
clean_x: torch.Tensor,
) -> torch.Tensor:
return model.predict_noise(latents,
timestep,
batch,
conditional=conditional,
cfg_uncond=None,
attn_kind=self._attn_kind,
clean_x=clean_x)
@torch.no_grad()
def _update_ema(self) -> None:
decay = self._ema_decay
for ema_p, p in zip(self.ema_model.transformer.parameters(), self.student.transformer.parameters(),
strict=True):
ema_p.mul_(decay).add_(p.detach().to(ema_p.dtype), alpha=1.0 - decay)
def _init_optimizers_and_schedulers(self) -> None:
tc = self.training_config
student_lr = float(tc.optimizer.learning_rate)
if student_lr <= 0.0:
raise ValueError("training.learning_rate must be > 0 for causal-cd")
student_params = [p for p in self.student.transformer.parameters() if p.requires_grad]
(
self._student_optimizer,
self._student_lr_scheduler,
) = build_optimizer_and_scheduler(
params=student_params,
optimizer_config=tc.optimizer,
loop_config=tc.loop,
learning_rate=student_lr,
betas=tc.optimizer.betas,
scheduler_name=str(tc.optimizer.lr_scheduler),
)
@@ -7,10 +7,12 @@ from typing import TYPE_CHECKING
if TYPE_CHECKING:
from fastvideo.train.methods.fine_tuning.dfsft import DiffusionForcingSFTMethod
from fastvideo.train.methods.fine_tuning.finetune import FineTuneMethod
from fastvideo.train.methods.fine_tuning.tfsft import TeacherForcingSFTMethod
__all__ = [
"DiffusionForcingSFTMethod",
"FineTuneMethod",
"TeacherForcingSFTMethod",
]
@@ -25,4 +27,9 @@ def __getattr__(name: str) -> object:
from fastvideo.train.methods.fine_tuning.finetune import FineTuneMethod
return FineTuneMethod
if name == "TeacherForcingSFTMethod":
from fastvideo.train.methods.fine_tuning.tfsft import (
TeacherForcingSFTMethod, )
return TeacherForcingSFTMethod
raise AttributeError(name)
+19 -3
View File
@@ -135,12 +135,11 @@ class DiffusionForcingSFTMethod(TrainingMethod):
t_inhom.flatten(),
)
pred = self.student.predict_noise(
pred = self._predict_noise(
noisy_latents,
t_inhom,
training_batch,
conditional=True,
attn_kind=self._attn_kind,
clean_latents,
)
if bool(self.training_config.model.precondition_outputs):
@@ -178,6 +177,23 @@ class DiffusionForcingSFTMethod(TrainingMethod):
metrics: dict[str, LogScalar] = {}
return loss_map, outputs, metrics
def _predict_noise(
self,
noisy_latents: torch.Tensor,
timestep: torch.Tensor,
training_batch: Any,
clean_latents: torch.Tensor,
) -> torch.Tensor:
# Unused here; the teacher-forcing subclass overrides this to pass clean_x.
del clean_latents
return self.student.predict_noise(
noisy_latents,
timestep,
training_batch,
conditional=True,
attn_kind=self._attn_kind,
)
# TrainingMethod override: backward
def backward(
self,
@@ -0,0 +1,30 @@
# SPDX-License-Identifier: Apache-2.0
"""Teacher-forcing SFT method (TFSFT; algorithm layer)."""
from __future__ import annotations
from typing import Any
import torch
from fastvideo.train.methods.fine_tuning.dfsft import (
DiffusionForcingSFTMethod, )
class TeacherForcingSFTMethod(DiffusionForcingSFTMethod):
def _predict_noise(
self,
noisy_latents: torch.Tensor,
timestep: torch.Tensor,
training_batch: Any,
clean_latents: torch.Tensor,
) -> torch.Tensor:
return self.student.predict_noise(
noisy_latents,
timestep,
training_batch,
conditional=True,
attn_kind=self._attn_kind,
clean_x=clean_latents,
)
+1 -31
View File
@@ -10,12 +10,6 @@ from typing import Any
import torch
import torch.distributed as dist
from torch.distributed.checkpoint.state_dict import (
StateDictOptions,
get_model_state_dict,
set_model_state_dict,
)
from torch.distributed.checkpoint.stateful import Stateful
from tqdm.auto import tqdm
from fastvideo.dataset.parquet_dataset_map_style import (
@@ -37,6 +31,7 @@ from fastvideo.train.methods.rl.common import (
validation_caption,
validation_shard_indices,
)
from fastvideo.train.utils.checkpoint import _FullModelState
from fastvideo.train.utils.config import (
get_optional_float,
get_optional_int,
@@ -78,31 +73,6 @@ class _DiffusionNFTEMAState:
self._method._ema_update_count = int(update_count)
class _FullModelState(Stateful):
"""DCP wrapper that saves frozen model parameters too.
The shared ``ModelWrapper`` intentionally filters to ``requires_grad``
parameters. DiffusionNFT's old policy is frozen but must be restored on
resume, so it needs full model state.
"""
def __init__(self, model: torch.nn.Module) -> None:
self.model = model
def state_dict(self) -> dict[str, Any]:
return get_model_state_dict(self.model) # type: ignore[no-any-return]
def load_state_dict(
self,
state_dict: dict[str, Any],
) -> None:
set_model_state_dict(
self.model,
model_state_dict=state_dict,
options=StateDictOptions(strict=False),
)
class DiffusionNFTMethod(TrainingMethod):
"""DiffusionNFT-style RL for diffusion models.
+15 -2
View File
@@ -320,6 +320,8 @@ class WanModel(ModelBase):
conditional: bool,
cfg_uncond: dict[str, Any] | None = None,
attn_kind: Literal["dense", "vsa"] = "dense",
clean_x: torch.Tensor | None = None,
aug_t: torch.Tensor | None = None,
) -> torch.Tensor:
device_type = self.device.type
dtype = self._get_training_dtype()
@@ -347,7 +349,11 @@ class WanModel(ModelBase):
current_timestep=batch.timesteps,
attn_metadata=attn_metadata,
):
input_kwargs = (self._build_distill_input_kwargs(noisy_latents, timestep, text_dict))
input_kwargs = (self._build_distill_input_kwargs(noisy_latents,
timestep,
text_dict,
clean_x=clean_x,
aug_t=aug_t))
transformer = self._get_transformer(timestep)
pred_noise = transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
return pred_noise
@@ -530,17 +536,24 @@ class WanModel(ModelBase):
noise_input: torch.Tensor,
timestep: torch.Tensor,
text_dict: dict[str, torch.Tensor] | None,
clean_x: torch.Tensor | None = None,
aug_t: torch.Tensor | None = None,
) -> dict[str, Any]:
if text_dict is None:
raise ValueError("text_dict cannot be None for "
"Wan distillation")
return {
kwargs: dict[str, Any] = {
"hidden_states": noise_input.permute(0, 2, 1, 3, 4),
"encoder_hidden_states": text_dict["encoder_hidden_states"],
"encoder_attention_mask": text_dict["encoder_attention_mask"],
"timestep": timestep,
"return_dict": False,
}
if clean_x is not None:
# Teacher forcing: clean context latents (+ optional aug timestep).
kwargs["clean_x"] = clean_x.permute(0, 2, 1, 3, 4)
kwargs["aug_t"] = aug_t
return kwargs
def _get_transformer(self, timestep: torch.Tensor) -> torch.nn.Module:
return self.transformer
+11
View File
@@ -49,6 +49,7 @@ class WanCausalModel(WanModel, CausalModelBase):
transformer_override_safetensor: str
| None = None,
lora: LoraConfig | dict[str, Any] | None = None,
num_frames_per_block: int | None = None,
) -> None:
super().__init__(
init_from=init_from,
@@ -62,6 +63,16 @@ class WanCausalModel(WanModel, CausalModelBase):
)
self._streaming_caches: (dict[tuple[int, str], _StreamingCaches]) = {}
if num_frames_per_block is not None:
num_frames_per_block = int(num_frames_per_block)
if not 1 <= num_frames_per_block <= 3:
# Same bound as CausalWanTransformer3DModel's config path
# (assert num_frame_per_block <= 3); this override must not
# bypass it.
raise ValueError("num_frames_per_block must be between 1 and 3, "
f"got {num_frames_per_block}")
self.transformer.num_frame_per_block = num_frames_per_block
# --- CausalModelBase override: clear_caches ---
def clear_caches(
self,
+33
View File
@@ -15,6 +15,13 @@ import numpy as np
import torch
import torch.distributed as dist
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import (
StateDictOptions,
get_model_state_dict,
set_model_state_dict,
)
from torch.distributed.checkpoint.stateful import Stateful
from fastvideo.logger import init_logger
logger = init_logger(__name__)
@@ -131,6 +138,32 @@ class _RoleModuleContainer(torch.nn.Module):
self.add_module(name, module)
class _FullModelState(Stateful):
"""DCP wrapper that saves frozen model parameters too.
The shared ``ModelWrapper`` intentionally filters to ``requires_grad``
parameters. Frozen-but-mutated roles (e.g. DiffusionNFT's old policy,
causal-CD's EMA target) must still be restored on resume, so they need
full model state.
"""
def __init__(self, model: torch.nn.Module) -> None:
self.model = model
def state_dict(self) -> dict[str, Any]:
return get_model_state_dict(self.model) # type: ignore[no-any-return]
def load_state_dict(
self,
state_dict: dict[str, Any],
) -> None:
set_model_state_dict(
self.model,
model_state_dict=state_dict,
options=StateDictOptions(strict=False),
)
class _CallbackStateWrapper:
"""Wraps a CallbackDict for DCP save/load."""