Compare commits

...
7 Commits
Author SHA1 Message Date
JerryZhou54 5ae1b1ca91 Update 2025-06-22 17:28:50 +00:00
JerryZhou54 7e6e37b6b1 Rename column name 2025-06-11 22:05:36 +00:00
JerryZhou54 64c9c780df Small change 2025-06-10 23:17:06 +00:00
JerryZhou54 bd7b7078ba Fix preprocess 2025-06-10 23:17:06 +00:00
Wei Zhou 3cde7934ab Update preprocess_pipeline_i2v.py 2025-06-10 23:17:05 +00:00
JerryZhou54 c0147f611e I2V Runnable 2025-06-10 23:17:03 +00:00
“BrianChen1129” f6314572e6 Add Encoded first frame to processed dataset 2025-06-10 23:15:25 +00:00
8 changed files with 784 additions and 21 deletions
@@ -35,6 +35,9 @@ pyarrow_schema_i2v = pa.schema([
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_dtype", pa.string()),
pa.field("encoded_first_frame_bytes", pa.binary()),
pa.field("encoded_first_frame_shape", pa.list_(pa.int64())),
pa.field("encoded_first_frame_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
@@ -30,6 +30,7 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
sp_world_size: int,
global_rank: int,
drop_last: bool = True,
drop_first_row: bool = False,
seed: int = 0,
):
self.batch_size = batch_size
@@ -44,6 +45,9 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
rng = torch.Generator().manual_seed(self.seed)
# Create a random permutation of all indices
global_indices = torch.randperm(self.dataset_size, generator=rng)
if drop_first_row:
# remove 0 from global_indices
global_indices = global_indices[global_indices != 0]
if self.drop_last:
# For drop_last=True, we:
@@ -145,7 +149,10 @@ class LatentsParquetMapStyleDataset(Dataset):
Using parquet for map style dataset is not efficient, we mainly keep it for backward compatibility and debugging.
"""
# Modify this in the future if we want to add more keys, for example, in image to video.
keys = ["vae_latent", "text_embedding"]
keys = [
"vae_latent", "text_embedding", "clip_feature", "first_frame_latent",
"pil_image"
]
def __init__(
self,
@@ -154,6 +161,7 @@ class LatentsParquetMapStyleDataset(Dataset):
cfg_rate: float = 0.0,
seed: int = 42,
drop_last: bool = True,
drop_first_row: bool = False,
text_padding_length: int = 512,
):
super().__init__()
@@ -168,15 +176,6 @@ class LatentsParquetMapStyleDataset(Dataset):
self.parquet_files, self.lengths = get_parquet_files_and_length(path)
self.batch = batch_size
self.text_padding_length = text_padding_length
self._cols = [
"vae_latent_bytes",
"vae_latent_shape",
"text_embedding_bytes",
"text_embedding_shape",
"text_embedding_dtype",
"height",
"width",
]
self.sampler = DP_SP_BatchSampler(
batch_size=batch_size,
dataset_size=sum(self.lengths),
@@ -184,6 +183,7 @@ class LatentsParquetMapStyleDataset(Dataset):
sp_world_size=get_sp_world_size(),
global_rank=get_world_rank(),
drop_last=drop_last,
drop_first_row=drop_first_row,
seed=seed,
)
logger.info("Dataset initialized with %d parquet files and %d rows",
@@ -196,15 +196,23 @@ class LatentsParquetMapStyleDataset(Dataset):
"""
return_dict = {}
for key in self.keys:
if f"{key}_shape" not in row_dict:
logger.warning_once(
f"Expected columns not found in row_dict for {key}.")
continue
shape = row_dict[f"{key}_shape"]
bytes = row_dict[f"{key}_bytes"]
# TODO (peiyuan): read precision
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
data = torch.from_numpy(data)
if len(bytes) > 0:
data = np.frombuffer(bytes,
dtype=np.float32).reshape(shape).copy()
data = torch.from_numpy(data)
else:
data = torch.zeros(shape, dtype=torch.float32)
return_dict[key] = data
return return_dict
def get_validation_negative_prompt(self) -> tuple[Any, Any, Any, Any]:
def get_validation_negative_prompt(self) -> tuple[Any, Any, Any]:
"""
Get the negative prompt for validation.
This method ensures the negative prompt is loaded and cached properly.
@@ -229,7 +237,10 @@ class LatentsParquetMapStyleDataset(Dataset):
padded_emb = padded_emb
mask = mask
return None, padded_emb, mask, None
# Get the negative prompt
negative_prompt = row_dict["caption"]
return padded_emb, mask, negative_prompt
def _pad(self, t: torch.Tensor, padding_length: int) -> torch.Tensor:
"""
@@ -263,25 +274,58 @@ class LatentsParquetMapStyleDataset(Dataset):
all_latents = []
all_embs = []
all_masks = []
all_clip_features = []
all_first_frame_latents = []
all_pil_images = []
all_infos = []
# Process each row individually
for i, row in enumerate(rows):
info_keys = [
"caption", "file_name", "media_type", "width", "height",
"num_frames", "duration_sec", "fps"
]
info = {}
for key in info_keys:
if key in row:
info[key] = row[key]
else:
info[key] = ""
info["prompt"] = info["caption"]
# Get tensors from row
data = self._get_torch_tensors_from_row_dict(row)
latents, emb = data["vae_latent"], data["text_embedding"]
padded_emb, mask = self._pad(emb, self.text_padding_length)
# Get extra latents
clip_features, first_frame_latents = data["clip_feature"], data[
"first_frame_latent"]
pil_image = data.get("pil_image", None)
# Store in batch tensors
all_latents.append(latents)
all_embs.append(padded_emb)
all_masks.append(mask)
all_clip_features.append(clip_features)
all_first_frame_latents.append(first_frame_latents)
all_pil_images.append(pil_image)
all_infos.append(info)
# Pin memory for faster transfer to GPU
all_latents = torch.stack(all_latents)
all_embs = torch.stack(all_embs)
all_masks = torch.stack(all_masks)
all_clip_features = torch.stack(all_clip_features)
all_first_frame_latents = torch.stack(all_first_frame_latents)
all_extra_latents = {
"encoder_hidden_states_image": all_clip_features,
"image_latents": all_first_frame_latents,
"pil_image": all_pil_images
}
return all_latents, all_embs, all_masks, indices
return all_latents, all_embs, all_masks, indices, all_extra_latents, all_infos
def __len__(self):
return sum(self.lengths)
@@ -300,6 +344,7 @@ def build_parquet_map_style_dataloader(
num_data_workers,
cfg_rate=0.0,
drop_last=True,
drop_first_row=False,
text_padding_length=512,
seed=42) -> Tuple[LatentsParquetMapStyleDataset, StatefulDataLoader]:
dataset = LatentsParquetMapStyleDataset(
@@ -307,6 +352,7 @@ def build_parquet_map_style_dataloader(
batch_size,
cfg_rate=cfg_rate,
drop_last=drop_last,
drop_first_row=drop_first_row,
text_padding_length=text_padding_length,
seed=seed)
@@ -318,4 +364,4 @@ def build_parquet_map_style_dataloader(
pin_memory=True,
persistent_workers=num_data_workers > 0,
)
return dataset, loader
return dataset, loader
+2
View File
@@ -618,6 +618,8 @@ class WanTransformer3DModel(CachableDiT):
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image)
if encoder_hidden_states.dim() == 2:
encoder_hidden_states = encoder_hidden_states.unsqueeze(0)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
@@ -9,6 +9,7 @@ from typing import Any, Dict, List, Optional
import numpy as np
import torch
import PIL
from PIL import Image
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_i2v
@@ -17,6 +18,7 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.pipelines.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.v1.models.vision_utils import numpy_to_pt, pil_to_numpy, normalize
class PreprocessPipeline_I2V(BasePreprocessPipeline):
@@ -32,12 +34,18 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
def get_extra_features(self, valid_data: Dict[str, Any],
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
features = {}
"""Get CLIP features from the first frame of each video."""
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
batch_size, _, num_frames, height, width = valid_data["pixel_values"].shape
latent_height = height // self.get_module("vae").spatial_compression_ratio
latent_width = width // self.get_module("vae").spatial_compression_ratio
processed_images = []
# Frame has values between -1 and 1
for frame in first_frame:
frame = (frame + 1) * 127.5
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
processed_img = self.get_module("image_processor")(
images=frame_pil, return_tensors="pt")
@@ -52,8 +60,69 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
with set_forward_context(current_timestep=0, attn_metadata=None):
clip_features = self.get_module("image_encoder")(**image_inputs)
clip_features = clip_features.last_hidden_state
features["clip_feature"] = clip_features
return {"clip_feature": clip_features}
"""Get VAE features from the first frame of each video"""
video_conditions = []
for frame in first_frame:
processed_img = frame.to(device="cpu", dtype=torch.float32)
processed_img = processed_img.unsqueeze(0).permute(0, 3, 1, 2).unsqueeze(2)
# (B, H, W, C) -> (B, C, 1, H, W)
video_condition = torch.cat([
processed_img,
processed_img.new_zeros(processed_img.shape[0], processed_img.shape[1],
num_frames - 1, height, width)
],
dim=2)
video_condition = video_condition.to(device=get_torch_device(),
dtype=torch.float32)
video_conditions.append(video_condition)
video_conditions = torch.cat(video_conditions, dim=0)
with torch.autocast(device_type="cuda",
dtype=torch.float32,
enabled=True):
encoder_outputs = self.get_module("vae").encode(video_conditions)
latent_condition = encoder_outputs.mean
if (hasattr(self.get_module("vae"), "shift_factor")
and self.get_module("vae").shift_factor is not None):
if isinstance(self.get_module("vae").shift_factor, torch.Tensor):
latent_condition -= self.get_module("vae").shift_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition -= self.get_module("vae").shift_factor
if isinstance(self.get_module("vae").scaling_factor, torch.Tensor):
latent_condition = latent_condition * self.get_module("vae").scaling_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition = latent_condition * self.get_module("vae").scaling_factor
mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
latent_width)
mask_lat_size[:, :, list(range(1, num_frames))] = 0
first_frame_mask = mask_lat_size[:, :, 0:1]
first_frame_mask = torch.repeat_interleave(
first_frame_mask,
dim=2,
repeats=self.get_module("vae").temporal_compression_ratio)
mask_lat_size = torch.concat(
[first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
mask_lat_size = mask_lat_size.view(batch_size, -1,
self.get_module("vae").temporal_compression_ratio,
latent_height, latent_width)
mask_lat_size = mask_lat_size.transpose(1, 2)
mask_lat_size = mask_lat_size.to(latent_condition.device)
image_latent = torch.concat([mask_lat_size, latent_condition],
dim=1)
features["first_frame_latent"] = image_latent
return features
def create_record(
self,
@@ -87,7 +156,37 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
"clip_feature_dtype": "",
})
if extra_features and "first_frame_latent" in extra_features:
first_frame_latent = extra_features["first_frame_latent"]
record.update({
"first_frame_latent_bytes": first_frame_latent.tobytes(),
"first_frame_latent_shape": list(first_frame_latent.shape),
"first_frame_latent_dtype": str(first_frame_latent.dtype),
})
else:
record.update({
"first_frame_latent_bytes": b"",
"first_frame_latent_shape": [],
"first_frame_latent_dtype": "",
})
return record
def preprocess(
self,
image: PIL.Image.Image
) -> torch.Tensor:
image = [image]
image = pil_to_numpy(image) # to np
image = numpy_to_pt(image) # to pt
do_normalize = True
if image.min() < 0:
do_normalize = False
if do_normalize:
image = normalize(image)
return image
EntryClass = PreprocessPipeline_I2V
+1 -1
View File
@@ -101,7 +101,7 @@ class DenoisingStage(PipelineStage):
n=sp_world_size).contiguous()
latents = latents[:, :, rank_in_sp_group, :, :, :]
batch.latents = latents
if batch.image_latent is not None:
if batch.image_latent is not None and latents.shape[2] != batch.image_latent.shape[2]:
image_latent = rearrange(batch.image_latent,
"b t (n s) h w -> b t n s h w",
n=sp_world_size).contiguous()
@@ -0,0 +1,558 @@
import os
import random
import sys
import time
from collections import deque
from copy import deepcopy
import gc
import torchvision
import imageio
from einops import rearrange
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from tqdm.auto import tqdm
from fastvideo import SamplingParam
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory, get_sp_group, get_torch_device,
get_world_group)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.wan.wan_pipeline import WanValidationPipeline
from fastvideo.v1.training.training_pipeline import TrainingPipeline
from fastvideo.v1.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, get_sigmas, load_checkpoint,
normalize_dit_input, save_checkpoint, shard_latents_across_sp)
from fastvideo.v1.dataset import build_parquet_map_style_dataloader
import wandb # isort: skip
logger = init_logger(__name__)
# Manual gradient checking flag - set to True to enable gradient verification
ENABLE_GRADIENT_CHECK = False
class WanI2VTrainingPipeline(TrainingPipeline):
"""
A training pipeline for Wan.
"""
_required_config_modules = ["scheduler", "transformer"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
args_copy.vae_config.load_encoder = False
validation_pipeline = WanValidationPipeline.from_pretrained(
training_args.model_path,
args=None,
inference_mode=True,
loaded_modules={"transformer": self.get_module("transformer")},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus
)
self.validation_pipeline = validation_pipeline
self.latents = None
def train_one_step(
self,
transformer,
model_type,
optimizer,
lr_scheduler,
loader_iter,
noise_scheduler,
noise_random_generator,
gradient_accumulation_steps,
sp_size,
precondition_outputs,
max_grad_norm,
weighting_scheme,
logit_mean,
logit_std,
mode_scale,
) -> tuple[float, float]:
assert self.training_args is not None
self.modules["transformer"].requires_grad_(True)
self.modules["transformer"].train()
total_loss = 0.0
optimizer.zero_grad()
for _ in range(gradient_accumulation_steps):
latents = self.latents
encoder_hidden_states = self.encoder_hidden_states
encoder_attention_mask = self.encoder_attention_mask
infos = self.infos
extra_latents = self.extra_latents
# logger.info("rank: %s, caption: %s",
# self.rank,
# infos['caption'],
# local_main_process_only=False)
# TODO(will): don't hardcode bfloat16
latents = latents.to(get_torch_device(),
dtype=torch.bfloat16)
latents = shard_latents_across_sp(latents, self.training_args.num_latent_t)
encoder_hidden_states = encoder_hidden_states.to(
get_torch_device(), dtype=torch.bfloat16)
latents = normalize_dit_input(model_type, latents)
batch_size = latents.shape[0]
noise = torch.randn_like(latents)
u = compute_density_for_timestep_sampling(
weighting_scheme=weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
logit_mean=logit_mean,
logit_std=logit_std,
mode_scale=mode_scale,
)
indices = (u * noise_scheduler.config.num_train_timesteps).long()
timesteps = noise_scheduler.timesteps[indices].to(
device=latents.device)
if sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
sp_group = get_sp_group()
sp_group.broadcast(timesteps, src=0)
sigmas = get_sigmas(
noise_scheduler,
latents.device,
timesteps,
n_dim=latents.ndim,
dtype=latents.dtype,
)
noisy_model_input = (1.0 - sigmas) * latents + sigmas * noise
input_kwargs = {}
# I2V
if extra_latents:
image_embeds, image_latents = extra_latents["encoder_hidden_states_image"], extra_latents["image_latents"]
# Image Embeds
assert torch.isnan(image_embeds).sum() == 0
image_embeds = image_embeds.to(get_torch_device(),
dtype=torch.bfloat16)
input_kwargs["encoder_hidden_states_image"] = image_embeds
# Image Latents
assert torch.isnan(image_latents).sum() == 0
image_latents = image_latents.to(get_torch_device(),
dtype=torch.bfloat16)
image_latents = shard_latents_across_sp(image_latents, self.training_args.num_latent_t)
noisy_model_input = torch.cat(
[noisy_model_input, image_latents],
dim=1)
with torch.autocast("cuda", dtype=torch.bfloat16):
input_kwargs.update({
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"encoder_attention_mask": encoder_attention_mask, # B, L
"return_dict": False,
})
if 'hunyuan' in model_type:
input_kwargs["guidance"] = torch.tensor(
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
with set_forward_context(current_timestep=timesteps,
attn_metadata=None):
model_pred = transformer(**input_kwargs)
if precondition_outputs:
model_pred = noisy_model_input - model_pred * sigmas
target = latents if precondition_outputs else noise - latents
loss = (torch.mean((model_pred.float() - target.float())**2) /
gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
# logger.info(f"rank: {self.rank}, avg_loss: {avg_loss.item()}",
# local_main_process_only=False)
world_group = get_world_group()
world_group.all_reduce(avg_loss, op=torch.distributed.ReduceOp.AVG)
total_loss += avg_loss.item()
# TODO(will): perhaps move this into transformer api so that we can do
# the following:
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
if max_grad_norm is not None:
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,
foreach=None,
)
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
else:
grad_norm = 0.0
optimizer.step()
lr_scheduler.step()
return total_loss, grad_norm
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
):
assert self.training_args is not None
# Set random seeds for deterministic training
seed = self.training_args.seed if self.training_args.seed is not None else 42
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
noise_random_generator = torch.Generator(device="cpu").manual_seed(seed)
logger.info("Initialized random seeds with seed: %s", seed)
noise_scheduler = FlowMatchEulerDiscreteScheduler()
# Train!
assert self.training_args.sp_size is not None
assert self.training_args.gradient_accumulation_steps is not None
total_batch_size = (self.world_size *
self.training_args.gradient_accumulation_steps /
self.training_args.sp_size *
self.training_args.train_sp_batch_size)
logger.info("***** Running training *****")
logger.info(" Num examples = %s", len(self.train_dataset))
logger.info(" Dataloader size = %s", len(self.train_dataloader))
logger.info(" Num Epochs = %s", self.num_train_epochs)
logger.info(" Resume training from step %s",
self.init_steps) # type: ignore
logger.info(" Instantaneous batch size per device = %s",
self.training_args.train_batch_size)
logger.info(
" Total train batch size (w. data & sequence parallel, accumulation) = %s",
total_batch_size)
logger.info(" Gradient Accumulation steps = %s",
self.training_args.gradient_accumulation_steps)
logger.info(" Total optimization steps = %s",
self.training_args.max_train_steps)
logger.info(
" Total training parameters per FSDP shard = %s B",
sum(p.numel()
for p in self.transformer.parameters() if p.requires_grad) /
1e9)
# print dtype
logger.info(" Master weight dtype: %s",
self.transformer.parameters().__next__().dtype)
if self.training_args.resume_from_checkpoint:
logger.info("Loading checkpoint from %s",
self.training_args.resume_from_checkpoint)
resumed_step = load_checkpoint(
self.transformer, self.global_rank,
self.training_args.resume_from_checkpoint, self.optimizer,
self.train_dataloader, self.lr_scheduler,
noise_random_generator)
if resumed_step > 0:
self.init_steps = resumed_step
logger.info("Successfully resumed from step %s", resumed_step)
else:
logger.warning(
"Failed to load checkpoint, starting from step 0")
self.init_steps = 0
progress_bar = tqdm(
range(0, self.training_args.max_train_steps),
initial=self.init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=self.local_rank > 0,
)
self.train_loader_iter = iter(self.train_dataloader)
(
self.latents,
self.encoder_hidden_states,
self.encoder_attention_mask,
self.indices,
self.extra_latents,
self.infos
) = next(self.train_loader_iter, None)
self.infos = self.infos[0]
step_times: deque[float] = deque(maxlen=100)
# TODO(will): fix this
# for i in range(self.init_steps):
# next(loader_iter)
# get gpu memory usage
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)
for step in range(self.init_steps + 1,
self.training_args.max_train_steps + 1):
start_time = time.perf_counter()
loss, grad_norm = self.train_one_step(
self.transformer,
# args.model_type,
"wan",
self.optimizer,
self.lr_scheduler,
self.train_loader_iter,
noise_scheduler,
noise_random_generator,
self.training_args.gradient_accumulation_steps,
self.training_args.sp_size,
self.training_args.precondition_outputs,
self.training_args.max_grad_norm,
self.training_args.weighting_scheme,
self.training_args.logit_mean,
self.training_args.logit_std,
self.training_args.mode_scale,
)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info("GPU memory usage after train_one_step: %s MB",
gpu_memory_usage)
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
# Manual gradient checking - only at first step
if step == 1 and ENABLE_GRADIENT_CHECK:
logger.info("Performing gradient check at step %s", step)
self.setup_gradient_check(args, self.train_loader_iter,
noise_scheduler,
noise_random_generator)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
})
progress_bar.update(1)
if self.global_rank == 0:
wandb.log(
{
"train_loss": loss,
"learning_rate": self.lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
},
step=step,
)
if step % self.training_args.checkpointing_steps == 0:
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir, step,
self.optimizer, self.train_dataloader,
self.lr_scheduler, noise_random_generator)
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)
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir,
self.training_args.max_train_steps, self.optimizer,
self.train_dataloader, self.lr_scheduler,
noise_random_generator)
if get_sp_group():
cleanup_dist_env_and_memory()
@torch.no_grad()
def _log_validation(self, transformer, training_args, global_step) -> None:
assert training_args is not None
training_args.inference_mode = True
training_args.use_cpu_offload = False
if not training_args.log_validation:
return
if self.validation_pipeline is None:
raise ValueError("Validation pipeline is not set")
logger.info("Starting validation")
# Create sampling parameters if not provided
sampling_param = SamplingParam.from_pretrained(training_args.model_path)
# Set deterministic seed for validation
validation_seed = training_args.seed if training_args.seed is not None else 42
torch.manual_seed(validation_seed)
torch.cuda.manual_seed_all(validation_seed)
logger.info("Using validation seed: %s", validation_seed)
# Prepare validation prompts
logger.info('fastvideo_args.validation_prompt_dir: %s',
training_args.validation_prompt_dir)
validation_dataset, validation_dataloader = build_parquet_map_style_dataloader(
training_args.validation_prompt_dir,
batch_size=1,
num_data_workers=0,
drop_last=False,
drop_first_row=sampling_param.negative_prompt is not None,
cfg_rate=training_args.cfg)
if sampling_param.negative_prompt:
negative_prompt_embeds, negative_prompt_attention_mask, _ = validation_dataset.get_validation_negative_prompt(
)
transformer.eval()
# Process each validation prompt
videos = []
captions = []
caption = self.infos['caption']
captions.extend(caption)
prompt_embeds = self.encoder_hidden_states.to(get_torch_device())
prompt_attention_mask = self.encoder_attention_mask.to(get_torch_device())
image_embeds = self.extra_latents["encoder_hidden_states_image"].to(get_torch_device())
image_latent = self.extra_latents["image_latents"].to(get_torch_device())
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8,
sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
# 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],
negative_attention_mask=[negative_prompt_attention_mask],
image_embeds=[image_embeds],
image_latent=shard_latents_across_sp(image_latent, self.training_args.num_latent_t),
# 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=sampling_param.guidance_scale,
n_tokens=n_tokens,
eta=0.0,
)
# Run validation inference
with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
# Re-enable gradients for training
transformer.requires_grad_(True)
transformer.train()
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
videos.append(frames)
# Log validation results
world_group = get_world_group()
num_sp_groups = world_group.world_size // self.sp_group.world_size
# Only sp_group leaders (rank_in_sp_group == 0) need to send their
# results to global rank 0
if self.rank_in_sp_group == 0:
if self.global_rank == 0:
# Global rank 0 collects results from all sp_group leaders
all_videos = videos # Start with own results
all_captions = captions
# Receive from other sp_group leaders
for sp_group_idx in range(1, num_sp_groups):
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
recv_videos = world_group.recv_object(src=src_rank)
recv_captions = world_group.recv_object(src=src_rank)
all_videos.extend(recv_videos)
all_captions.extend(recv_captions)
video_filenames = []
for i, (video,
caption) in enumerate(zip(all_videos, all_captions)):
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
f"validation_step_{global_step}_video_{i}.mp4")
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
logs = {
"validation_videos": [
wandb.Video(filename, caption=caption) for filename,
caption in zip(video_filenames, all_captions)
]
}
wandb.log(logs, step=global_step)
else:
# Other sp_group leaders send their results to global rank 0
world_group.send_object(videos, dst=0)
world_group.send_object(captions, dst=0)
gc.collect()
torch.cuda.empty_cache()
def main(args) -> None:
logger.info("Starting training pipeline...")
pipeline = WanI2VTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.forward(None, args)
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.v1.fastvideo_args import TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.use_cpu_offload = False
main(args)
+54
View File
@@ -0,0 +1,54 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export HOME="/mnt/weka/home/hao.zhang/wei"
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
DATA_DIR=$HOME/FastVideo/data/wei-i2v-dataset/crush-smol_preprocessed/combined_parquet_dataset
VALIDATION_DIR=$HOME/FastVideo/data/wei-i2v-dataset/crush-smol_preprocessed/validation_parquet_dataset
NUM_GPUS=4
# IP=[MASTER NODE IP]
CHECKPOINT_PATH="$DATA_DIR/outputs/wan_i2v_finetune/checkpoint-50"
# 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_i2v_training_pipeline.py\
--model_path Wan-AI/Wan2.1-I2V-14B-480P-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-I2V-14B-480P-Diffusers \
--cache_dir "$HOME/ray/.cache"\
--data_path "$DATA_DIR"\
--validation_prompt_dir "$VALIDATION_DIR"\
--train_batch_size=1\
--num_latent_t 8 \
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--dp_shards $NUM_GPUS \
--train_sp_batch_size 1\
--dataloader_num_workers 1\
--gradient_accumulation_steps=8 \
--max_train_steps=5000 \
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=500000 \
--validation_steps 20\
--validation_sampling_steps "2,4,8" \
--log_validation \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune"\
--tracker_project_name wan_i2v_finetune \
--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 \
# --resume_from_checkpoint "$CHECKPOINT_PATH"
@@ -1,9 +1,10 @@
# export WANDB_MODE="offline"
export WANDB_MODE="offline"
export HOME="/mnt/weka/home/hao.zhang/wei"
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="finetrainers/crush-smol/merge.txt"
OUTPUT_DIR="crush-smol_preprocess"
DATA_MERGE_PATH="$HOME/FastVideo/data/wei-i2v-dataset/crush-smol_raw/merge.txt"
OUTPUT_DIR="$HOME/FastVideo/data/wei-i2v-dataset/crush-smol_preprocessed"
VALIDATION_PATH="assets/prompt.txt"
torchrun --nproc_per_node=$GPU_NUM \