Compare commits
13
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8db5dff76f | ||
|
|
4388fa043d | ||
|
|
d6ef6c6ae4 | ||
|
|
da485fbe40 | ||
|
|
f657eb40dc | ||
|
|
19d75b9af3 | ||
|
|
2ecdc2bb8d | ||
|
|
43cb9075f2 | ||
|
|
2768c94977 | ||
|
|
2c35841a39 | ||
|
|
0f0285d1ee | ||
|
|
9f6b0ddc27 | ||
|
|
bb96fa2003 |
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
@@ -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 \
|
||||
|
||||
Reference in New Issue
Block a user