Compare commits

..
Author SHA1 Message Date
Will Lin 8db5dff76f exp 2025-06-06 11:01:14 -07:00
Will Lin 4388fa043d exp 2025-06-05 22:50:54 -07:00
Will Lin d6ef6c6ae4 Revert "Revert "[STA] Implement mask search for V1's Wan2.1 (#415)""
This reverts commit f657eb40dc.
2025-06-05 16:08:57 -07:00
Will Lin da485fbe40 fix 2025-06-05 16:04:14 -07:00
Will Lin f657eb40dc Revert "[STA] Implement mask search for V1's Wan2.1 (#415)"
This reverts commit e3d0cbe185.
2025-06-05 14:34:53 -07:00
Will Lin 19d75b9af3 add todo 2025-06-05 13:43:05 -07:00
Will Lin 2ecdc2bb8d cleanup 2025-06-05 13:43:05 -07:00
Will Lin 43cb9075f2 clean up 2025-06-05 13:43:04 -07:00
Will Lin 2768c94977 clean up 2025-06-05 13:43:04 -07:00
Will Lin 2c35841a39 update 2025-06-05 13:43:04 -07:00
Will Lin 0f0285d1ee update 2025-06-05 13:43:04 -07:00
“BrianChen1129” 9f6b0ddc27 update 2025-06-05 13:43:03 -07:00
“BrianChen1129” bb96fa2003 misc 2025-06-05 13:43:03 -07:00
10 changed files with 150 additions and 29 deletions
+25 -12
View File
@@ -9,6 +9,19 @@ import torch.nn as nn
from fastvideo.v1.layers.custom_op import CustomOp
class FP32LayerNorm(nn.LayerNorm):
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
origin_dtype = inputs.dtype
return torch.nn.functional.layer_norm(
inputs.float(),
self.normalized_shape,
self.weight.float() if self.weight is not None else None,
self.bias.float() if self.bias is not None else None,
self.eps,
).to(origin_dtype)
@CustomOp.register("rms_norm")
class RMSNorm(CustomOp):
"""Root mean square normalization.
@@ -121,10 +134,9 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
eps=eps,
dtype=dtype)
elif norm_type == "layer":
self.norm = nn.LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps,
dtype=dtype)
self.norm = FP32LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps)
else:
raise NotImplementedError(f"Norm type {norm_type} not implemented")
@@ -144,9 +156,11 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
# Apply residual connection with gating
residual_output = residual + x * gate
# Apply normalization
normalized = self.norm(residual_output)
normalized = self.norm(residual_output.float()).to(
residual_output.dtype)
# Apply scale and shift
modulated = normalized * (1.0 + scale) + shift
modulated = (normalized.float() * (1.0 + scale) + shift).to(
residual_output.dtype)
return modulated, residual_output
@@ -171,15 +185,14 @@ class LayerNormScaleShift(nn.Module):
has_weight=elementwise_affine,
eps=eps)
elif norm_type == "layer":
self.norm = nn.LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps,
dtype=dtype)
self.norm = FP32LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps)
else:
raise NotImplementedError(f"Norm type {norm_type} not implemented")
def forward(self, x: torch.Tensor, shift: torch.Tensor,
scale: torch.Tensor) -> torch.Tensor:
"""Apply ln followed by scale and shift in a single fused operation."""
normalized = self.norm(x)
return normalized * (1.0 + scale) + shift
normalized = self.norm(x.float()).to(x.dtype)
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
+4 -3
View File
@@ -13,8 +13,8 @@ from fastvideo.v1.configs.sample.wan import WanTeaCacheParams
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.forward_context import get_forward_context
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, RMSNorm,
ScaleResidual,
from fastvideo.v1.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.v1.layers.linear import ReplicatedLinear
# from torch.nn import RMSNorm
@@ -229,7 +229,8 @@ class WanTransformerBlock(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
# self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = FP32LayerNorm(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)
@@ -7,8 +7,7 @@ This module defines the dataclasses used to pass state between pipeline componen
in a functional manner, reducing the need for explicit parameter passing.
"""
import pprint
from dataclasses import asdict, dataclass, field
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Union
import torch
@@ -129,6 +128,3 @@ class ForwardBatch:
self.do_classifier_free_guidance = True
if self.negative_prompt_embeds is None:
self.negative_prompt_embeds = []
def __str__(self):
return pprint.pformat(asdict(self), indent=2, width=120)
@@ -36,7 +36,6 @@ class ConditioningStage(PipelineStage):
Returns:
The batch with applied conditioning.
"""
# TODO!!
if not batch.do_classifier_free_guidance:
return batch
else:
+2 -2
View File
@@ -47,8 +47,6 @@ class DenoisingStage(PipelineStage):
super().__init__()
self.transformer = transformer
self.scheduler = scheduler
# when used for validation, transformer is None as it is taking from the
# training loop
if transformer is not None:
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
self.attn_backend = get_attn_backend(
@@ -74,6 +72,8 @@ class DenoisingStage(PipelineStage):
Returns:
The batch with denoised latents.
"""
self.transformer.to(fastvideo_args.device)
# Prepare extra step kwargs for scheduler
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
+4 -2
View File
@@ -208,6 +208,8 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
for _, embeddings, masks, infos in validation_dataloader:
caption = infos['caption']
captions.extend(caption)
print(f"rank {self.rank} is running validation")
print(f"rank {self.rank} file_name: {infos['file_name']}")
prompt_embeds = embeddings.to(training_args.device)
prompt_attention_mask = masks.to(training_args.device)
@@ -221,13 +223,13 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
logger.info(f"rank {self.rank} num_frames: {num_frames}")
# Prepare batch for validation
batch = ForwardBatch(
data_type="video",
latents=None,
seed=validation_seed, # Use deterministic seed
generator=torch.Generator(
device="cpu").manual_seed(validation_seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
@@ -198,7 +198,8 @@ class WanTrainingPipeline(TrainingPipeline):
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
noise_random_generator = torch.Generator(device="cpu").manual_seed(seed)
noise_random_generator = torch.Generator(device="cpu")
noise_random_generator.manual_seed(seed)
logger.info("Initialized random seeds with seed: %s", seed)
@@ -270,7 +271,10 @@ class WanTrainingPipeline(TrainingPipeline):
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info("GPU memory usage before train_one_step: %s MB",
gpu_memory_usage)
self._log_validation(self.transformer, self.training_args, 1)
# Do validation at the beginning of training
# self._log_validation(self.transformer, self.training_args, 0)
for step in range(self.init_steps + 1,
self.training_args.max_train_steps + 1):
start_time = time.perf_counter()
+53
View File
@@ -0,0 +1,53 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
DATA_DIR=data/cats_480_2_latents_parq_neg/combined_parquet_dataset
VALIDATION_DIR=data/cats_480_2_latents_parq_neg/validation_parquet_dataset
NUM_GPUS=4
CUDA_VISIBLE_DEVICES=4,5,6,7
# IP=[MASTER NODE IP]
# 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/v1/training/wan_training_pipeline.py\
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache"\
--data_path "$DATA_DIR"\
--validation_prompt_dir "$VALIDATION_DIR"\
--train_batch_size=1 \
--num_latent_t 20 \
--num_gpus 4 \
--sp_size 4 \
--tp_size 4 \
--dp_size 1 \
--dp_shards 4 \
--train_sp_batch_size 1\
--dataloader_num_workers 1\
--gradient_accumulation_steps=1 \
--max_train_steps=5000 \
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=6000 \
--validation_steps 10\
--validation_sampling_steps "2,4,8" \
--log_validation \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--output_dir="data/wan_finetune_crush"\
--tracker_project_name wan_finetune \
--num_height 480 \
--num_width 832 \
--num_frames 81 \
--shift 3 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--weight_decay 0.01 \
--not_apply_cfg_solver \
--master_weight_type "fp32" \
--max_grad_norm 1.0
+53
View File
@@ -0,0 +1,53 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
DATA_DIR=data/crush-smol_parq/combined_parquet_dataset
VALIDATION_DIR=data/crush-smol_parq/validation_parquet_dataset
NUM_GPUS=1
CUDA_VISIBLE_DEVICES=4,5,6,7
# IP=[MASTER NODE IP]
# 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/v1/training/wan_training_pipeline.py\
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache"\
--data_path "$DATA_DIR"\
--validation_prompt_dir "$VALIDATION_DIR"\
--train_batch_size=1 \
--num_latent_t 14 \
--num_gpus 1 \
--sp_size 1 \
--tp_size 1 \
--dp_size 1 \
--dp_shards 1 \
--train_sp_batch_size 1\
--dataloader_num_workers 1\
--gradient_accumulation_steps=1 \
--max_train_steps=5000 \
--learning_rate=1e-5\
--mixed_precision="bf16"\
--checkpointing_steps=6000 \
--validation_steps 100\
--validation_sampling_steps "2,4,8" \
--log_validation \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--output_dir="data/wan_finetune_crush"\
--tracker_project_name finetrainers-wan \
--num_height 480 \
--num_width 832 \
--num_frames 77 \
--shift 3 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--weight_decay 0.01 \
--not_apply_cfg_solver \
--master_weight_type "fp32" \
--max_grad_norm 1.0
@@ -2,8 +2,8 @@
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol/latents"
DATA_MERGE_PATH="your/path/to/Mixkit-Src/merge.txt"
OUTPUT_DIR="your/path"
VALIDATION_PATH="assets/prompt.txt"
torchrun --nproc_per_node=$GPU_NUM \