Compare commits

...
Author SHA1 Message Date
SolitaryThinker e26b389f37 i2v validation 2025-06-11 21:35:10 -07:00
JerryZhou54 86dc4c4bb3 Small change 2025-06-11 13:29:11 -07:00
JerryZhou54 41d0400832 Fix preprocess 2025-06-11 13:29:11 -07:00
Wei Zhou e97aa17d33 Update preprocess_pipeline_i2v.py 2025-06-11 13:29:11 -07:00
JerryZhou54 0a40b79b36 I2V Runnable 2025-06-11 13:29:11 -07:00
“BrianChen1129” f8c69045d6 Add Encoded first frame to processed dataset 2025-06-11 13:29:10 -07:00
20 changed files with 1111 additions and 28 deletions
@@ -0,0 +1,31 @@
{
"data": [
{
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
"image_path": null,
"video_path": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-034.mp4",
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-027.mp4",
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-030.mp4",
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
+2 -2
View File
@@ -18,7 +18,7 @@ logger = init_logger(__name__)
def main(args):
args.model_path = maybe_download_model(args.model_path)
maybe_init_distributed_environment_and_model_parallel(args.tp_size, args.sp_size)
maybe_init_distributed_environment_and_model_parallel(1, 1)
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
@@ -43,7 +43,7 @@ if __name__ == "__main__":
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--model_type", type=str, default="mochi")
parser.add_argument("--data_merge_path", type=str, required=True)
parser.add_argument("--validation_prompt_txt", type=str)
parser.add_argument("--validation_dataset_file", type=str)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
@@ -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()),
@@ -268,6 +268,8 @@ class LatentsParquetMapStyleDataset(Dataset):
for i, row in enumerate(rows):
# Get tensors from row
data = self._get_torch_tensors_from_row_dict(row)
print(data)
import pdb; pdb.set_trace()
latents, emb = data["vae_latent"], data["text_embedding"]
padded_emb, mask = self._pad(emb, self.text_padding_length)
+7 -6
View File
@@ -523,7 +523,8 @@ class TrainingArgs(FastVideoArgs):
precondition_outputs: bool = False
# validation & logs
validation_prompt_dir: str = ""
validation_dataset_file: str = ""
validation_path: str = ""
validation_sampling_steps: str = ""
validation_guidance_scale: str = ""
validation_steps: float = 0.0
@@ -576,9 +577,6 @@ class TrainingArgs(FastVideoArgs):
# master_weight_type
master_weight_type: str = ""
# For fast checking in LoRA pipeline
training_mode: bool = True
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
# Get all fields from the dataclass
@@ -671,9 +669,12 @@ class TrainingArgs(FastVideoArgs):
help="Whether to precondition the outputs of the model")
# Validation and logging
parser.add_argument("--validation-prompt-dir",
parser.add_argument("--validation-dataset-file",
type=str,
help="Directory containing validation prompts")
help="File containing validation dataset")
parser.add_argument("--validation-path",
type=str,
help="Path to validation dataset")
parser.add_argument("--validation-sampling-steps",
type=str,
help="Validation sampling steps")
+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:
@@ -12,6 +12,7 @@ from dataclasses import asdict, dataclass, field
from typing import Any, Dict, List, Optional, Union
import torch
import PIL.Image
from fastvideo.v1.configs.sample.teacache import (TeaCacheParams,
WanTeaCacheParams)
@@ -36,6 +37,8 @@ class ForwardBatch:
# Image inputs
image_path: Optional[str] = None
image_embeds: List[torch.Tensor] = field(default_factory=list)
preprocessed_image: Optional[torch.Tensor] = None
pil_image: Optional[PIL.Image.Image] = None
# Text inputs
prompt: Optional[Union[str, List[str]]] = None
@@ -9,7 +9,15 @@ from typing import Any, Dict, List, Optional
import numpy as np
import torch
import PIL
from PIL import Image
import os
from tqdm import tqdm
from concurrent.futures import ProcessPoolExecutor
import gc
import pyarrow as pa
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_i2v
from fastvideo.v1.distributed import get_torch_device
@@ -17,6 +25,13 @@ 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
from fastvideo.v1.dataset.validation_dataset import ValidationDataset
from fastvideo.v1.pipelines.stages import TextEncodingStage, ImageEncodingStage
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class PreprocessPipeline_I2V(BasePreprocessPipeline):
@@ -26,18 +41,235 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
"text_encoder", "tokenizer", "vae", "image_encoder", "image_processor"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
def preprocess_validation_text(self, fastvideo_args: FastVideoArgs, args):
"""Process validation text prompts and save them to parquet files.
This base implementation handles the common validation text processing logic.
Subclasses can override this method to add pipeline-specific features.
"""
# Create Parquet dataset directory for validation
validation_parquet_dir = os.path.join(args.output_dir,
"validation_parquet_dataset")
os.makedirs(validation_parquet_dir, exist_ok=True)
validation_dataset = ValidationDataset(args.validation_dataset_file)
from itertools import chain
# for data in validation_dataset:
# print(data)
# assert False
# with open(args.validation_prompt_txt, encoding="utf-8") as file:
# lines = file.readlines()
# prompts = [line.strip() for line in lines]
# Prepare batch data for Parquet dataset
batch_data = []
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
if sampling_param.negative_prompt:
negative_prompt = {
'caption': sampling_param.negative_prompt,
'image_path': None,
'video_path': None,
}
else:
negative_prompt = None
validation_dataset = chain([negative_prompt], validation_dataset)
# Add progress bar for validation text preprocessing
pbar = tqdm(enumerate(validation_dataset),
desc="Processing validation dataset",
unit="sample")
for idx, row in pbar:
print(idx)
# print(type(row))
# print(row)
prompt = row['caption']
print(row)
print(prompt)
with torch.inference_mode():
# Text Encoder
batch = ForwardBatch(
data_type="video",
prompt=prompt,
prompt_embeds=[],
prompt_attention_mask=[],
)
assert hasattr(self, "prompt_encoding_stage")
result_batch = self.prompt_encoding_stage(batch, fastvideo_args)
if idx != 0:
image = None
if 'image' in row:
image = row['image']
else:
assert 'video' in row
image = row['video'][0]
assert image is not None
result_batch.pil_image = image
# assert image is not None
assert hasattr(self, "image_encoding_stage")
result_batch = self.image_encoding_stage(result_batch, fastvideo_args)
prompt_embeds = result_batch.prompt_embeds[0]
prompt_attention_mask = result_batch.prompt_attention_mask[0]
file_name = prompt.split(".")[0]
# Get the sequence length from attention mask (number of 1s)
seq_len = prompt_attention_mask.sum().item()
text_embedding = prompt_embeds[0, :seq_len].cpu().numpy()
text_attention_mask = prompt_attention_mask[
0, :seq_len].cpu().numpy().astype(np.uint8)
if idx != 0:
image_embeds = result_batch.image_embeds[0]
image_embeds = image_embeds.cpu().numpy()
else:
image_embeds = None
# Log the shapes after removing padding
logger.info(
"Shape after removing padding - Embeddings: %s, Mask: %s",
text_embedding.shape, text_attention_mask.shape)
# image embedding
if idx != 0:
clip_feature = {
'clip_feature': image_embeds,
}
else:
clip_feature = None
# Create record for Parquet dataset
record = self.create_record(video_name=file_name,
vae_latent=np.array([],
dtype=np.float32),
text_embedding=text_embedding,
text_attention_mask=text_attention_mask,
valid_data=None,
idx=0,
extra_features=clip_feature)
batch_data.append(record)
logger.info("Saved validation sample: %s", file_name)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = []
for field in self.get_schema_fields():
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
elif field in ['width', 'height', 'num_frames']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.int32()))
elif field in ['duration_sec', 'fps']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.float32()))
else:
arrays.append(
pa.array([record[field] for record in batch_data]))
table = pa.Table.from_arrays(arrays, names=self.get_schema_fields())
write_pbar.update(1)
write_pbar.close()
logger.info("Total validation samples: %s", len(table))
work_range = (0, 1, table, 0, validation_parquet_dir, len(table))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=1) as executor:
futures = {
executor.submit(self.process_chunk_range, work_range):
work_range
}
for future in tqdm(futures, desc="Processing chunks"):
try:
total_written += future.result()
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error("Failed to process range %s-%s: %s",
work_range[0], work_range[1], str(e))
if failed_ranges:
logger.warning("Retrying %s failed ranges sequentially",
len(failed_ranges))
for work_range in failed_ranges:
try:
total_written += self.process_chunk_range(work_range)
except Exception as e:
logger.error(
"Failed to process range %s-%s after retry: %s",
work_range[0], work_range[1], str(e))
logger.info("Total validation samples written: %s", total_written)
# Clear memory
del table
gc.collect() # Force garbage collection
def get_schema_fields(self) -> List[str]:
"""Get the schema fields for I2V pipeline."""
return [f.name for f in pyarrow_schema_i2v]
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 +284,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["encoded_first_frame"] = image_latent
return features
def create_record(
self,
@@ -87,6 +380,20 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
"clip_feature_dtype": "",
})
if extra_features and "encoded_first_frame" in extra_features:
encoded_first_frame = extra_features["encoded_first_frame"]
record.update({
"encoded_first_frame_bytes": encoded_first_frame.tobytes(),
"encoded_first_frame_shape": list(encoded_first_frame.shape),
"encoded_first_frame_dtype": str(encoded_first_frame.dtype),
})
else:
record.update({
"encoded_first_frame_bytes": b"",
"encoded_first_frame_shape": [],
"encoded_first_frame_dtype": "",
})
return record
@@ -419,7 +419,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
# Clear memory
del table
gc.collect() # Force garbage collection
def _flush_tables(self, num_processed_samples: int, args,
combined_parquet_dir: str):
"""Flush collected tables to disk."""
+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()
+15 -7
View File
@@ -49,22 +49,30 @@ class EncodingStage(PipelineStage):
"""
self.vae = self.vae.to(get_torch_device())
image_path = batch.image_path
assert batch.pil_image is not None, "Image must be provided"
image = batch.pil_image
# image_path = batch.image_path
# TODO(will): remove this once we add input/output validation for stages
if image_path is None:
raise ValueError("Image Path must be provided")
# if image_path is None:
# raise ValueError("Image Path must be provided")
assert batch.height is not None
assert batch.width is not None
latent_height = batch.height // self.vae.spatial_compression_ratio
latent_width = batch.width // self.vae.spatial_compression_ratio
image = load_image(image_path)
image = self.preprocess(
# image = load_image(image_path)
image = self.preprocess_image(
image,
vae_scale_factor=self.vae.spatial_compression_ratio,
height=batch.height,
width=batch.width).to(get_torch_device(), dtype=torch.float32)
image = image.unsqueeze(2)
# if batch.preprocessed_image is not None:
# image = batch.preprocessed_image
video_condition = torch.cat([
image,
image.new_zeros(image.shape[0], image.shape[1],
@@ -148,7 +156,7 @@ class EncodingStage(PipelineStage):
raise AttributeError(
"Could not access latents of provided encoder_output")
def preprocess(
def preprocess_image(
self,
image: PIL.Image.Image,
vae_scale_factor: int,
@@ -172,4 +180,4 @@ class EncodingStage(PipelineStage):
if do_normalize:
image = normalize(image)
return image
return image
@@ -56,7 +56,9 @@ class ImageEncodingStage(PipelineStage):
if fastvideo_args.use_cpu_offload:
self.image_encoder = self.image_encoder.to(get_torch_device())
image = load_image(batch.image_path)
# image = load_image(batch.image_path)
assert batch.pil_image is not None, "Image must be provided"
image = batch.pil_image
image_inputs = self.image_processor(
images=image, return_tensors="pt").to(get_torch_device())
@@ -3,12 +3,21 @@
Input validation stage for diffusion pipelines.
"""
import torch
from typing import Optional
import torch
from fastvideo.v1.models.vision_utils import (load_image, pil_to_numpy,
numpy_to_pt, normalize,
get_default_height_width,
resize)
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
import PIL.Image
logger = init_logger(__name__)
@@ -85,5 +94,44 @@ class InputValidationStage(PipelineStage):
raise ValueError(
f"Guidance scale must be positive, but got {batch.guidance_scale}"
)
# for i2v, get image from image_path
if batch.image_path is not None:
image = load_image(batch.image_path)
batch.pil_image = image
# image = self.preprocess_image(
# image,
# vae_scale_factor=self.vae.spatial_compression_ratio,
# height=batch.height,
# width=batch.width).to(get_torch_device(), dtype=torch.float32)
# image = image.unsqueeze(2)
# batch.preprocessed_image = image
return batch
def preprocess_image(
self,
image: PIL.Image.Image,
vae_scale_factor: int,
height: Optional[int] = None,
width: Optional[int] = None,
resize_mode: str = "default", # "default", "fill", "crop"
) -> torch.Tensor:
image = [image]
height, width = get_default_height_width(image[0], vae_scale_factor,
height, width)
image = [
resize(i, height, width, resize_mode=resize_mode) for i in 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
@@ -76,4 +76,36 @@ class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
stage=DecodingStage(vae=self.get_module("vae")))
class WanImageToVideoValidationPipeline(ComposedPipelineBase):
"""
I2V Validation pipeline for Wan2.1, assumes that the input are preprocess latents.
"""
_required_config_modules = ["vae", "scheduler", "transformer"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=EncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = WanImageToVideoPipeline
@@ -0,0 +1,588 @@
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_world_group, get_torch_device)
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_i2v_pipeline import WanImageToVideoValidationPipeline
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)
# from fastvideo.v1.dataset.parquet_datasets import ParquetVideoTextDataset
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
args_copy.log_validation = False
validation_pipeline = WanImageToVideoValidationPipeline(
training_args.model_path,
fastvideo_args=args_copy,
# inference_mode=True,
loaded_modules={"transformer": self.get_module("transformer")},
)
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):
(
self.latents,
self.encoder_hidden_states,
self.encoder_attention_mask,
self.infos,
self.extra_latents
) = next(self.train_loader_iter, None)
if self.latents is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
(
self.latents,
self.encoder_hidden_states,
self.encoder_attention_mask,
self.infos,
self.extra_latents
) = next(self.train_loader_iter)
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)
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["img_embed"], extra_latents["img_lat"]
# 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)
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)
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_path: %s',
training_args.validation_path)
# validation_dataset, validation_dataloader = build_parquet_map_style_dataloader(
# training_args.validation_path,
# batch_size=1,
# num_data_workers=0,
# drop_last=False,
# cfg_rate=training_args.cfg)
if sampling_param.negative_prompt:
_, negative_prompt_embeds, negative_prompt_attention_mask, _ = validation_dataset.get_validation_negative_prompt(
)
# validation_dataset = ParquetVideoTextDataset(
# training_args.validation_prompt_dir,
# batch_size=1,
# cfg_rate=training_args.cfg,
# 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_loader_iter = iter(validation_dataloader)
transformer.eval()
# Process each validation prompt
videos = []
captions = []
for batch in validation_loader_iter:
logger.info(batch)
(
latents,
encoder_hidden_states,
encoder_attention_mask,
infos,
extra_latents
) = batch
# caption = infos['caption']
# captions.extend(caption)
prompt_embeds = encoder_hidden_states.to(get_torch_device())
prompt_attention_mask = encoder_attention_mask.to(get_torch_device())
image_embeds = extra_latents["img_embed"].to(get_torch_device())
image_latent = extra_latents["img_lat"].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=image_latent,
# 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,16 +1,18 @@
# export WANDB_MODE="offline"
export WANDB_MODE="offline"
export HOME="/mnt/user_storage/src/"
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"
VALIDATION_PATH="assets/prompt.txt"
DATA_MERGE_PATH="$HOME/FastVideo/data/crush-smol/merge.txt"
OUTPUT_DIR="$HOME/FastVideo/data/crush-smol_parq_i2v"
# VALIDATION_PATH="assets/prompt.txt"
VALIDATION_PATH="examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation.json"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/data_preprocess/preprocess.py \
fastvideo/data_preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
--preprocess_video_batch_size 4 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
@@ -18,7 +20,7 @@ torchrun --nproc_per_node=$GPU_NUM \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_prompt_txt $VALIDATION_PATH \
--samples_per_file 16 \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 32 \
--flush_frequency 32 \
--preprocess_task "i2v"