Compare commits

...
Author SHA1 Message Date
RandNMR73 44840ac49d wan2.2 training doesn't crash 2025-09-20 05:08:58 +00:00
JerryZhou54 c99b1d4d97 Change Wan DiT to have 0 numerical diff with SF's Wan 2025-09-18 22:50:07 +00:00
SolitaryThinker edfe4dd1bf fix rotary embedding 2025-09-13 08:31:53 +00:00
12 changed files with 1685 additions and 475 deletions
@@ -129,7 +129,7 @@ dmd_args=(
# Self-forcing specific arguments
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks False # Whether to use same denoising step across all blocks
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
--validate_cache_structure False # Set to True for debugging KV cache issues
@@ -0,0 +1,151 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:1
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=dmd_t2v_output/t2v_%j.out
#SBATCH --error=dmd_t2v_output/t2v_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29503
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="2f25ad37933894dbf0966c838c0b8494987f9f2f"
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# Configs
NUM_GPUS=8
# Model paths for Self-Forcing DMD distillation with Wan2.2:
GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name SFwan2.2_t2v_distill_self_forcing_dmd # Updated for Wan2.2
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan2.2_t2v_finetune"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 16
--num_height 448 # Updated to match Wan2.2 config
--num_width 832 # Updated to match Wan2.2 config
--num_frames 61 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frame_per_block 4 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 16
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS # 64
--sp_size 4
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim 8
)
# Model arguments
model_args=(
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--generator_model_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 100
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
)
# Self-forcing DMD arguments
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 3.0
--fake_score_learning_rate 8e-6
--fake_score_betas '0.0,0.999'
--warp_denoising_step
)
# Self-forcing specific arguments
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
--validate_cache_structure False # Set to True for debugging KV cache issues
)
torchrun \
--nnodes 1 \
--master_port $MASTER_PORT \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
+5 -5
View File
@@ -212,9 +212,9 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
frame_seqlen = normalized.shape[1] // num_frames
modulated = (
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1.0 + scale) + shift).flatten(1, 2)
(1 + scale) + shift).flatten(1, 2)
else:
modulated = normalized * (1.0 + scale) + shift
modulated = normalized * (1 + scale) + shift
return modulated, residual_output
@@ -267,13 +267,13 @@ class LayerNormScaleShift(nn.Module):
frame_seqlen = normalized.shape[1] // num_frames
output = (
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1.0 + scale) + shift).flatten(1, 2)
(1 + scale) + shift).flatten(1, 2)
else:
# scale.shape: [batch_size, 1, inner_dim]
# shift.shape: [batch_size, 1, inner_dim]
output = normalized * (1.0 + scale) + shift
output = normalized * (1 + scale) + shift
if self.compute_dtype == torch.float32:
output = output.to(x.dtype)
return output
return output
+62 -40
View File
@@ -147,8 +147,8 @@ class CausalWanSelfAttention(nn.Module):
# Assign new keys/values directly up to current_end
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
kv_cache["k"] = kv_cache["k"].detach()
kv_cache["v"] = kv_cache["v"].detach()
# kv_cache["k"] = kv_cache["k"].detach()
# kv_cache["v"] = kv_cache["v"].detach()
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
@@ -179,7 +179,7 @@ class CausalWanTransformerBlock(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -212,8 +212,7 @@ class CausalWanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 2. Cross-attention
# Only T2V for now
@@ -226,8 +225,7 @@ class CausalWanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -252,29 +250,34 @@ class CausalWanTransformerBlock(nn.Module):
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
num_frames = temb.shape[1]
frame_seqlen = hidden_states.shape[1] // num_frames
frame_seqlen = hidden_states.shape[1] // num_frames
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
e = self.scale_shift_table + temb
# e.shape: [batch_size, num_frames, 6, inner_dim]
assert e.shape == (bs, num_frames, 6, self.hidden_dim)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=2)
# *_msa.shape: [batch_size, num_frames, 1, inner_dim]
assert shift_msa.dtype == torch.float32
# assert shift_msa.dtype == torch.float32
# logger.info("temb sum: %s, dtype: %s", temb.float().sum().item(), temb.dtype)
# logger.info("scale_msa sum: %s, dtype: %s", scale_msa.float().sum().item(), scale_msa.dtype)
# logger.info("shift_msa sum: %s, dtype: %s", shift_msa.float().sum().item(), shift_msa.dtype)
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).flatten(1, 2).to(orig_dtype)
norm_hidden_states = (self.norm1(hidden_states).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).flatten(1, 2)
# logger.info("norm_hidden_states sum: %s, shape: %s", norm_hidden_states.float().sum().item(), norm_hidden_states.shape)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k(key)
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -288,8 +291,6 @@ class CausalWanTransformerBlock(nn.Module):
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
@@ -298,13 +299,10 @@ class CausalWanTransformerBlock(nn.Module):
crossattn_cache=crossattn_cache)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
@@ -367,8 +365,7 @@ class CausalWanTransformer3DModel(BaseDiT):
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
@@ -378,7 +375,7 @@ class CausalWanTransformer3DModel(BaseDiT):
# Causal-specific
self.block_mask = None
self.num_frame_per_block = 1
self.num_frame_per_block = 3
self.independent_first_frame = False
self.__post_init__()
@@ -490,12 +487,16 @@ class CausalWanTransformer3DModel(BaseDiT):
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
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)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
@@ -542,14 +543,9 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
output = self.unpatchify(hidden_states, grid_sizes)
return output
return torch.stack(output)
def _forward_train(self,
hidden_states: torch.Tensor,
@@ -590,8 +586,8 @@ class CausalWanTransformer3DModel(BaseDiT):
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
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:
@@ -604,8 +600,12 @@ class CausalWanTransformer3DModel(BaseDiT):
)
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
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)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
@@ -640,14 +640,9 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
output = self.unpatchify(hidden_states, grid_sizes)
return output
return torch.stack(output)
def forward(
self,
@@ -658,3 +653,30 @@ class CausalWanTransformer3DModel(BaseDiT):
return self._forward_inference(*args, **kwargs)
else:
return self._forward_train(*args, **kwargs)
def unpatchify(self, x, grid_sizes):
r"""
Args:
x (List[Tensor]):
List of patchified features, each with shape [L, C_out * prod(patch_size)]
grid_sizes (Tensor):
Original spatial-temporal grid dimensions before patching,
Returns:
Tensor:
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
"""
c = self.out_channels
out = []
for u, v in zip(x, grid_sizes.tolist()):
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
u = u.permute(6, 0, 3, 1, 4, 2, 5)
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
out.append(u)
return out
+80 -73
View File
@@ -1,3 +1,5 @@
import torch
import torch.nn as nn
# SPDX-License-Identifier: Apache-2.0
import math
@@ -37,16 +39,14 @@ class WanImageEmbedding(torch.nn.Module):
def __init__(self, in_features: int, out_features: int):
super().__init__()
self.norm1 = FP32LayerNorm(in_features)
self.norm1 = nn.LayerNorm(in_features)
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
self.norm2 = FP32LayerNorm(out_features)
self.norm2 = nn.LayerNorm(out_features)
def forward(self,
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
dtype = encoder_hidden_states_image.dtype
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
hidden_states = self.norm1(encoder_hidden_states_image)
hidden_states = self.ff(hidden_states)
hidden_states = self.norm2(hidden_states).to(dtype)
hidden_states = self.norm2(hidden_states)
return hidden_states
@@ -62,7 +62,7 @@ class WanTimeTextImageEmbedding(nn.Module):
super().__init__()
self.time_embedder = TimestepEmbedder(
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
dim, frequency_embedding_size=time_freq_dim, act_layer="silu", freq_dtype=torch.float64)
self.time_modulation = ModulateProjection(dim,
factor=6,
act_layer="silu")
@@ -156,12 +156,12 @@ class WanT2VCrossAttention(WanSelfAttention):
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
if crossattn_cache is not None:
if not crossattn_cache["is_init"]:
crossattn_cache["is_init"] = True
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
crossattn_cache["k"] = k
crossattn_cache["v"] = v
@@ -169,7 +169,7 @@ class WanT2VCrossAttention(WanSelfAttention):
k = crossattn_cache["k"]
v = crossattn_cache["v"]
else:
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
# compute attention
@@ -213,10 +213,10 @@ class WanI2VCrossAttention(WanSelfAttention):
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
k_img = self.norm_added_k.forward_native(self.add_k_proj(context_img)[0]).view(
b, -1, n, d)
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
img_x = self.attn(q, k_img, v_img)
@@ -247,7 +247,7 @@ class WanTransformerBlock(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -278,29 +278,29 @@ class WanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
# I2V
self.attn2 = WanI2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
self.attn2 = WanI2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
else:
# T2V
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dim,
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -319,12 +319,11 @@ class WanTransformerBlock(nn.Module):
hidden_states = hidden_states.squeeze(1)
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
if temb.dim() == 4:
# temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
self.scale_shift_table.unsqueeze(0) + temb.float()
self.scale_shift_table.unsqueeze(0) + temb
).chunk(6, dim=2)
# batch_size, seq_len, 1, inner_dim
shift_msa = shift_msa.squeeze(2)
@@ -335,22 +334,20 @@ class WanTransformerBlock(nn.Module):
c_gate_msa = c_gate_msa.squeeze(2)
else:
# temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B)
e = self.scale_shift_table + temb.float()
e = self.scale_shift_table + temb
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
norm_hidden_states = self.norm1(hidden_states) * (1 + scale_msa) + shift_msa
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k(key)
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -370,26 +367,20 @@ class WanTransformerBlock(nn.Module):
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
context=encoder_hidden_states,
attn_output = self.attn2(norm_hidden_states,
context=encoder_hidden_states,
context_lens=None)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
class WanTransformerBlock_VSA(nn.Module):
def __init__(self,
@@ -406,7 +397,7 @@ class WanTransformerBlock_VSA(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -438,8 +429,7 @@ class WanTransformerBlock_VSA(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
@@ -459,8 +449,7 @@ class WanTransformerBlock_VSA(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -480,23 +469,22 @@ class WanTransformerBlock_VSA(nn.Module):
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
e = self.scale_shift_table + temb
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
norm_hidden_states = (self.norm1(hidden_states) *
(1 + scale_msa) + shift_msa)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
gate_compress, _ = self.to_gate_compress(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k(key)
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -521,8 +509,6 @@ class WanTransformerBlock_VSA(nn.Module):
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
@@ -530,17 +516,15 @@ class WanTransformerBlock_VSA(nn.Module):
context_lens=None)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
class WanTransformer3DModel(CachableDiT):
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
_compile_conditions = WanVideoConfig()._compile_conditions
@@ -598,8 +582,7 @@ class WanTransformer3DModel(CachableDiT):
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
@@ -659,10 +642,12 @@ class WanTransformer3DModel(CachableDiT):
rope_theta=10000)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
@@ -672,6 +657,8 @@ class WanTransformer3DModel(CachableDiT):
else:
ts_seq_len = None
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)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len)
if ts_seq_len is not None:
@@ -728,14 +715,35 @@ class WanTransformer3DModel(CachableDiT):
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
output = self.unpatchify(hidden_states, grid_sizes)
return output
return torch.stack(output)
def unpatchify(self, x, grid_sizes):
r"""
Args:
x (List[Tensor]):
List of patchified features, each with shape [L, C_out * prod(patch_size)]
grid_sizes (Tensor):
Original spatial-temporal grid dimensions before patching,
Returns:
Tensor:
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
"""
c = self.out_channels
out = []
for u, v in zip(x, grid_sizes.tolist()):
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
u = u.permute(6, 0, 3, 1, 4, 2, 5)
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
out.append(u)
return out
def maybe_cache_states(self, hidden_states: torch.Tensor,
original_hidden_states: torch.Tensor) -> None:
@@ -827,5 +835,4 @@ class WanTransformer3DModel(CachableDiT):
if self.is_even:
return hidden_states + self.previous_residual_even
else:
return hidden_states + self.previous_residual_odd
return hidden_states + self.previous_residual_odd
@@ -0,0 +1,292 @@
# SPDX-License-Identifier: Apache-2.0
import os
import math
import numpy as np
import pytest
import torch
from fastvideo.wan.modules.causal_model import CausalWanModel
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.models.dits.causal_wanvideo import CausalWanTransformer3DModel
from fastvideo.utils import maybe_download_model
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
@pytest.mark.usefixtures("distributed_setup")
def test_ori_causal_wan_transformer():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
dit_cpu_offload=True,
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
model1 = CausalWanModel.from_pretrained(
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
causal_state_dict = torch.load("/mnt/weka/home/hao.zhang/wei/Self-Forcing/checkpoints/self_forcing_dmd.pt")["generator_ema"]
new_state_dict = {}
for k, v in causal_state_dict.items():
if k.startswith("model."):
new_state_dict[k.replace("model.", "")] = v
causal_state_dict = new_state_dict
model1.load_state_dict(causal_state_dict)
total_params = sum(p.numel() for p in model1.parameters())
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
weight_sum_model1 = sum(
p.to(torch.float64).sum().item() for p in model1.parameters())
# Also calculate mean for more stable comparison
weight_mean_model1 = weight_sum_model1 / total_params
logger.info("Model 1 weight sum: %s", weight_sum_model1)
logger.info("Model 1 weight mean: %s", weight_mean_model1)
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
total_params_model2 = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(
p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params_model2
logger.info("Model 2 weight sum: %s", weight_sum_model2)
logger.info("Model 2 weight mean: %s", weight_mean_model2)
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
logger.info("Weight sum difference: %s", weight_sum_diff)
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
logger.info("Weight mean difference: %s", weight_mean_diff)
# Set both models to eval mode
model1 = model1.eval()
model2 = model2.eval()
# Create identical inputs for both models
batch_size = 1
text_seq_len = 30
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(batch_size,
16,
12,
160,
90,
device=device,
dtype=precision)
block_sizes = [3 for _ in range(4)]
timesteps = [1000, 750, 500, 250]
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(batch_size,
text_seq_len + 1,
4096,
device=device,
dtype=precision)
output1 = _causal_inference(model1, hidden_states.clone(), encoder_hidden_states.clone(), block_sizes, timesteps, precision)
logger.info("Finish inference for model1")
output2 = _causal_inference(model2, hidden_states.clone(), encoder_hidden_states.clone(), block_sizes, timesteps, precision)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
logger.info("Output 1 Sum: %s", output1.float().sum().item())
logger.info("Output 2 Sum: %s", output2.float().sum().item())
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
def _causal_inference(transformer, latents, prompt_embeds, block_sizes, timesteps, target_dtype):
forward_batch = ForwardBatch(
data_type="dummy",
)
start_index = 0
pos_start_base = 0
frame_seq_length = latents.shape[-1] * latents.shape[-2] // (WanVideoConfig().arch_config.patch_size[-1] * WanVideoConfig().arch_config.patch_size[-2])
seq_len = frame_seq_length * latents.shape[2]
kv_cache1 = _initialize_kv_cache(transformer, batch_size=latents.shape[0],
kv_cache_size=frame_seq_length * latents.shape[2],
dtype=target_dtype,
device=latents.device)
crossattn_cache = _initialize_crossattn_cache(
transformer,
batch_size=latents.shape[0],
max_text_len=WanVideoConfig().arch_config.text_len,
dtype=target_dtype,
device=latents.device)
for current_num_frames, t_cur in zip(block_sizes, timesteps):
# logger.info(f"Current frame idx: {start_index}, Current timestep: {t_cur}")
# logger.info(f"k cache sum: {sum(kv_cache['k'].float().sum().item() for kv_cache in kv_cache1)}, v cache sum: {sum(kv_cache['v'].float().sum().item() for kv_cache in kv_cache1)}")
# logger.info(f"latents sum: {latents.float().sum().item()}, encoder_hidden_states sum: {prompt_embeds.float().sum().item()}")
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
attn_metadata = None
with set_forward_context(current_timestep=0,
attn_metadata=attn_metadata,
forward_batch=forward_batch):
# Run transformer; follow DMD stage pattern
t_expanded_noise = t_cur * torch.ones(
(current_latents.shape[0], 1),
device=current_latents.device,
dtype=torch.long)
if isinstance(transformer, CausalWanModel):
pred_noise_btchw = transformer(
x=current_latents,
context=prompt_embeds,
t=t_expanded_noise,
seq_len=seq_len,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
frame_seq_length
)
elif isinstance(transformer, CausalWanTransformer3DModel):
pred_noise_btchw = transformer(
current_latents,
prompt_embeds,
t_expanded_noise,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
frame_seq_length,
start_frame=start_index
)
# Write back and advance
latents[:, :, start_index:start_index +
current_num_frames, :, :] = pred_noise_btchw.clone()
# Re-run with context timestep to update KV cache using clean context
context_noise = 0
t_context = torch.ones([latents.shape[0]],
device=latents.device,
dtype=torch.long) * int(context_noise)
context_bcthw = pred_noise_btchw.to(target_dtype)
with set_forward_context(current_timestep=0,
attn_metadata=attn_metadata,
forward_batch=forward_batch):
t_expanded_context = t_context.unsqueeze(1)
if isinstance(transformer, CausalWanModel):
_ = transformer(
x=context_bcthw,
context=prompt_embeds,
t=t_expanded_context,
seq_len=seq_len,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
frame_seq_length
)
elif isinstance(transformer, CausalWanTransformer3DModel):
_ = transformer(
context_bcthw,
prompt_embeds,
t_expanded_context,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
frame_seq_length,
start_frame=start_index
)
start_index += current_num_frames
return latents
def _initialize_kv_cache(transformer, batch_size, kv_cache_size, dtype, device) -> None:
"""
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
"""
kv_cache1 = []
if isinstance(transformer, CausalWanModel):
num_attention_heads = transformer.num_heads
attention_head_dim = transformer.dim // transformer.num_heads
elif isinstance(transformer, CausalWanTransformer3DModel):
num_attention_heads = transformer.num_attention_heads
attention_head_dim = transformer.attention_head_dim
for _ in range(len(transformer.blocks)):
kv_cache1.append({
"k":
torch.zeros([
batch_size, kv_cache_size, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"v":
torch.zeros([
batch_size, kv_cache_size, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
"local_end_index":
torch.tensor([0], dtype=torch.long, device=device),
})
return kv_cache1
def _initialize_crossattn_cache(transformer, batch_size, max_text_len, dtype,
device) -> None:
"""
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
"""
crossattn_cache = []
if isinstance(transformer, CausalWanModel):
num_attention_heads = transformer.num_heads
attention_head_dim = transformer.dim // transformer.num_heads
elif isinstance(transformer, CausalWanTransformer3DModel):
num_attention_heads = transformer.num_attention_heads
attention_head_dim = transformer.attention_head_dim
for _ in range(len(transformer.blocks)):
crossattn_cache.append({
"k":
torch.zeros([
batch_size, max_text_len, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"v":
torch.zeros([
batch_size, max_text_len, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"is_init":
False,
})
return crossattn_cache
@@ -0,0 +1,133 @@
# SPDX-License-Identifier: Apache-2.0
import os
import math
import numpy as np
import pytest
import torch
from fastvideo.wan.modules.model import WanModel
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.utils import maybe_download_model
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
@pytest.mark.usefixtures("distributed_setup")
def test_ori_wan_transformer():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
dit_cpu_offload=True,
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
model1 = WanModel.from_pretrained(
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
total_params = sum(p.numel() for p in model1.parameters())
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
weight_sum_model1 = sum(
p.to(torch.float64).sum().item() for p in model1.parameters())
# Also calculate mean for more stable comparison
weight_mean_model1 = weight_sum_model1 / total_params
logger.info("Model 1 weight sum: %s", weight_sum_model1)
logger.info("Model 1 weight mean: %s", weight_mean_model1)
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
total_params_model2 = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(
p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params_model2
logger.info("Model 2 weight sum: %s", weight_sum_model2)
logger.info("Model 2 weight mean: %s", weight_mean_model2)
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
logger.info("Weight sum difference: %s", weight_sum_diff)
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
logger.info("Weight mean difference: %s", weight_mean_diff)
# Set both models to eval mode
model1 = model1.eval()
model2 = model2.eval()
# Create identical inputs for both models
batch_size = 1
text_seq_len = 30
seq_len = math.ceil((160 * 90) /
(2 * 2) *
21)
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(batch_size,
16,
21,
160,
90,
device=device,
dtype=precision)
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(batch_size,
text_seq_len + 1,
4096,
device=device,
dtype=precision)
# Timestep
timestep = torch.tensor([500], device=device, dtype=precision)
forward_batch = ForwardBatch(
data_type="dummy",
)
# with torch.amp.autocast('cuda', dtype=precision):
output1 = model1(
x=hidden_states,
context=encoder_hidden_states,
t=timestep,
seq_len=seq_len,
)
with set_forward_context(
current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch,
):
output2 = model2(hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
@@ -0,0 +1,144 @@
# SPDX-License-Identifier: Apache-2.0
import os
import math
import numpy as np
import pytest
import torch
from fastvideo.wan.modules.causal_model import CausalWanModel
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.utils import maybe_download_model
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
@pytest.mark.usefixtures("distributed_setup")
def test_train_ori_causal_wan_transformer():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
dit_cpu_offload=True,
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
model1 = CausalWanModel.from_pretrained(
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
causal_state_dict = torch.load("/mnt/weka/home/hao.zhang/wei/Self-Forcing/checkpoints/self_forcing_dmd.pt")["generator_ema"]
new_state_dict = {}
for k, v in causal_state_dict.items():
if k.startswith("model."):
new_state_dict[k.replace("model.", "")] = v
causal_state_dict = new_state_dict
model1.load_state_dict(causal_state_dict)
model1.num_frame_per_block = 3
model2.num_frame_per_block = 3
total_params = sum(p.numel() for p in model1.parameters())
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
weight_sum_model1 = sum(
p.to(torch.float64).sum().item() for p in model1.parameters())
# Also calculate mean for more stable comparison
weight_mean_model1 = weight_sum_model1 / total_params
logger.info("Model 1 weight sum: %s", weight_sum_model1)
logger.info("Model 1 weight mean: %s", weight_mean_model1)
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
total_params_model2 = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(
p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params_model2
logger.info("Model 2 weight sum: %s", weight_sum_model2)
logger.info("Model 2 weight mean: %s", weight_mean_model2)
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
logger.info("Weight sum difference: %s", weight_sum_diff)
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
logger.info("Weight mean difference: %s", weight_mean_diff)
# Set both models to eval mode
model1 = model1.eval()
model2 = model2.eval()
# Create identical inputs for both models
batch_size = 1
text_seq_len = 30
seq_len = math.ceil((160 * 90) /
(2 * 2) *
21)
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(batch_size,
16,
21,
160,
90,
device=device,
dtype=precision)
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(batch_size,
text_seq_len + 1,
4096,
device=device,
dtype=precision)
# Timestep
timestep = torch.randint(0, 1000, (batch_size, 21), device=device, dtype=torch.long)
logger.info("timestep: %s", timestep)
forward_batch = ForwardBatch(
data_type="dummy",
)
# with torch.amp.autocast('cuda', dtype=precision):
output1 = model1(
x=hidden_states,
context=encoder_hidden_states,
t=timestep,
seq_len=seq_len,
)
with set_forward_context(
current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch,
):
output2 = model2(hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
+79 -12
View File
@@ -40,7 +40,7 @@ from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
EMA_FSDP, clip_grad_norm_while_handling_failing_dtensor_cases,
get_scheduler, load_distillation_checkpoint, save_distillation_checkpoint,
shift_timestep)
shift_timestep, compute_density_for_timestep_sampling, get_sigmas)
from fastvideo.utils import (is_vsa_available, maybe_download_model,
set_random_seed, verify_model_config_and_directory)
@@ -91,6 +91,9 @@ class DistillationPipeline(TrainingPipeline):
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
shift=self.timestep_shift)
self.transformer_2 = self.get_module("transformer_2", None)
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
if training_args.real_score_model_path:
logger.info(
f"Loading real score transformer from: {training_args.real_score_model_path}"
@@ -127,6 +130,39 @@ class DistillationPipeline(TrainingPipeline):
self.real_score_transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
if self.transformer_2 is not None:
self.transformer_2 = apply_activation_checkpointing(
self.transformer_2,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
if self.transformer_2 is not None:
self.transformer_2.train()
self.transformer_2.requires_grad_(True)
params_to_optimize_2 = self.transformer_2.parameters()
params_to_optimize_2 = list(
filter(lambda p: p.requires_grad, params_to_optimize_2))
betas_str = training_args.betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.optimizer_2 = torch.optim.AdamW(
params_to_optimize_2,
lr=training_args.learning_rate,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
self.lr_scheduler_2 = get_scheduler(
training_args.lr_scheduler,
optimizer=self.optimizer_2,
num_warmup_steps=training_args.lr_warmup_steps,
num_training_steps=training_args.max_train_steps,
num_cycles=training_args.lr_num_cycles,
power=training_args.lr_power,
min_lr_ratio=training_args.min_lr_ratio,
last_epoch=self.init_steps - 1,
)
# Initialize optimizers
fake_score_params = list(
@@ -281,6 +317,9 @@ class DistillationPipeline(TrainingPipeline):
"""Prepare training environment for distillation."""
self.transformer.requires_grad_(True)
self.transformer.train()
if self.transformer_2 is not None:
self.transformer_2.requires_grad_(True)
self.transformer_2.train()
self.fake_score_transformer.requires_grad_(True)
self.fake_score_transformer.train()
@@ -436,7 +475,9 @@ class DistillationPipeline(TrainingPipeline):
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.conditional_dict,
training_batch)
pred_noise = self.transformer(**training_batch.input_kwargs).permute(
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
pred_noise = current_model(**training_batch.input_kwargs).permute(
0, 2, 1, 3, 4)
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
@@ -478,6 +519,7 @@ class DistillationPipeline(TrainingPipeline):
max_target_idx = len(self.denoising_step_list) - 1
noise_latents = []
noise_latent_index = target_timestep_idx_int - 1
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
if max_target_idx > 0:
# Run student model for all steps before the target timestep
with torch.no_grad():
@@ -489,7 +531,7 @@ class DistillationPipeline(TrainingPipeline):
training_batch_temp = self._build_distill_input_kwargs(
current_noise_latents, current_timestep_tensor,
training_batch.conditional_dict, training_batch)
pred_flow = self.transformer(
pred_flow = current_model(
**training_batch_temp.input_kwargs).permute(
0, 2, 1, 3, 4)
pred_clean = pred_noise_to_pred_video(
@@ -532,7 +574,7 @@ class DistillationPipeline(TrainingPipeline):
training_batch = self._build_distill_input_kwargs(
noisy_input, target_timestep, training_batch.conditional_dict,
training_batch)
pred_noise = self.transformer(**training_batch.input_kwargs).permute(
pred_noise = current_model(**training_batch.input_kwargs).permute(
0, 2, 1, 3, 4)
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
@@ -774,6 +816,9 @@ class DistillationPipeline(TrainingPipeline):
batches.append(batch)
self.optimizer.zero_grad()
# TODO: confirm this
if self.transformer_2 is not None:
self.optimizer_2.zero_grad()
total_dmd_loss = 0.0
dmd_latent_vis_dict = {}
fake_score_latent_vis_dict = {}
@@ -802,15 +847,31 @@ class DistillationPipeline(TrainingPipeline):
attn_metadata=batch_gen.attn_metadata_vsa):
(dmd_loss / gradient_accumulation_steps).backward()
total_dmd_loss += dmd_loss.detach().item()
self._clip_model_grad_norm_(batch_gen, self.transformer)
for param in self.transformer.parameters():
# check if the gradient is not None and not zero
assert param.grad is not None and param.grad.abs().sum() > 0
self.optimizer.step()
self.optimizer.zero_grad(set_to_none=True)
# Only clip gradients for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
self._clip_model_grad_norm_(batch_gen, self.transformer_2)
for param in self.transformer_2.parameters():
# check if the gradient is not None and not zero
assert param.grad is not None and param.grad.abs().sum() > 0
self.optimizer_2.step()
self.optimizer_2.zero_grad(set_to_none=True)
else:
self._clip_model_grad_norm_(batch_gen, self.transformer)
for param in self.transformer.parameters():
# check if the gradient is not None and not zero
assert param.grad is not None and param.grad.abs().sum() > 0
self.optimizer.step()
self.optimizer.zero_grad(set_to_none=True)
if self.generator_ema is not None:
self.generator_ema.update(self.transformer)
# TODO: support EMA for transformer_2?
if self.train_transformer_2 and self.transformer_2 is not None:
# Note: EMA currently only supports the main transformer
# Could be extended to support transformer_2 in the future
pass
else:
self.generator_ema.update(self.transformer)
avg_dmd_loss = torch.tensor(total_dmd_loss /
gradient_accumulation_steps,
@@ -840,7 +901,13 @@ class DistillationPipeline(TrainingPipeline):
assert param.grad is not None and param.grad.abs().sum() > 0
self.fake_score_optimizer.step()
self.fake_score_lr_scheduler.step()
self.lr_scheduler.step()
# Step the appropriate scheduler
if self.train_transformer_2 and self.transformer_2 is not None:
self.lr_scheduler_2.step()
else:
self.lr_scheduler.step()
self.fake_score_optimizer.zero_grad(set_to_none=True)
avg_fake_score_loss = torch.tensor(total_fake_score_loss /
gradient_accumulation_steps,
File diff suppressed because it is too large Load Diff
+123 -41
View File
@@ -1,4 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
import gc
import dataclasses
import math
import os
@@ -69,6 +70,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
train_dataloader: StatefulDataLoader
train_loader_iter: Iterator[dict[str, Any]]
current_epoch: int = 0
train_transformer_2: bool = False
def __init__(
self,
@@ -104,6 +106,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.sp_world_size = self.sp_group.world_size
self.local_rank = world_group.local_rank
self.transformer = self.get_module("transformer")
self.transformer_2 = self.get_module("transformer_2", None)
self.seed = training_args.seed
self.set_schemas()
@@ -116,6 +119,11 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
if self.transformer_2 is not None:
self.transformer_2 = apply_activation_checkpointing(
self.transformer_2,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
noise_scheduler = self.modules["scheduler"]
self.set_trainable()
@@ -147,6 +155,30 @@ class TrainingPipeline(LoRAPipeline, ABC):
min_lr_ratio=training_args.min_lr_ratio,
last_epoch=self.init_steps - 1,
)
if self.transformer_2 is not None:
# Ensure transformer_2 has trainable parameters before creating optimizer
self.transformer_2.train()
self.transformer_2.requires_grad_(True)
params_to_optimize_2 = self.transformer_2.parameters()
params_to_optimize_2 = list(
filter(lambda p: p.requires_grad, params_to_optimize_2))
self.optimizer_2 = torch.optim.AdamW(
params_to_optimize_2,
lr=training_args.learning_rate,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
self.lr_scheduler_2 = get_scheduler(
training_args.lr_scheduler,
optimizer=self.optimizer_2,
num_warmup_steps=training_args.lr_warmup_steps,
num_training_steps=training_args.max_train_steps,
num_cycles=training_args.lr_num_cycles,
power=training_args.lr_power,
min_lr_ratio=training_args.min_lr_ratio,
last_epoch=self.init_steps - 1,
)
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
training_args.data_path,
@@ -161,7 +193,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
seed=self.seed)
self.noise_scheduler = noise_scheduler
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
self.num_update_steps_per_epoch = math.ceil(
len(self.train_dataloader) /
training_args.gradient_accumulation_steps * training_args.sp_size /
@@ -187,9 +219,25 @@ class TrainingPipeline(LoRAPipeline, ABC):
def _prepare_training(self, training_batch: TrainingBatch) -> TrainingBatch:
self.transformer.train()
self.optimizer.zero_grad()
if self.transformer_2 is not None:
self.transformer_2.train()
self.optimizer_2.zero_grad()
training_batch.total_loss = 0.0
return training_batch
def _enable_training(self, model: torch.nn.Module, optimizer: torch.optim.Optimizer) -> None:
"""Enable training mode and gradients for the specified model."""
for param in model.parameters():
param.requires_grad = True
model.train()
optimizer.zero_grad()
def _disable_training(self, model: torch.nn.Module, optimizer: torch.optim.Optimizer) -> None:
"""Disable training mode and gradients for the specified model."""
for param in model.parameters():
param.requires_grad = False
optimizer.zero_grad(set_to_none=True)
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
@@ -233,17 +281,17 @@ class TrainingPipeline(LoRAPipeline, ABC):
generator=self.noise_gen_cuda,
device=latents.device,
dtype=latents.dtype)
u = compute_density_for_timestep_sampling(
weighting_scheme=self.training_args.weighting_scheme,
batch_size=batch_size,
generator=self.noise_random_generator,
logit_mean=self.training_args.logit_mean,
logit_std=self.training_args.logit_std,
mode_scale=self.training_args.mode_scale,
)
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
timesteps = self.noise_scheduler.timesteps[indices].to(
device=latents.device)
timesteps = self._sample_timesteps(batch_size, latents.device)
# Enable training for the model that will be trained next and disable the other
if self.train_transformer_2:
self._enable_training(self.transformer_2, self.optimizer_2)
self._disable_training(self.transformer, self.optimizer)
else:
self._enable_training(self.transformer, self.optimizer)
if self.transformer_2 is not None:
self._disable_training(self.transformer_2, self.optimizer_2)
if self.training_args.sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
sp_group = get_sp_group()
@@ -266,6 +314,38 @@ class TrainingPipeline(LoRAPipeline, ABC):
return training_batch
def _sample_timesteps(self, batch_size, device):
# Determine which model to train based on the boundary timestep
if (self.transformer_2 is not None and self.boundary_timestep is not None and
torch.rand(1, generator=self.noise_random_generator).item() <= self.training_args.boundary_ratio):
self.train_transformer_2 = True
else:
self.train_transformer_2 = False
# Broadcast the decision to all processes
decision = torch.tensor(1.0 if self.train_transformer_2 else 0.0, device=self.device)
dist.broadcast(decision, src=0)
self.train_transformer_2 = decision.item() == 1.0
# Sample u from the appropriate range
u = compute_density_for_timestep_sampling(
weighting_scheme=self.training_args.weighting_scheme,
batch_size=batch_size,
generator=self.noise_random_generator,
logit_mean=self.training_args.logit_mean,
logit_std=self.training_args.logit_std,
mode_scale=self.training_args.mode_scale,
)
boundary_ratio = self.training_args.boundary_ratio
if self.train_transformer_2:
u = (1 - boundary_ratio) + u * boundary_ratio # min: 1 - boundary_ratio, max: 1
else:
u = u * (1 - boundary_ratio) # min: 0, max: 1 - boundary_ratio
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
return self.noise_scheduler.timesteps[indices].to(device=device)
def _build_attention_metadata(
self, training_batch: TrainingBatch) -> TrainingBatch:
latents_shape = training_batch.raw_latent_shape
@@ -281,20 +361,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
patch_size=patch_size,
VSA_sparsity=current_vsa_sparsity,
device=get_local_torch_device())
# elif vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
# moba_params = self.training_args.moba_config.copy()
# moba_params.update({
# "current_timestep":
# training_batch.timesteps,
# "raw_latent_shape":
# training_batch.raw_latent_shape[2:5],
# "patch_size":
# self.training_args.pipeline_config.dit_config.patch_size,
# "device":
# get_local_torch_device(),
# })
# training_batch.attn_metadata = VideoMobaAttentionMetadataBuilder(
# ).build(**moba_params)
else:
training_batch.attn_metadata = None
@@ -319,7 +385,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
def _transformer_forward_and_compute_loss(
self, training_batch: TrainingBatch) -> TrainingBatch:
# if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN" or vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
assert training_batch.attn_metadata is not None
else:
@@ -331,11 +396,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
# [1000.0],
# device=training_batch.noisy_model_input.device,
# dtype=torch.bfloat16)
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
with set_forward_context(
current_timestep=training_batch.current_timestep,
attn_metadata=training_batch.attn_metadata):
model_pred = self.transformer(**input_kwargs)
model_pred = current_model(**input_kwargs)
if self.training_args.precondition_outputs:
assert training_batch.sigmas is not None
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
@@ -366,7 +432,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
# the following:
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
if max_grad_norm is not None:
model_parts = [self.transformer]
# Only clip gradients for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
model_parts = [self.transformer_2]
else:
model_parts = [self.transformer]
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for m in model_parts for p in m.parameters()],
max_grad_norm,
@@ -411,9 +482,14 @@ class TrainingPipeline(LoRAPipeline, ABC):
training_batch = self._clip_grad_norm(training_batch)
self.optimizer.step()
self.lr_scheduler.step()
# Only step the optimizer and scheduler for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
self.optimizer_2.step()
self.lr_scheduler_2.step()
else:
self.optimizer.step()
self.lr_scheduler.step()
training_batch.total_loss = training_batch.total_loss
training_batch.grad_norm = training_batch.grad_norm
return training_batch
@@ -445,6 +521,11 @@ class TrainingPipeline(LoRAPipeline, ABC):
logger.info("Starting training with %s B trainable parameters",
round(num_trainable_params / 1e9, 3))
if getattr(self, "transformer_2", None) is not None:
num_trainable_params = _get_trainable_params(self.transformer_2)
logger.info("Transformer 2: Starting training with %s B trainable parameters",
round(num_trainable_params / 1e9, 3))
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
@@ -465,7 +546,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
self._log_training_info()
self._log_validation(self.transformer, self.training_args,
self._log_validation(self.training_args,
self.init_steps)
# Train!
@@ -486,9 +567,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
current_decay_times = min(step // vsa_decay_interval_steps,
vsa_sparsity // vsa_decay_rate)
current_vsa_sparsity = current_decay_times * vsa_decay_rate
# elif vmoba_available:
# # TODO: add vmoba sparsity scheduling here
# pass
else:
current_vsa_sparsity = 0.0
@@ -530,7 +608,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.transformer.train()
self.sp_group.barrier()
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
self._log_validation(self.transformer, self.training_args, step)
self._log_validation(self.training_args, step)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
trainable_params = round(
_get_trainable_params(self.transformer) / 1e9, 3)
@@ -611,12 +689,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
return batch
@torch.no_grad()
def _log_validation(self, transformer, training_args, global_step) -> None:
def _log_validation(self, training_args, global_step) -> None:
"""
Generate a validation video and log it to wandb to check the quality during training.
"""
training_args.inference_mode = True
training_args.dit_cpu_offload = True
training_args.dit_cpu_offload = False
if not training_args.log_validation:
return
if self.validation_pipeline is None:
@@ -638,7 +716,9 @@ class TrainingPipeline(LoRAPipeline, ABC):
batch_size=None,
num_workers=0)
transformer.eval()
self.transformer.eval()
if getattr(self, "transformer_2", None) is not None:
self.transformer_2.eval()
validation_steps = training_args.validation_sampling_steps.split(",")
validation_steps = [int(step) for step in validation_steps]
@@ -730,4 +810,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Re-enable gradients for training
training_args.inference_mode = False
transformer.train()
self.transformer.train()
if getattr(self, "transformer_2", None) is not None:
self.transformer_2.train()
+57
View File
@@ -0,0 +1,57 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export WANDB_API_KEY='2f25ad37933894dbf0966c838c0b8494987f9f2f'
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
export TRITON_CACHE_DIR=/tmp/triton_cache
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/
DATA_DIR=/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol_processed_t2v/combined_parquet_dataset
# VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/mixkit/validation_8.json
VALIDATION_DIR=/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol-single_processed_t2v/validation.json
NUM_GPUS=8
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/training/wan_training_pipeline.py \
--model_path Wan-AI/Wan2.2-T2V-A14B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.2-T2V-A14B-Diffusers \
--cache_dir "/home/ray/.cache" \
--data_path "$DATA_DIR" \
--validation_dataset_file "$VALIDATION_DIR" \
--train_batch_size 1 \
--num_latent_t 16\
--sp_size 4 \
--tp_size 1 \
--num_gpus $NUM_GPUS \
--hsdp_replicate_dim 1 \
--hsdp-shard-dim 8 \
--train_sp_batch_size 1 \
--dataloader_num_workers 4 \
--gradient_accumulation_steps 1 \
--max_train_steps 30000 \
--learning_rate 1e-5 \
--mixed_precision "bf16" \
--checkpointing_steps 1000 \
--validation_steps 30 \
--validation_sampling_steps "40" \
--log_validation True \
--checkpoints_total_limit 3 \
--ema_start_step 0 \
--training_cfg_rate 0.1 \
--seed 1024 \
--output_dir "outputs_train_test/wan_finetune_v1" \
--tracker_project_name VSA_finetune \
--num_height 448 \
--num_width 832 \
--num_frames 61 \
--flow_shift 5 \
--validation_guidance_scale "5.0" \
--num_euler_timesteps 50 \
--master_weight_type "fp32" \
--dit_precision "fp32" \
--weight_decay 0.01 \
--max_grad_norm 1.0