Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5ae1b1ca91 | ||
|
|
7e6e37b6b1 | ||
|
|
64c9c780df | ||
|
|
bd7b7078ba | ||
|
|
3cde7934ab | ||
|
|
c0147f611e | ||
|
|
f6314572e6 |
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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 \
|
||||
|
||||
Reference in New Issue
Block a user