Compare commits

..
Author SHA1 Message Date
JerryZhou54 32711250ad checkpoint 2025-09-09 02:48:02 +00:00
JerryZhou54 3d6cac57fc Fix runtime errors during training 2025-09-09 01:15:18 +00:00
Matthew Noto 57fd3d8159 new branch 2025-09-09 01:15:18 +00:00
JerryZhou54 1eacdd80de new branch 2025-09-09 01:15:18 +00:00
Matthew Noto 2ca24b3288 new branch 2025-09-08 19:20:32 +00:00
9 changed files with 219 additions and 53 deletions
@@ -0,0 +1,139 @@
# 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=29500
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# Configs
NUM_GPUS=8
# Model paths for DMD distillation:
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/SFWan2.1-I2V/validation_better.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd
--output_dir "/mnt/sharefs/users/hao.zhang/wl/sf_checkpoints/ode0_SFwan_t2v_finetune_sf_${lr}_c${critic_lr}"
--wandb_run_name "DEBUG${lr}_c${critic_lr}"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81
--warp_denoising_step
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frame_per_block 3
--enable_gradient_masking
--gradient_mask_last_n_frames 21
)
# Parallel arguments
parallel_args=(
--num_gpus 8 # 64
--sp_size 1
--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 "3"
--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 50
--weight_only_checkpointing_steps 50
--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'
)
# 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
--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 \
--nproc_per_node $NUM_GPUS \
--master_port $MASTER_PORT \
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[@]}"
+2 -2
View File
@@ -57,7 +57,7 @@ class WanT2V480PConfig(PipelineConfig):
# WanConfig-specific added parameters
def __post_init__(self):
self.vae_config.load_encoder = True
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@@ -148,4 +148,4 @@ class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
is_causal: bool = True
flow_shift: int = 5
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 750, 500, 250])
default_factory=lambda: [1000, 750, 500, 250])
+10 -12
View File
@@ -147,10 +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"].clone()
kv_cache["v"] = kv_cache["v"].clone()
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
x = self.attn(
@@ -248,27 +246,27 @@ class CausalWanTransformerBlock(nn.Module):
current_start: int = 0,
cache_start: int | None = None,
) -> torch.Tensor:
logger.info("temb.shape: %s", temb.shape)
# logger.info("temb.shape: %s", temb.shape)
num_frames = temb.shape[1]
logger.info("first hidden_states.shape: %s", hidden_states.shape)
logger.info("num_frames: %s", num_frames)
# logger.info("first hidden_states.shape: %s", hidden_states.shape)
# logger.info("num_frames: %s", num_frames)
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
frame_seqlen = hidden_states.shape[1] // temb.shape[1]
logger.info("frame_seqlen: %s", frame_seqlen)
# logger.info("frame_seqlen: %s", frame_seqlen)
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
logger.info("e.shape: %s", e.shape)
# logger.info("e.shape: %s", e.shape)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=2)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
logger.info("hidden_states.shape: %s", hidden_states.shape)
logger.info("scale_msa.shape: %s", scale_msa.shape)
logger.info("shift_msa.shape: %s", shift_msa.shape)
# logger.info("hidden_states.shape: %s", hidden_states.shape)
# logger.info("scale_msa.shape: %s", scale_msa.shape)
# logger.info("shift_msa.shape: %s", shift_msa.shape)
norm_hidden_states_unflattened = self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen))
# logger.info("norm_hidden_states_unflattened.shape: %s", norm_hidden_states_unflattened.shape)
@@ -320,7 +318,7 @@ class CausalWanTransformerBlock(nn.Module):
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
# logger.info("after mlp_residual norm_hidden_states.shape: %s", norm_hidden_states.shape)
logger.info("after mlp_residual hidden_states.shape: %s", hidden_states.shape)
# logger.info("after mlp_residual hidden_states.shape: %s", hidden_states.shape)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
@@ -517,7 +515,7 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# logger.info("forward inference flattened and transposed hidden_states.shape: %s", hidden_states.shape)
logger.info("timestep shape: %s", timestep.shape)
# logger.info("timestep shape: %s", timestep.shape)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
@@ -635,15 +635,32 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
noise: torch.Tensor,
timestep: torch.IntTensor,
) -> torch.Tensor:
"""
Args:
clean_latent: the clean latent with shape [B, C, H, W],
where B is batch_size or batch_size * num_frames
noise: the noise with shape [B, C, H, W]
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
Returns:
the corrupted latent with shape [B, C, H, W]
"""
# If timestep is [bs, num_frames]
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
assert timestep.numel() == clean_latent.shape[0]
elif timestep.ndim == 1:
# If timestep is [1]
if timestep.shape[0] == 1:
timestep = timestep.expand(clean_latent.shape[0])
else:
assert timestep.numel() == clean_latent.shape[0]
else:
raise ValueError(f"[add_noise] Invalid timestep shape: {timestep.shape}")
# timestep shape should be [B]
self.sigmas = self.sigmas.to(noise.device)
# TODO: hack
if timestep.ndim == 2 and timestep.shape[1] == 1:
timestep = timestep.expand(clean_latent.shape[0])
self.timesteps = self.timesteps.to(noise.device)
logger.info("self.timesteps shape: %s", self.timesteps.shape)
logger.info("timestep shape: %s", timestep.shape)
if timestep.ndim > 1:
timestep = timestep.squeeze(0)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
@@ -656,4 +673,4 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
return sample
def __len__(self) -> int:
return self.config.num_train_timesteps
return self.config.num_train_timesteps
+23 -4
View File
@@ -148,11 +148,30 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
scheduler: Any) -> torch.Tensor:
"""
Convert predicted noise to clean latent.
Args:
pred_noise: the predicted noise with shape [B, C, H, W]
where B is batch_size or batch_size * num_frames
noise_input_latent: the noisy latent with shape [B, C, H, W],
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
scheduler: the scheduler
Returns:
the predicted video with shape [B, C, H, W]
"""
logger.info(f"timestep: {timestep.shape}")
logger.info(f"noise_input_latent: {noise_input_latent.shape}")
logger.info(f"pred_noise: {pred_noise.shape}")
timestep = timestep.expand(noise_input_latent.shape[0])
# If timestep is [bs, num_frames]
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
assert timestep.numel() == noise_input_latent.shape[0]
elif timestep.ndim == 1:
# If timestep is [1]
if timestep.shape[0] == 1:
timestep = timestep.expand(noise_input_latent.shape[0])
else:
assert timestep.numel() == noise_input_latent.shape[0]
else:
raise ValueError(f"[pred_noise_to_pred_video] Invalid timestep shape: {timestep.shape}")
# timestep shape should be [B]
dtype = pred_noise.dtype
device = pred_noise.device
pred_noise = pred_noise.float().to(device)
-3
View File
@@ -27,7 +27,6 @@ from fastvideo.layers.activation import get_act_fn
from fastvideo.models.vaes.common import (DiagonalGaussianDistribution,
ParallelTiledVAE)
from fastvideo.platforms import current_platform
from fastvideo.logger import init_logger
CACHE_T = 2
@@ -36,7 +35,6 @@ feat_cache = contextvars.ContextVar("feat_cache", default=None)
feat_idx = contextvars.ContextVar("feat_idx", default=0)
first_chunk = contextvars.ContextVar("first_chunk", default=None)
logger = init_logger(__name__)
@contextmanager
def forward_context(first_frame_arg=False,
@@ -1131,7 +1129,6 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
self._conv_idx = 0
self._feat_map = [None] * self._conv_num
# cache encode
logger.info("self.config.load_encoder: %s", self.config.load_encoder)
if self.config.load_encoder:
self._enc_conv_num = _count_conv3d(self.encoder)
self._enc_conv_idx = 0
+7 -5
View File
@@ -12,6 +12,7 @@ from typing import Any
import imageio
import numpy as np
import torch
import torch.distributed as dist
import torch.nn.functional as F
import torchvision
from einops import rearrange
@@ -728,16 +729,17 @@ class DistillationPipeline(TrainingPipeline):
"encoder_hidden_states": training_batch.encoder_hidden_states,
"encoder_attention_mask": training_batch.encoder_attention_mask,
}
unconditional_dict = {
"encoder_hidden_states": self.negative_prompt_embeds,
"encoder_attention_mask": self.negative_prompt_attention_mask,
}
if getattr(self, "negative_prompt_embeds", None) is not None:
unconditional_dict = {
"encoder_hidden_states": self.negative_prompt_embeds,
"encoder_attention_mask": self.negative_prompt_attention_mask,
}
training_batch.unconditional_dict = unconditional_dict
training_batch.dmd_latent_vis_dict = {}
training_batch.fake_score_latent_vis_dict = {}
training_batch.conditional_dict = conditional_dict
training_batch.unconditional_dict = unconditional_dict
training_batch.raw_latent_shape = training_batch.latents.shape
training_batch.latents = training_batch.latents.permute(0, 2, 1, 3, 4)
self.video_latent_shape = training_batch.latents.shape
@@ -52,14 +52,17 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
super().initialize_training_pipeline(training_args)
self.dfake_gen_update_ratio = getattr(training_args, 'dfake_gen_update_ratio', 5)
# Self-forcing specific properties
self.num_frame_per_block = getattr(training_args, 'num_frame_per_block', 3)
self.independent_first_frame = getattr(training_args, 'independent_first_frame', False)
self.same_step_across_blocks = getattr(training_args, 'same_step_across_blocks', False)
self.last_step_only = getattr(training_args, 'last_step_only', False)
self.context_noise = getattr(training_args, 'context_noise', 0)
# Calculate frame sequence length - this will be set properly in _prepare_dit_inputs
self.frame_seq_length = 1560 # TODO: Calculate this dynamically based on patch size
# Cache references (will be initialized per forward pass)
self.kv_cache1 = None
self.crossattn_cache = None
@@ -155,9 +158,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
noisy_latent = self.noise_scheduler.add_noise(latents.flatten(0, 1),
noise.flatten(0, 1),
timestep).unflatten(
0,
(1, latents.shape[1]))
torch.tensor([timestep], device=noise.device))
# Step 4: Build input kwargs with KV cache support
training_batch = self._build_distill_input_kwargs(
@@ -186,7 +187,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
noise_input_latent=noisy_latent.flatten(0, 1),
timestep=timestep,
timestep=torch.tensor([timestep], device=noisy_latent.device),
scheduler=self.noise_scheduler).unflatten(0, pred_noise.shape[:2])
self._reset_simulation_caches(kv_cache, crossattn_cache)
@@ -206,7 +207,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
initial_latent = getattr(training_batch, 'image_latent', None)
# Dynamic frame generation logic (adapted from _run_generator)
num_training_frames = getattr(self.training_args, 'num_frames', 21)
num_training_frames = getattr(self.training_args, 'num_latent_t', 21)
# During training, the number of generated frames should be uniformly sampled from
# [21, self.num_training_frames], but still being a multiple of self.num_frame_per_block
@@ -287,7 +288,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
all_num_frames = [1] + all_num_frames
num_denoising_steps = len(self.denoising_step_list)
exit_flags = self.generate_and_sync_list(len(all_num_frames), num_denoising_steps, device=noise.device)
start_gradient_frame_index = num_output_frames - 21
start_gradient_frame_index = max(0, num_output_frames - 21)
for block_index, current_num_frames in enumerate(all_num_frames):
noisy_input = noise[:, current_start_frame - num_input_frames:current_start_frame + current_num_frames - num_input_frames]
@@ -300,7 +301,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
exit_flag = (index == exit_flags[block_index])
timestep = torch.ones([batch_size, current_num_frames], device=noise.device, dtype=torch.int64) * current_timestep
logger.info("timestep shape at initalization: %s", timestep.shape)
if not exit_flag:
with torch.no_grad():
@@ -308,7 +308,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
training_batch_temp = self._build_distill_input_kwargs(
noisy_input, timestep, training_batch.conditional_dict, training_batch)
logger.info("timestep shape in generator: %s", timestep.shape)
pred_flow = self.transformer(
hidden_states=training_batch_temp.input_kwargs['hidden_states'],
encoder_hidden_states=training_batch_temp.input_kwargs['encoder_hidden_states'],
@@ -322,21 +321,18 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
denoised_pred = pred_noise_to_pred_video(
pred_noise=pred_flow.flatten(0, 1),
noise_input_latent=noisy_input.flatten(0, 1),
timestep=timestep.flatten(),
timestep=timestep,
scheduler=self.noise_scheduler).unflatten(0, pred_flow.shape[:2])
next_timestep = self.denoising_step_list[index + 1]
logger.info("denoised_pred shape before: %s", denoised_pred.shape)
noisy_input = self.noise_scheduler.add_noise(
denoised_pred.flatten(0, 1),
torch.randn_like(denoised_pred.flatten(0, 1)),
next_timestep * torch.ones([batch_size * current_num_frames], device=noise.device, dtype=torch.long)
).unflatten(0, denoised_pred.shape[:2])
logger.info("denoised_pred shape after: %s", denoised_pred.shape)
else:
# Final prediction with gradient control
if current_start_frame < start_gradient_frame_index:
logger.info("timestep shape in generator: %s", timestep.shape)
with torch.no_grad():
training_batch_temp = self._build_distill_input_kwargs(
noisy_input, timestep, training_batch.conditional_dict, training_batch)
@@ -351,7 +347,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
current_start=current_start_frame * self.frame_seq_length
).permute(0, 2, 1, 3, 4)
else:
logger.info("timestep shape in generator: %s", timestep.shape)
training_batch_temp = self._build_distill_input_kwargs(
noisy_input, timestep, training_batch.conditional_dict, training_batch)
@@ -368,7 +363,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
denoised_pred = pred_noise_to_pred_video(
pred_noise=pred_flow.flatten(0, 1),
noise_input_latent=noisy_input.flatten(0, 1),
timestep=timestep.flatten(),
timestep=timestep,
scheduler=self.noise_scheduler).unflatten(0, pred_flow.shape[:2])
break
@@ -377,13 +372,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
# Step 3.3: rerun with timestep zero to update the cache
context_timestep = torch.ones_like(timestep) * self.context_noise
logger.info("denoised_pred shape before: %s", denoised_pred.shape)
denoised_pred = self.noise_scheduler.add_noise(
denoised_pred.flatten(0, 1),
torch.randn_like(denoised_pred.flatten(0, 1)),
context_timestep * torch.ones([batch_size * current_num_frames], device=noise.device, dtype=torch.long)
context_timestep
).unflatten(0, denoised_pred.shape[:2])
logger.info("denoised_pred shape after: %s", denoised_pred.shape)
with torch.no_grad():
training_batch_temp = self._build_distill_input_kwargs(
@@ -433,7 +426,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
frame = pixels[:, :, -1:, :, :].to(dtype) # Last frame [B, C, 1, H, W]
# Encode frame back to get image latent
image_latent = self.vae.encode(frame).mean.to(dtype)
image_latent = self.vae.encode(frame).to(dtype)
image_latent = image_latent.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
pred_image_or_video_last_21 = torch.cat([image_latent, pred_image_or_video[:, -20:, ...]], dim=1)
+2 -1
View File
@@ -2,7 +2,8 @@
counter=0
# for pair in "1e-5 1e-5" "1e-5 8e-6" "1e-5 6e-6" "1e-5 4e-6" "1e-5 2e-6" "1e-5 1e-6"; do
for pair in "1e-5 8e-6" "1e-5 4e-6" "1e-5 2e-6" "8e-6 8e-6" "8e-6 6e-6" "8e-6 2e-6"; do
# for pair in "1e-5 8e-6" "1e-5 4e-6" "1e-5 2e-6" "8e-6 8e-6" "8e-6 6e-6" "8e-6 2e-6"; do
for pair in "1e-5 8e-6"; do
# for pair in "2e-6 4e-7" "2e-6 6e-7" "4e-6 6e-7" "4e-6 8e-7" "4e-6 1e-6" "6e-6 4e-7" "6e-6 6e-7" "6e-6 8e-7" "6e-6 1e-6"; do
# for pair in "2e-6 4e-7" "2e-6 6e-7" "4e-6 6e-7" "4e-6 8e-7" "4e-6 1e-6" "6e-6 4e-7" "6e-6 6e-7" "6e-6 8e-7" "6e-6 1e-6"; do
port=$((29500 + counter))