Compare commits
6
Commits
reason
...
will/i2v_val
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e26b389f37 | ||
|
|
86dc4c4bb3 | ||
|
|
41d0400832 | ||
|
|
e97aa17d33 | ||
|
|
0a40b79b36 | ||
|
|
f8c69045d6 |
@@ -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
|
||||
}
|
||||
]
|
||||
}
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user