Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4d8713f6e4 | ||
|
|
fda43fc610 | ||
|
|
5a8329ddf6 | ||
|
|
fcb5b465c5 | ||
|
|
4f13fb0fee | ||
|
|
d96ec99b4b | ||
|
|
2555f25cce | ||
|
|
fbe56fce8d | ||
|
|
5064bcdc47 | ||
|
|
2aa824b615 | ||
|
|
f934efc58d |
@@ -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
|
||||
@@ -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")
|
||||
@@ -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,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)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user