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
12 changed files with 268 additions and 30 deletions
+93 -1
View File
@@ -33,7 +33,8 @@ class ParquetVideoTextDataset(Dataset):
world_size: int = 1,
cfg_rate: float = 0.0,
num_latent_t: int = 2,
seed: int = 0):
seed: int = 0,
validation: bool = False):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
@@ -47,6 +48,12 @@ class ParquetVideoTextDataset(Dataset):
self.cfg_rate = cfg_rate
self.num_latent_t = num_latent_t
self.local_indices = None
self.validation = validation
# Negative prompt caching
self.neg_metadata = None
self.cached_neg_prompt: Dict[str, Any] | None = None
self.plan_output_dir = os.path.join(
self.path,
f"data_plan_{self.world_size}_{self.sp_world_size}_{self.dp_world_size}.json"
@@ -75,6 +82,12 @@ class ParquetVideoTextDataset(Dataset):
for row_idx in range(num_rows):
metadatas.append((file_path, row_idx))
# the negative prompt is always the first row in the first
# parquet file
if validation:
self.neg_metadata = metadatas[0]
metadatas = metadatas[1:]
# Generate the plan that distribute rows among workers
random.seed(seed)
random.shuffle(metadatas)
@@ -93,9 +106,88 @@ class ParquetVideoTextDataset(Dataset):
for global_rank in group_ranks_list[sp_group_idx]:
plan[global_rank].append(metadata)
if validation:
assert self.neg_metadata is not None
plan["negative_prompt"] = [self.neg_metadata]
with open(self.plan_output_dir, "w") as f:
json.dump(plan, f)
else:
pass
dist.barrier()
if validation:
with open(self.plan_output_dir) as f:
plan = json.load(f)
self.neg_metadata = plan["negative_prompt"][0]
def _load_and_cache_negative_prompt(self) -> None:
"""Load and cache the negative prompt. Only rank 0 in each SP group should call this."""
if not self.validation or self.neg_metadata is None:
return
if self.cached_neg_prompt is not None:
return
# Only rank 0 in each SP group should read the negative prompt
try:
file_path, row_idx = self.neg_metadata
parquet_file = pq.ParquetFile(file_path)
# Since negative prompt is always the first row (row_idx = 0),
# it's always in the first row group
row_group_index = 0
local_index = row_idx # This will be 0 for the negative prompt
row_group = parquet_file.read_row_group(row_group_index).to_pydict()
row_dict = {k: v[local_index] for k, v in row_group.items()}
del row_group
# Process the negative prompt row
self.cached_neg_prompt = self._process_row(row_dict)
except Exception as e:
logger.error("Failed to load negative prompt: %s", e)
self.cached_neg_prompt = None
def get_validation_negative_prompt(
self
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, Dict[str, Any]]:
"""
Get the negative prompt for validation.
This method ensures the negative prompt is loaded and cached properly.
Returns the processed negative prompt data (latents, embeddings, masks, info).
"""
if not self.validation:
raise ValueError(
"get_validation_negative_prompt() can only be called in validation mode"
)
# Load and cache if needed (only rank 0 in SP group will actually load)
if self.cached_neg_prompt is None:
self._load_and_cache_negative_prompt()
if self.cached_neg_prompt is None:
raise RuntimeError(
f"Rank {self.rank} (SP rank {self.local_rank}): Could not retrieve negative prompt data"
)
# Extract the components
lat, emb, mask, info = (self.cached_neg_prompt["latents"],
self.cached_neg_prompt["embeddings"],
self.cached_neg_prompt["masks"],
self.cached_neg_prompt["info"])
# Apply the same processing as in __getitem__
if lat.numel() == 0: # Validation parquet
return lat, emb, mask, info
else:
lat = lat[:, -self.num_latent_t:]
if self.sp_world_size > 1:
lat = rearrange(lat,
"t (n s) h w -> t n s h w",
n=self.sp_world_size).contiguous()
lat = lat[:, self.local_rank, :, :, :]
return lat, emb, mask, info
def __len__(self):
if self.local_indices is None:
+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)
@@ -161,6 +161,7 @@ class ComposedPipelineBase(ABC):
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
fastvideo_args.num_gpus = int(os.environ.get("WORLD_SIZE", 1))
fastvideo_args.use_cpu_offload = False
# make sure we are in training mode
fastvideo_args.inference_mode = False
@@ -126,4 +126,5 @@ class ForwardBatch:
# Set do_classifier_free_guidance based on guidance scale and negative prompt
if self.guidance_scale > 1.0:
self.do_classifier_free_guidance = True
if self.negative_prompt_embeds is None:
self.negative_prompt_embeds = []
@@ -12,6 +12,7 @@ from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset import getdataset
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
@@ -300,7 +301,10 @@ class BasePreprocessPipeline(ComposedPipelineBase):
# Prepare batch data for Parquet dataset
batch_data = []
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
if sampling_param.negative_prompt:
prompts = [sampling_param.negative_prompt] + prompts
# Add progress bar for validation text preprocessing
pbar = tqdm(enumerate(prompts),
desc="Processing validation prompts",
+11 -8
View File
@@ -47,14 +47,15 @@ class DenoisingStage(PipelineStage):
super().__init__()
self.transformer = transformer
self.scheduler = scheduler
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA) # hack
)
if transformer is not None:
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA) # hack
)
def forward(
self,
@@ -71,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,
@@ -30,6 +30,7 @@ class WanPipeline(LoRAPipeline, ComposedPipelineBase):
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# We use UniPCMScheduler from Wan2.1 official repo, not the one in diffusers.
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
+18 -5
View File
@@ -178,7 +178,11 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
rank=self.rank,
world_size=self.world_size,
cfg_rate=training_args.cfg,
num_latent_t=training_args.num_latent_t)
num_latent_t=training_args.num_latent_t,
validation=True)
if sampling_param.negative_prompt:
_, negative_prompt_embeds, negative_prompt_attention_mask, _ = validation_dataset.get_validation_negative_prompt(
)
validation_dataloader = StatefulDataLoader(
validation_dataset,
@@ -194,6 +198,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
# Add the transformer to the validation pipeline
self.validation_pipeline.add_module("transformer", transformer)
# TODO(Peiyuan): those logic should be inside add_module
self.validation_pipeline.latent_preparation_stage.transformer = transformer # type: ignore[attr-defined]
self.validation_pipeline.denoising_stage.transformer = transformer # type: ignore[attr-defined]
@@ -203,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)
@@ -216,6 +223,8 @@ 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",
@@ -223,22 +232,26 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
seed=validation_seed, # Use deterministic seed
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
# make sure we use the same height, width, and num_frames as the training pipeline
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
# TODO(will): validation_sampling_steps and
# validation_guidance_scale are actually passed in as a list of
# values, like "10,20,30". The validation should be run for each
# combination of values.
# num_inference_steps=fastvideo_args.validation_sampling_steps,
num_inference_steps=sampling_param.num_inference_steps,
# guidance_scale=fastvideo_args.validation_guidance_scale,
guidance_scale=1,
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
do_classifier_free_guidance=False,
eta=0.0,
)
# Run validation inference
with torch.inference_mode(), torch.autocast("cuda",
dtype=torch.bfloat16):
with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
@@ -272,6 +272,9 @@ class WanTrainingPipeline(TrainingPipeline):
logger.info("GPU memory usage before train_one_step: %s MB",
gpu_memory_usage)
# 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