Compare commits
22
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
04adb96905 | ||
|
|
b4665e360f | ||
|
|
00d4065b1a | ||
|
|
338cd6d12b | ||
|
|
1b1b7e28a2 | ||
|
|
8e0417fa7c | ||
|
|
8039b29624 | ||
|
|
36f845b465 | ||
|
|
cfb2fe3bde | ||
|
|
8c88fa5d6a | ||
|
|
87e45872ff | ||
|
|
ede88863a6 | ||
|
|
d83db1ea4a | ||
|
|
1904947319 | ||
|
|
5f8fb3cf5f | ||
|
|
57aea577c2 | ||
|
|
329b93b0f6 | ||
|
|
3a127bf455 | ||
|
|
686b91b94e | ||
|
|
52c3b1d627 | ||
|
|
da66702631 | ||
|
|
513d6513aa |
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,25 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
|
||||
DATASET_PATH="data/crush-smol"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v_hunyuan15/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM --master_port=29513 \
|
||||
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
|
||||
--model_path $MODEL_PATH \
|
||||
--mode preprocess \
|
||||
--workload_type t2v \
|
||||
--preprocess.video_loader_type torchvision \
|
||||
--preprocess.dataset_type merged \
|
||||
--preprocess.dataset_path $DATASET_PATH \
|
||||
--preprocess.dataset_output_dir $OUTPUT_DIR \
|
||||
--preprocess.preprocess_video_batch_size 2 \
|
||||
--preprocess.dataloader_num_workers 0 \
|
||||
--preprocess.max_height 480 \
|
||||
--preprocess.max_width 832 \
|
||||
--preprocess.num_frames 61 \
|
||||
--preprocess.train_fps 15 \
|
||||
--preprocess.samples_per_file 8 \
|
||||
--preprocess.flush_frequency 8 \
|
||||
--preprocess.video_length_tolerance_range 5
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,25 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="hunyuanvideo-community/HunyuanVideo"
|
||||
DATASET_PATH="data/crush-smol"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v_hunyuan/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
|
||||
--model_path $MODEL_PATH \
|
||||
--mode preprocess \
|
||||
--workload_type t2v \
|
||||
--preprocess.video_loader_type torchvision \
|
||||
--preprocess.dataset_type merged \
|
||||
--preprocess.dataset_path $DATASET_PATH \
|
||||
--preprocess.dataset_output_dir $OUTPUT_DIR \
|
||||
--preprocess.preprocess_video_batch_size 2 \
|
||||
--preprocess.dataloader_num_workers 0 \
|
||||
--preprocess.max_height 480 \
|
||||
--preprocess.max_width 832 \
|
||||
--preprocess.num_frames 77 \
|
||||
--preprocess.train_fps 16 \
|
||||
--preprocess.samples_per_file 8 \
|
||||
--preprocess.flush_frequency 8 \
|
||||
--preprocess.video_length_tolerance_range 5
|
||||
@@ -44,6 +44,15 @@ class CLIPTextArchConfig(TextEncoderArchConfig):
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [_is_transformer_layer, _is_embeddings])
|
||||
|
||||
tokenizer_kwargs: dict = field(
|
||||
default_factory=lambda: {
|
||||
"padding": "max_length",
|
||||
"truncation": True,
|
||||
"max_length": 77,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class CLIPVisionArchConfig(ImageEncoderArchConfig):
|
||||
|
||||
@@ -56,6 +56,15 @@ class LlamaArchConfig(TextEncoderArchConfig):
|
||||
default_factory=lambda:
|
||||
[_is_transformer_layer, _is_embeddings, _is_final_norm])
|
||||
|
||||
tokenizer_kwargs: dict = field(
|
||||
default_factory=lambda: {
|
||||
"padding": "max_length",
|
||||
"truncation": True,
|
||||
"max_length": 256,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LlamaConfig(TextEncoderConfig):
|
||||
|
||||
@@ -72,9 +72,16 @@ def i2v_record_creator(batch: PreprocessBatch) -> list[dict[str, Any]]:
|
||||
"""Create a record for the Parquet dataset with CLIP features."""
|
||||
records = basic_t2v_record_creator(batch)
|
||||
|
||||
assert len(
|
||||
batch.image_embeds) == 1, "image embedding should be a single tensor"
|
||||
image_embeds = batch.image_embeds[0]
|
||||
# Adapt for model doesn't have image encoder, e.g., Hunyuan
|
||||
if len(batch.image_embeds) == 0:
|
||||
image_embeds = None
|
||||
elif len(batch.image_embeds) == 1:
|
||||
image_embeds = batch.image_embeds[0]
|
||||
else:
|
||||
raise ValueError(
|
||||
"Unexpected number of image_embeds in batch: expected 0 or 1, got {}".format(
|
||||
len(batch.image_embeds)))
|
||||
|
||||
image_latent = batch.image_latent
|
||||
pil_image = batch.pil_image
|
||||
|
||||
|
||||
@@ -30,12 +30,13 @@ def get_torch_tensors_from_row_dict(row_dict,
|
||||
"""
|
||||
return_dict = {}
|
||||
for key in keys:
|
||||
shape, bytes = None, None
|
||||
shape, bytes, dtype_str = None, None, None
|
||||
if isinstance(key, tuple):
|
||||
for k in key:
|
||||
try:
|
||||
shape = row_dict[f"{k}_shape"]
|
||||
bytes = row_dict[f"{k}_bytes"]
|
||||
dtype_str = row_dict.get(f"{k}_dtype", 'float32')
|
||||
except KeyError:
|
||||
continue
|
||||
key = key[0]
|
||||
@@ -44,13 +45,25 @@ def get_torch_tensors_from_row_dict(row_dict,
|
||||
else:
|
||||
shape = row_dict[f"{key}_shape"]
|
||||
bytes = row_dict[f"{key}_bytes"]
|
||||
dtype_str = row_dict.get(f"{key}_dtype", 'float32')
|
||||
|
||||
# Convert dtype string to numpy dtype
|
||||
if 'float16' in dtype_str:
|
||||
np_dtype = np.float16
|
||||
elif 'float32' in dtype_str:
|
||||
np_dtype = np.float32
|
||||
elif 'float64' in dtype_str:
|
||||
np_dtype = np.float64
|
||||
else:
|
||||
# Default to float32 if dtype is unrecognized
|
||||
np_dtype = np.float32
|
||||
|
||||
# TODO (peiyuan): read precision
|
||||
if key == 'text_embedding' and (rng.random()
|
||||
if rng else random.random()) < cfg_rate:
|
||||
data = np.zeros((512, 4096), dtype=np.float32)
|
||||
else:
|
||||
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
|
||||
data = np.frombuffer(bytes, dtype=np_dtype).reshape(shape).copy()
|
||||
data = torch.from_numpy(data)
|
||||
if len(data.shape) == 3:
|
||||
B, L, D = data.shape
|
||||
@@ -141,21 +154,35 @@ def collate_rows_from_parquet_schema(rows,
|
||||
# Get tensor data from row using the existing helper function pattern
|
||||
shape_key = f"{tensor_name}_shape"
|
||||
bytes_key = f"{tensor_name}_bytes"
|
||||
dtype_key = f"{tensor_name}_dtype"
|
||||
|
||||
if shape_key in row and bytes_key in row:
|
||||
shape = row[shape_key]
|
||||
bytes_data = row[bytes_key]
|
||||
|
||||
# Get dtype from row, default to float32 for backward compatibility
|
||||
dtype_str = row.get(dtype_key, 'float32')
|
||||
# Convert dtype string to numpy dtype
|
||||
if 'float16' in dtype_str:
|
||||
np_dtype = np.float16
|
||||
elif 'float32' in dtype_str:
|
||||
np_dtype = np.float32
|
||||
elif 'float64' in dtype_str:
|
||||
np_dtype = np.float64
|
||||
else:
|
||||
# Default to float32 if dtype is unrecognized
|
||||
np_dtype = np.float32
|
||||
|
||||
if len(bytes_data) == 0:
|
||||
tensor = torch.zeros(0, dtype=torch.bfloat16)
|
||||
else:
|
||||
# Convert bytes to tensor using float32 as default
|
||||
# Convert bytes to tensor using the correct dtype
|
||||
if tensor_name == 'text_embedding' and (rng.random(
|
||||
) if rng else random.random()) < cfg_rate:
|
||||
data = np.zeros((512, 4096), dtype=np.float32)
|
||||
else:
|
||||
data = np.frombuffer(
|
||||
bytes_data, dtype=np.float32).reshape(shape).copy()
|
||||
bytes_data, dtype=np_dtype).reshape(shape).copy()
|
||||
tensor = torch.from_numpy(data)
|
||||
# if len(data.shape) == 3:
|
||||
# B, L, D = tensor.shape
|
||||
|
||||
@@ -53,6 +53,9 @@ def sequence_model_parallel_all_gather_with_unpad(
|
||||
Tensor: Gathered and unpadded tensor
|
||||
"""
|
||||
|
||||
# NCCL all_gather expects contiguous inputs.
|
||||
if not input_.is_contiguous():
|
||||
input_ = input_.contiguous()
|
||||
# First gather across all ranks
|
||||
gathered = get_sp_group().all_gather(input_, dim)
|
||||
|
||||
|
||||
@@ -194,14 +194,14 @@ class RotaryEmbedding(CustomOp):
|
||||
cos, sin = cos_sin.chunk(2, dim=-1)
|
||||
|
||||
query_shape = query.shape
|
||||
query = query.view(num_tokens, -1, self.head_size)
|
||||
query = query.reshape(num_tokens, -1, self.head_size)
|
||||
query_rot = query[..., :self.rotary_dim]
|
||||
query_pass = query[..., self.rotary_dim:]
|
||||
query_rot = _apply_rotary_emb(query_rot, cos, sin, self.is_neox_style)
|
||||
query = torch.cat((query_rot, query_pass), dim=-1).reshape(query_shape)
|
||||
|
||||
key_shape = key.shape
|
||||
key = key.view(num_tokens, -1, self.head_size)
|
||||
key = key.reshape(num_tokens, -1, self.head_size)
|
||||
key_rot = key[..., :self.rotary_dim]
|
||||
key_pass = key[..., self.rotary_dim:]
|
||||
key_rot = _apply_rotary_emb(key_rot, cos, sin, self.is_neox_style)
|
||||
|
||||
@@ -690,12 +690,15 @@ class SingleTokenRefiner(nn.Module):
|
||||
timestep_aware_representations = self.t_embedder(t)
|
||||
|
||||
# Get context-aware representations
|
||||
original_dtype = x.dtype
|
||||
target_dtype = self.c_embedder.fc_in.weight.dtype
|
||||
if x.dtype != target_dtype:
|
||||
x = x.to(dtype=target_dtype)
|
||||
if mask is None:
|
||||
context_aware_representations = x.mean(dim=1)
|
||||
else:
|
||||
mask_float = mask.float().unsqueeze(-1) # [B, L, 1]
|
||||
context_aware_representations = (x * mask_float).sum(dim=1) / mask_float.sum(dim=1)
|
||||
mask_float = mask.to(dtype=target_dtype).unsqueeze(-1) # [B, L, 1]
|
||||
context_aware_representations = (
|
||||
x * mask_float).sum(dim=1) / mask_float.sum(dim=1)
|
||||
|
||||
context_aware_representations = self.c_embedder(
|
||||
context_aware_representations)
|
||||
@@ -850,4 +853,4 @@ class FinalLayer(nn.Module):
|
||||
scale, shift = self.adaLN_modulation(c).chunk(2, dim=-1)
|
||||
x = self.norm_final(x) * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
x, _ = self.linear(x)
|
||||
return x
|
||||
return x
|
||||
|
||||
@@ -114,11 +114,14 @@ class HunyuanVideo15AttnBlock(nn.Module):
|
||||
"""
|
||||
seq_len = n_frame * n_hw
|
||||
mask = torch.full((seq_len, seq_len), float("-inf"), dtype=dtype, device=device)
|
||||
|
||||
for i in range(seq_len):
|
||||
i_frame = i // n_hw
|
||||
mask[i, : (i_frame + 1) * n_hw] = 0
|
||||
if batch_size is not None:
|
||||
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
|
||||
# mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
|
||||
mask = mask.unsqueeze(0).unsqueeze(1).expand(batch_size, 1, -1, -1)
|
||||
|
||||
return mask
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
@@ -131,6 +134,7 @@ class HunyuanVideo15AttnBlock(nn.Module):
|
||||
value = self.to_v(x)
|
||||
|
||||
batch_size, channels, frames, height, width = query.shape
|
||||
|
||||
|
||||
query = query.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous()
|
||||
key = key.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous()
|
||||
|
||||
@@ -92,6 +92,11 @@ class HunyuanVAEAttention(nn.Module):
|
||||
value = value.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
|
||||
|
||||
# Perform scaled dot-product attention
|
||||
# If attention_mask is provided, it needs to be expanded for multi-head attention
|
||||
if attention_mask is not None:
|
||||
# Expand mask from [batch_size, seq_len, seq_len] to [batch_size, num_heads, seq_len, seq_len]
|
||||
attention_mask = attention_mask.unsqueeze(1).expand(-1, self.heads, -1, -1)
|
||||
|
||||
hidden_states = F.scaled_dot_product_attention(query,
|
||||
key,
|
||||
value,
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.preprocess.preprocess_stages import (
|
||||
TextTransformStage, VideoTransformStage)
|
||||
from fastvideo.pipelines.stages import (EncodingStage, ImageEncodingStage,
|
||||
TextEncodingStage)
|
||||
from fastvideo.pipelines.stages.image_encoding import ImageVAEEncodingStage
|
||||
|
||||
|
||||
class PreprocessPipelineT2V(ComposedPipelineBase):
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "text_encoder_2", "tokenizer_2", "vae"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
assert fastvideo_args.preprocess_config is not None
|
||||
|
||||
self.add_stage(
|
||||
stage_name="text_transform_stage",
|
||||
stage=TextTransformStage(
|
||||
cfg_uncondition_drop_rate=fastvideo_args.preprocess_config.training_cfg_rate,
|
||||
seed=fastvideo_args.preprocess_config.seed,
|
||||
)
|
||||
)
|
||||
|
||||
text_encoders = [
|
||||
self.get_module("text_encoder"),
|
||||
self.get_module("text_encoder_2")
|
||||
]
|
||||
tokenizers = [
|
||||
self.get_module("tokenizer"),
|
||||
self.get_module("tokenizer_2")
|
||||
]
|
||||
|
||||
self.add_stage(
|
||||
stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=text_encoders,
|
||||
tokenizers=tokenizers,
|
||||
)
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="video_transform_stage",
|
||||
stage=VideoTransformStage(
|
||||
train_fps=fastvideo_args.preprocess_config.train_fps,
|
||||
num_frames=fastvideo_args.preprocess_config.num_frames,
|
||||
max_height=fastvideo_args.preprocess_config.max_height,
|
||||
max_width=fastvideo_args.preprocess_config.max_width,
|
||||
do_temporal_sample=fastvideo_args.preprocess_config.do_temporal_sample,
|
||||
)
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="video_encoding_stage",
|
||||
stage=EncodingStage(vae=self.get_module("vae"))
|
||||
)
|
||||
|
||||
EntryClass = [PreprocessPipelineT2V]
|
||||
@@ -0,0 +1,59 @@
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.preprocess.preprocess_stages import (
|
||||
TextTransformStage, VideoTransformStage)
|
||||
from fastvideo.pipelines.stages import (EncodingStage, ImageEncodingStage,
|
||||
TextEncodingStage)
|
||||
from fastvideo.pipelines.stages.image_encoding import ImageVAEEncodingStage
|
||||
|
||||
|
||||
class PreprocessPipelineT2V(ComposedPipelineBase):
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "text_encoder_2", "tokenizer_2", "vae"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
assert fastvideo_args.preprocess_config is not None
|
||||
|
||||
self.add_stage(
|
||||
stage_name="text_transform_stage",
|
||||
stage=TextTransformStage(
|
||||
cfg_uncondition_drop_rate=fastvideo_args.preprocess_config.training_cfg_rate,
|
||||
seed=fastvideo_args.preprocess_config.seed,
|
||||
)
|
||||
)
|
||||
|
||||
text_encoders = [
|
||||
self.get_module("text_encoder"),
|
||||
self.get_module("text_encoder_2")
|
||||
]
|
||||
tokenizers = [
|
||||
self.get_module("tokenizer"),
|
||||
self.get_module("tokenizer_2")
|
||||
]
|
||||
|
||||
self.add_stage(
|
||||
stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=text_encoders,
|
||||
tokenizers=tokenizers,
|
||||
)
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="video_transform_stage",
|
||||
stage=VideoTransformStage(
|
||||
train_fps=fastvideo_args.preprocess_config.train_fps,
|
||||
num_frames=fastvideo_args.preprocess_config.num_frames,
|
||||
max_height=fastvideo_args.preprocess_config.max_height,
|
||||
max_width=fastvideo_args.preprocess_config.max_width,
|
||||
do_temporal_sample=fastvideo_args.preprocess_config.do_temporal_sample,
|
||||
)
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="video_encoding_stage",
|
||||
stage=EncodingStage(vae=self.get_module("vae"))
|
||||
)
|
||||
|
||||
EntryClass = [PreprocessPipelineT2V]
|
||||
@@ -5,10 +5,6 @@ from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
from fastvideo.workflow.workflow_base import WorkflowBase
|
||||
|
||||
import os
|
||||
|
||||
os.environ["MASTER_PORT"] = "29513"
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,186 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
from huggingface_hub import snapshot_download
|
||||
from fastvideo.tests.ssim.test_inference_similarity import compute_video_ssim_torchvision
|
||||
|
||||
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
|
||||
|
||||
NUM_NODES = "1"
|
||||
MODEL_PATH = "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
|
||||
|
||||
DATA_DIR = "data"
|
||||
LOCAL_RAW_DATA_DIR = Path(DATA_DIR) / "cats15"
|
||||
LOCAL_PREPROCESSED_DATA_DIR = Path(DATA_DIR) / "cats_processed_t2v_hunyuan15"
|
||||
LOCAL_OUTPUT_DIR = Path(DATA_DIR) / "outputs_hunyuan15"
|
||||
|
||||
NUM_GPUS_PER_NODE_PREPROCESSING = "1"
|
||||
NUM_GPUS_PER_NODE_TRAINING = "8"
|
||||
|
||||
# entrypoints (adjust to what hunyuan15 scripts use)
|
||||
PREPROCESS_ENTRY = ["-m", "fastvideo.pipelines.preprocess.v1_preprocessing_new"]
|
||||
TRAIN_ENTRY_FILE_PATH = "fastvideo/training/hunyuan15_training_pipeline.py" # change if hunyuan15 uses a different file
|
||||
|
||||
LOCAL_TRAINING_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "training_dataset", "worker_0", "worker_0")
|
||||
LOCAL_VALIDATION_DATASET_FILE = os.path.join(LOCAL_RAW_DATA_DIR, "validation_prompt.json")
|
||||
|
||||
|
||||
def download_data():
|
||||
os.makedirs(DATA_DIR, exist_ok=True)
|
||||
snapshot_download(
|
||||
repo_id="wlsaidhi/cats-overfit-merged",
|
||||
local_dir=str(LOCAL_RAW_DATA_DIR),
|
||||
repo_type="dataset",
|
||||
resume_download=True,
|
||||
token=os.environ.get("HF_TOKEN"),
|
||||
)
|
||||
|
||||
# normalize dataset layout like your hunyuan test
|
||||
video_dir = LOCAL_RAW_DATA_DIR / "video"
|
||||
videos_dir = LOCAL_RAW_DATA_DIR / "videos"
|
||||
if video_dir.exists() and not videos_dir.exists():
|
||||
video_dir.rename(videos_dir)
|
||||
|
||||
src_val = LOCAL_RAW_DATA_DIR / "validation_prompt_1_sample.json"
|
||||
shutil.copy2(src_val, LOCAL_VALIDATION_DATASET_FILE)
|
||||
|
||||
src_v2c = LOCAL_RAW_DATA_DIR / "videos2caption_1_sample.json"
|
||||
shutil.copy2(src_v2c, LOCAL_RAW_DATA_DIR / "videos2caption.json")
|
||||
|
||||
|
||||
def run_preprocessing():
|
||||
if LOCAL_PREPROCESSED_DATA_DIR.exists():
|
||||
shutil.rmtree(LOCAL_PREPROCESSED_DATA_DIR)
|
||||
|
||||
env = os.environ.copy()
|
||||
env["PYTHONPATH"] = os.getcwd()
|
||||
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE_PREPROCESSING,
|
||||
*PREPROCESS_ENTRY,
|
||||
"--model-path", MODEL_PATH,
|
||||
"--mode", "preprocess",
|
||||
"--workload-type", "t2v",
|
||||
"--preprocess.video_loader_type", "torchvision",
|
||||
"--preprocess.dataset_type", "merged",
|
||||
"--preprocess.dataset_path", str(LOCAL_RAW_DATA_DIR),
|
||||
"--preprocess.dataset_output_dir", str(LOCAL_PREPROCESSED_DATA_DIR),
|
||||
"--preprocess.preprocess_video_batch_size", "1",
|
||||
"--preprocess.dataloader_num_workers", "0",
|
||||
"--preprocess.max_height", "480",
|
||||
"--preprocess.max_width", "832",
|
||||
"--preprocess.num_frames", "61",
|
||||
"--preprocess.train_fps", "24",
|
||||
"--preprocess.samples_per_file", "1",
|
||||
"--preprocess.flush_frequency", "1",
|
||||
"--preprocess.video_length_tolerance_range", "5",
|
||||
|
||||
# IMPORTANT: add a dtype/precision flag here if available in your codebase,
|
||||
# so saved arrays are fp16/fp32 instead of bf16.
|
||||
# Example (ONLY if supported by args):
|
||||
# "--preprocess.save_dtype", "fp16",
|
||||
# "--preprocess.embedding_dtype", "fp16",
|
||||
]
|
||||
|
||||
subprocess.run(cmd, check=True, env=env)
|
||||
|
||||
|
||||
def run_training():
|
||||
env = os.environ.copy()
|
||||
env["PYTHONPATH"] = os.getcwd()
|
||||
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE_TRAINING,
|
||||
TRAIN_ENTRY_FILE_PATH,
|
||||
"--model_path", MODEL_PATH,
|
||||
"--inference_mode", "False", # Required: must be False for training mode
|
||||
"--pretrained_model_name_or_path", MODEL_PATH,
|
||||
"--data_path", LOCAL_TRAINING_DATA_DIR,
|
||||
"--validation_dataset_file", LOCAL_VALIDATION_DATASET_FILE,
|
||||
|
||||
"--train_batch_size", "1",
|
||||
"--num_latent_t", "16", # (61-1)//4 + 1 = 16 for 61 frames with temporal_compression=4
|
||||
"--max_train_steps", "901",
|
||||
"--learning_rate", "1e-5",
|
||||
"--gradient_accumulation_steps", "1", # Required: default 0 causes empty training loop
|
||||
|
||||
# If your pipeline produces bf16 artifacts that later go to numpy/parquet,
|
||||
# switch mixed_precision to fp16 to avoid bf16 -> numpy issues.
|
||||
"--mixed_precision", "fp16",
|
||||
|
||||
# Distributed training configuration - must match nproc_per_node
|
||||
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--sp_size", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--tp_size", "1",
|
||||
"--hsdp_replicate_dim", "1",
|
||||
"--hsdp_shard_dim", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--train_sp_batch_size", "1",
|
||||
|
||||
# Dataloader
|
||||
"--dataloader_num_workers", "4",
|
||||
|
||||
# Validation & logging
|
||||
"--log_validation",
|
||||
"--validation_steps", "100",
|
||||
"--validation_sampling_steps", "50",
|
||||
"--validation_guidance_scale", "6.0", # Match Hunyuan15_480P_SamplingParam default
|
||||
|
||||
# Checkpointing
|
||||
"--weight_only_checkpointing_steps", "6000",
|
||||
"--training_state_checkpointing_steps", "6000",
|
||||
"--checkpoints_total_limit", "3",
|
||||
|
||||
# Optimizer settings
|
||||
"--weight_decay", "0.01",
|
||||
"--max_grad_norm", "1.0",
|
||||
|
||||
# Training config
|
||||
"--ema_start_step", "0",
|
||||
"--training_cfg_rate", "0.0",
|
||||
"--dit_precision", "fp32",
|
||||
"--enable_gradient_checkpointing_type", "full",
|
||||
|
||||
# Output
|
||||
"--output_dir", str(LOCAL_OUTPUT_DIR),
|
||||
"--tracker_project_name", "hunyuan15_finetune_overfit_ci",
|
||||
"--num_height", "480",
|
||||
"--num_width", "832",
|
||||
"--num_frames", "61",
|
||||
"--embedded_cfg_scale", "6.0",
|
||||
"--num_euler_timesteps", "50",
|
||||
]
|
||||
|
||||
subprocess.run(cmd, check=True, env=env)
|
||||
|
||||
|
||||
def test_e2e_hunyuan15_overfit_single_sample():
|
||||
os.environ["WANDB_MODE"] = "online"
|
||||
os.environ["WANDB_API_KEY"] = "8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde"
|
||||
|
||||
# download_data()
|
||||
# run_preprocessing()
|
||||
run_training()
|
||||
|
||||
reference_video_file = os.path.join(os.path.dirname(__file__), "reference_video_1_sample_hy15.mp4")
|
||||
final_validation_video_file = os.path.join(LOCAL_OUTPUT_DIR, "validation_step_900_inference_steps_50_video_0.mp4")
|
||||
|
||||
assert os.path.exists(reference_video_file)
|
||||
assert os.path.exists(final_validation_video_file)
|
||||
|
||||
mean_ssim, min_ssim, max_ssim = compute_video_ssim_torchvision(
|
||||
reference_video_file,
|
||||
final_validation_video_file,
|
||||
use_ms_ssim=True,
|
||||
)
|
||||
assert max_ssim > 0.5, f"Max SSIM is below threshold: {max_ssim}"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_e2e_hunyuan15_overfit_single_sample()
|
||||
@@ -0,0 +1,218 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
from huggingface_hub import snapshot_download
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
from fastvideo.tests.ssim.test_inference_similarity import compute_video_ssim_torchvision
|
||||
|
||||
# Example of using Hunyuan preprocessing and training pipeline
|
||||
|
||||
# Import the training pipeline
|
||||
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
|
||||
|
||||
NUM_NODES = "1"
|
||||
MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
|
||||
|
||||
# preprocessing
|
||||
DATA_DIR = "data"
|
||||
LOCAL_RAW_DATA_DIR = Path(os.path.join(DATA_DIR, "cats"))
|
||||
NUM_GPUS_PER_NODE_PREPROCESSING = "1"
|
||||
PREPROCESSING_ENTRY_FILE_PATH = "fastvideo/pipelines/preprocess/v1_preprocessing_new.py"
|
||||
|
||||
LOCAL_PREPROCESSED_DATA_DIR = Path(os.path.join(DATA_DIR, "cats_processed_t2v_hunyuan"))
|
||||
|
||||
|
||||
# training
|
||||
NUM_GPUS_PER_NODE_TRAINING = "4"
|
||||
TRAINING_ENTRY_FILE_PATH = "fastvideo/training/hunyuan_training_pipeline.py"
|
||||
# New preprocessing pipeline creates files in training_dataset/worker_0/worker_0/
|
||||
LOCAL_TRAINING_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "training_dataset", "worker_0", "worker_0")
|
||||
LOCAL_VALIDATION_DATASET_FILE = os.path.join(LOCAL_RAW_DATA_DIR, "validation_prompt.json")
|
||||
LOCAL_OUTPUT_DIR = Path(os.path.join(DATA_DIR, "outputs_hunyuan"))
|
||||
|
||||
def download_data():
|
||||
# create the data dir if it doesn't exist
|
||||
data_dir = Path(DATA_DIR)
|
||||
|
||||
print(f"Creating data directory at {data_dir}")
|
||||
os.makedirs(data_dir, exist_ok=True)
|
||||
|
||||
print(f"Downloading raw dataset to {LOCAL_RAW_DATA_DIR}...")
|
||||
try:
|
||||
result = snapshot_download(
|
||||
repo_id="wlsaidhi/cats-overfit-merged",
|
||||
local_dir=str(LOCAL_RAW_DATA_DIR),
|
||||
repo_type="dataset",
|
||||
resume_download=True,
|
||||
token=os.environ.get("HF_TOKEN"), # In case authentication is needed
|
||||
)
|
||||
print(f"Download completed successfully. Files downloaded to: {result}")
|
||||
|
||||
# Verify the download
|
||||
if not LOCAL_RAW_DATA_DIR.exists():
|
||||
raise RuntimeError(f"Download appeared to succeed but {LOCAL_RAW_DATA_DIR} does not exist")
|
||||
|
||||
# List downloaded files
|
||||
print("Downloaded files:")
|
||||
for file in LOCAL_RAW_DATA_DIR.rglob("*"):
|
||||
if file.is_file():
|
||||
print(f" - {file.relative_to(LOCAL_RAW_DATA_DIR)}")
|
||||
|
||||
# Rename video directory if needed (dataset has 'video' but preprocessing expects 'videos')
|
||||
video_dir = os.path.join(LOCAL_RAW_DATA_DIR, "video")
|
||||
videos_dir = os.path.join(LOCAL_RAW_DATA_DIR, "videos")
|
||||
if os.path.exists(video_dir) and not os.path.exists(videos_dir):
|
||||
print(f"Renaming video directory to videos...")
|
||||
os.rename(video_dir, videos_dir)
|
||||
|
||||
# Copy validation file to expected name
|
||||
source_validation_file = os.path.join(LOCAL_RAW_DATA_DIR, "validation_prompt_1_sample.json")
|
||||
target_validation_file = LOCAL_VALIDATION_DATASET_FILE
|
||||
|
||||
if os.path.exists(source_validation_file):
|
||||
print(f"Copying {source_validation_file} to {target_validation_file}...")
|
||||
shutil.copy2(source_validation_file, target_validation_file)
|
||||
else:
|
||||
raise FileNotFoundError(f"Source validation file not found: {source_validation_file}")
|
||||
|
||||
# Override videos2caption.json with the 1-sample version for preprocessing
|
||||
# The new preprocessing pipeline automatically reads videos2caption.json from the dataset path
|
||||
source_videos2caption = os.path.join(LOCAL_RAW_DATA_DIR, "videos2caption_1_sample.json")
|
||||
target_videos2caption = os.path.join(LOCAL_RAW_DATA_DIR, "videos2caption.json")
|
||||
|
||||
if os.path.exists(source_videos2caption):
|
||||
print(f"Overriding videos2caption.json with 1-sample version...")
|
||||
shutil.copy2(source_videos2caption, target_videos2caption)
|
||||
else:
|
||||
raise FileNotFoundError(f"Source videos2caption file not found: {source_videos2caption}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during download: {str(e)}")
|
||||
raise
|
||||
|
||||
|
||||
def run_preprocessing():
|
||||
# remove the local_preprocessed_data_dir if it exists
|
||||
if LOCAL_PREPROCESSED_DATA_DIR.exists():
|
||||
print(f"Removing local_preprocessed_data_dir: {LOCAL_PREPROCESSED_DATA_DIR}")
|
||||
shutil.rmtree(LOCAL_PREPROCESSED_DATA_DIR)
|
||||
env = os.environ.copy()
|
||||
env['PYTHONPATH'] = os.getcwd()
|
||||
# Run torchrun command using the new preprocessing pipeline
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE_PREPROCESSING,
|
||||
"-m", "fastvideo.pipelines.preprocess.v1_preprocessing_new",
|
||||
"--model-path", MODEL_PATH,
|
||||
"--mode", "preprocess",
|
||||
"--workload-type", "t2v",
|
||||
"--preprocess.video_loader_type", "torchvision",
|
||||
"--preprocess.dataset_type", "merged",
|
||||
"--preprocess.dataset_path", str(LOCAL_RAW_DATA_DIR),
|
||||
"--preprocess.dataset_output_dir", str(LOCAL_PREPROCESSED_DATA_DIR),
|
||||
"--preprocess.preprocess_video_batch_size", "1",
|
||||
"--preprocess.dataloader_num_workers", "0",
|
||||
"--preprocess.max_height", "480",
|
||||
"--preprocess.max_width", "832",
|
||||
"--preprocess.num_frames", "77",
|
||||
"--preprocess.train_fps", "16",
|
||||
"--preprocess.samples_per_file", "1",
|
||||
"--preprocess.flush_frequency", "1",
|
||||
"--preprocess.video_length_tolerance_range", "5",
|
||||
]
|
||||
|
||||
process = subprocess.run(cmd, check=True, env=env)
|
||||
|
||||
|
||||
def run_training():
|
||||
env = os.environ.copy()
|
||||
env['PYTHONPATH'] = os.getcwd()
|
||||
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE_TRAINING,
|
||||
TRAINING_ENTRY_FILE_PATH,
|
||||
"--model_path", MODEL_PATH,
|
||||
"--inference_mode", "False",
|
||||
"--pretrained_model_name_or_path", MODEL_PATH,
|
||||
"--data_path", LOCAL_TRAINING_DATA_DIR,
|
||||
"--validation_dataset_file", LOCAL_VALIDATION_DATASET_FILE,
|
||||
"--train_batch_size", "1",
|
||||
"--num_latent_t", "8",
|
||||
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--sp_size", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--tp_size", "1",
|
||||
"--hsdp_replicate_dim", "1",
|
||||
"--hsdp_shard_dim", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--train_sp_batch_size", "1",
|
||||
"--dataloader_num_workers", "10",
|
||||
"--gradient_accumulation_steps", "1",
|
||||
"--max_train_steps", "901",
|
||||
"--learning_rate", "1e-5",
|
||||
"--mixed_precision", "bf16",
|
||||
"--weight_only_checkpointing_steps", "6000",
|
||||
"--training_state_checkpointing_steps", "6000",
|
||||
"--validation_steps", "100",
|
||||
"--validation_sampling_steps", "50",
|
||||
"--log_validation",
|
||||
"--checkpoints_total_limit", "3",
|
||||
"--ema_start_step", "0",
|
||||
"--training_cfg_rate", "0.0",
|
||||
"--output_dir", str(LOCAL_OUTPUT_DIR),
|
||||
"--tracker_project_name", "hunyuan_finetune_overfit_ci",
|
||||
"--num_height", "480",
|
||||
"--num_width", "832",
|
||||
"--num_frames", "81",
|
||||
"--validation_guidance_scale", "1.0",
|
||||
"--embedded_cfg_scale", "6.0",
|
||||
"--num_euler_timesteps", "50",
|
||||
"--multi_phased_distill_schedule", "4000-1",
|
||||
"--weight_decay", "0.01",
|
||||
"--not_apply_cfg_solver",
|
||||
"--dit_precision", "fp32",
|
||||
"--max_grad_norm", "1.0",
|
||||
"--flow_shift", "7",
|
||||
]
|
||||
|
||||
print(f"Running training with command: {cmd}")
|
||||
process = subprocess.run(cmd, check=True, env=env)
|
||||
|
||||
|
||||
def test_e2e_hunyuan_overfit_single_sample():
|
||||
os.environ["WANDB_MODE"] = "online"
|
||||
|
||||
#download_data()
|
||||
run_preprocessing()
|
||||
run_training()
|
||||
|
||||
reference_video_file = os.path.join(os.path.dirname(__file__), "reference_video_1_sample_v0.mp4")
|
||||
print(f"reference_video_file: {reference_video_file}")
|
||||
final_validation_video_file = os.path.join(LOCAL_OUTPUT_DIR, "validation_step_900_inference_steps_50_video_0.mp4")
|
||||
print(f"final_validation_video_file: {final_validation_video_file}")
|
||||
|
||||
|
||||
# Ensure both files exist
|
||||
assert os.path.exists(reference_video_file), f"Reference video not found at {reference_video_file}"
|
||||
assert os.path.exists(final_validation_video_file), f"Validation video not found at {final_validation_video_file}"
|
||||
|
||||
#Compute SSIM
|
||||
mean_ssim, min_ssim, max_ssim = compute_video_ssim_torchvision(
|
||||
reference_video_file,
|
||||
final_validation_video_file,
|
||||
use_ms_ssim=True # Using MS-SSIM for better quality assessment
|
||||
)
|
||||
|
||||
print("\n===== SSIM Results for Step 900 Validation (Hunyuan) =====")
|
||||
print(f"Mean MS-SSIM: {mean_ssim:.4f}")
|
||||
print(f"Min MS-SSIM: {min_ssim:.4f}")
|
||||
print(f"Max MS-SSIM: {max_ssim:.4f}")
|
||||
|
||||
assert max_ssim > 0.5, f"Max SSIM is below 0.5: {max_ssim}"
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_e2e_hunyuan_overfit_single_sample()
|
||||
@@ -19,12 +19,16 @@ if "A40" in device_name:
|
||||
device_reference_folder = "A40" + device_reference_folder_suffix
|
||||
elif "L40S" in device_name:
|
||||
device_reference_folder = "L40S" + device_reference_folder_suffix
|
||||
elif "H100" in device_name or "NVIDIA H100" in device_name:
|
||||
device_reference_folder = "H100" + device_reference_folder_suffix
|
||||
|
||||
else:
|
||||
# device_reference_folder = "L40S" + device_reference_folder_suffix
|
||||
logger.warning(f"Unsupported device for ssim tests: {device_name}")
|
||||
# raise ValueError(f"Unsupported device for ssim tests: {device_name}")
|
||||
|
||||
# Base parameters from the shell script
|
||||
|
||||
HUNYUAN_PARAMS = {
|
||||
"num_gpus": 4,
|
||||
"model_path": "FastVideo/FastHunyuan-diffusers",
|
||||
|
||||
@@ -0,0 +1,81 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
|
||||
from fastvideo.pipelines.basic.hunyuan15.hunyuan15_pipeline import HunyuanVideo15Pipeline
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class Hunyuan15TrainingPipeline(TrainingPipeline):
|
||||
"""
|
||||
A training pipeline for HunyuanVideo-1.5.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer", "vae"]
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
if training_args.enable_gradient_checkpointing_type is None:
|
||||
training_args.enable_gradient_checkpointing_type = "full"
|
||||
super().initialize_training_pipeline(training_args)
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
pass
|
||||
# self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
|
||||
# shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
# reserved for future refactors
|
||||
pass
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline (Hunyuan15)...")
|
||||
args_copy = deepcopy(training_args)
|
||||
args_copy.inference_mode = True
|
||||
|
||||
validation_pipeline = HunyuanVideo15Pipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=args_copy, # type: ignore
|
||||
inference_mode=True,
|
||||
loaded_modules={
|
||||
# reuse the training transformer weights for validation sampling
|
||||
"transformer": self.get_module("transformer"),
|
||||
},
|
||||
tp_size=training_args.tp_size,
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus,
|
||||
pin_cpu_memory=training_args.pin_cpu_memory,
|
||||
dit_cpu_offload=True,
|
||||
# dit_layerwise_offload=True,
|
||||
use_fsdp_inference=True,
|
||||
)
|
||||
|
||||
self.validation_pipeline = validation_pipeline
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting Hunyuan15 training pipeline...")
|
||||
|
||||
pipeline = Hunyuan15TrainingPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path,
|
||||
args=args,
|
||||
)
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("Training pipeline done")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
|
||||
args.dit_cpu_offload = False
|
||||
main(args)
|
||||
@@ -0,0 +1,74 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.pipelines.basic.hunyuan.hunyuan_pipeline import HunyuanVideoPipeline
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class HunyuanTrainingPipeline(TrainingPipeline):
|
||||
"""
|
||||
A training pipeline for Hunyuan.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer", "vae"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
|
||||
shift=fastvideo_args.pipeline_config.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
|
||||
validation_pipeline = HunyuanVideoPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=args_copy, # type: ignore
|
||||
inference_mode=True,
|
||||
loaded_modules={
|
||||
"transformer": self.get_module("transformer"),
|
||||
},
|
||||
tp_size=training_args.tp_size,
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus,
|
||||
pin_cpu_memory=training_args.pin_cpu_memory,
|
||||
dit_cpu_offload=True)
|
||||
|
||||
self.validation_pipeline = validation_pipeline
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting training pipeline...")
|
||||
|
||||
pipeline = HunyuanTrainingPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("Training pipeline done")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.fastvideo_args import TrainingArgs
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
args.dit_cpu_offload = False
|
||||
main(args)
|
||||
@@ -1,5 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import asdict
|
||||
import inspect
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
@@ -174,6 +175,16 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
|
||||
text_padding_length = training_args.pipeline_config.text_encoder_configs[
|
||||
0].arch_config.text_len # type: ignore[attr-defined]
|
||||
if not text_padding_length:
|
||||
text_max_lengths = getattr(training_args.pipeline_config,
|
||||
"text_encoder_max_lengths", None)
|
||||
if text_max_lengths:
|
||||
text_padding_length = text_max_lengths[0]
|
||||
else:
|
||||
text_padding_length = 512
|
||||
|
||||
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
|
||||
training_args.data_path,
|
||||
training_args.train_batch_size,
|
||||
@@ -181,9 +192,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
num_data_workers=training_args.dataloader_num_workers,
|
||||
cfg_rate=training_args.training_cfg_rate,
|
||||
drop_last=True,
|
||||
text_padding_length=training_args.pipeline_config.
|
||||
text_encoder_configs[0].arch_config.
|
||||
text_len, # type: ignore[attr-defined]
|
||||
text_padding_length=text_padding_length,
|
||||
seed=self.seed)
|
||||
|
||||
self.noise_scheduler = noise_scheduler
|
||||
@@ -272,11 +281,27 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
def _normalize_dit_input(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
# TODO(will): support other models
|
||||
# Automatically detect model type based on VAE class name
|
||||
vae = self.get_module("vae")
|
||||
vae_class_name = vae.__class__.__name__
|
||||
|
||||
# Map VAE class names to model types
|
||||
if 'Hunyuan' in vae_class_name:
|
||||
model_type = 'hunyuan'
|
||||
elif 'Wan' in vae_class_name or 'AutoencoderKLCausal3D' in vae_class_name:
|
||||
model_type = 'wan'
|
||||
else:
|
||||
# Default to checking for latents_mean attribute
|
||||
if hasattr(vae, 'latents_mean'):
|
||||
model_type = 'wan'
|
||||
else:
|
||||
model_type = 'hunyuan'
|
||||
|
||||
with self.tracker.timed("timing/normalize_input"):
|
||||
training_batch.latents = normalize_dit_input(
|
||||
'wan',
|
||||
model_type,
|
||||
training_batch.latents,
|
||||
self.get_module("vae"),
|
||||
vae,
|
||||
)
|
||||
return training_batch
|
||||
|
||||
@@ -428,6 +453,87 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
# device=training_batch.noisy_model_input.device,
|
||||
# dtype=torch.bfloat16)
|
||||
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
|
||||
if "return_dict" in input_kwargs:
|
||||
try:
|
||||
sig = inspect.signature(current_model.forward)
|
||||
except (TypeError, ValueError):
|
||||
sig = None
|
||||
if sig is not None:
|
||||
params = sig.parameters.values()
|
||||
accepts_kwargs = any(param.kind == inspect.Parameter.VAR_KEYWORD
|
||||
for param in params)
|
||||
if ("return_dict" not in sig.parameters and not accepts_kwargs):
|
||||
input_kwargs = dict(input_kwargs)
|
||||
input_kwargs.pop("return_dict", None)
|
||||
base_model = current_model
|
||||
for attr in ("module", "_orig_mod"):
|
||||
if hasattr(base_model, attr):
|
||||
base_model = getattr(base_model, attr)
|
||||
if hasattr(base_model, "_fsdp_wrapped_module"):
|
||||
base_model = base_model._fsdp_wrapped_module
|
||||
if hasattr(base_model, "module"):
|
||||
base_model = base_model.module
|
||||
model_config = getattr(base_model, "config", None)
|
||||
if model_config is None:
|
||||
model_config = getattr(current_model, "config", None)
|
||||
if model_config is not None and hasattr(model_config,
|
||||
"text_embed_2_dim"):
|
||||
input_kwargs = dict(input_kwargs)
|
||||
text_states = training_batch.encoder_hidden_states
|
||||
text_mask = training_batch.encoder_attention_mask
|
||||
if text_states is None or text_mask is None:
|
||||
raise ValueError(
|
||||
"HunyuanVideo15 training requires text embeddings and masks"
|
||||
)
|
||||
if not isinstance(input_kwargs.get("encoder_hidden_states"),
|
||||
(list, tuple)):
|
||||
batch_size = text_states.shape[0]
|
||||
text_2_len = 256
|
||||
if hasattr(self.training_args.pipeline_config,
|
||||
"text_encoder_max_lengths"):
|
||||
text_2_len = self.training_args.pipeline_config.text_encoder_max_lengths[
|
||||
1]
|
||||
text_2_dim = model_config.text_embed_2_dim
|
||||
text_2 = torch.zeros((batch_size, text_2_len, text_2_dim),
|
||||
device=text_states.device,
|
||||
dtype=text_states.dtype)
|
||||
input_kwargs["encoder_hidden_states"] = [text_states, text_2]
|
||||
if not isinstance(input_kwargs.get("encoder_attention_mask"),
|
||||
(list, tuple)):
|
||||
batch_size = text_mask.shape[0]
|
||||
text_2_len = 256
|
||||
if hasattr(self.training_args.pipeline_config,
|
||||
"text_encoder_max_lengths"):
|
||||
text_2_len = self.training_args.pipeline_config.text_encoder_max_lengths[
|
||||
1]
|
||||
text_mask_2 = torch.zeros((batch_size, text_2_len),
|
||||
device=text_mask.device,
|
||||
dtype=text_mask.dtype)
|
||||
input_kwargs["encoder_attention_mask"] = [
|
||||
text_mask, text_mask_2
|
||||
]
|
||||
if "encoder_hidden_states_image" not in input_kwargs:
|
||||
image_embeds = training_batch.image_embeds
|
||||
if image_embeds is None or image_embeds.numel() == 0:
|
||||
batch_size = text_states.shape[0]
|
||||
image_len = 729
|
||||
image_dim = model_config.image_embed_dim
|
||||
image_embeds = torch.zeros(
|
||||
(batch_size, image_len, image_dim),
|
||||
device=text_states.device,
|
||||
dtype=text_states.dtype,
|
||||
)
|
||||
input_kwargs["encoder_hidden_states_image"] = [image_embeds]
|
||||
hidden_states = input_kwargs.get("hidden_states")
|
||||
if (isinstance(hidden_states, torch.Tensor)
|
||||
and hidden_states.shape[1] != model_config.in_channels):
|
||||
batch_size, _, t, h, w = hidden_states.shape
|
||||
video_latent = torch.zeros((batch_size, 1, t, h, w),
|
||||
device=hidden_states.device,
|
||||
dtype=hidden_states.dtype)
|
||||
zeros_latent = torch.zeros_like(hidden_states)
|
||||
input_kwargs["hidden_states"] = torch.cat(
|
||||
[hidden_states, video_latent, zeros_latent], dim=1)
|
||||
|
||||
with self.tracker.timed("timing/forward_backward"), set_forward_context(
|
||||
current_timestep=training_batch.current_timestep,
|
||||
@@ -650,7 +756,13 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
}
|
||||
metrics["batch_size"] = int(training_batch.raw_latent_shape[0])
|
||||
|
||||
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
|
||||
patch_size = self.training_args.pipeline_config.dit_config.patch_size
|
||||
if isinstance(patch_size, tuple):
|
||||
patch_t, patch_h, patch_w = patch_size
|
||||
else:
|
||||
patch_t = patch_size
|
||||
patch_h = patch_size
|
||||
patch_w = patch_size
|
||||
seq_len = (training_batch.raw_latent_shape[2] // patch_t) * (
|
||||
training_batch.raw_latent_shape[3] //
|
||||
patch_h) * (training_batch.raw_latent_shape[4] // patch_w)
|
||||
@@ -663,7 +775,13 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
metrics["hidden_dim"] = arch_config.hidden_size
|
||||
metrics["num_layers"] = arch_config.num_layers
|
||||
metrics["ffn_dim"] = arch_config.ffn_dim
|
||||
ffn_dim = getattr(arch_config, "ffn_dim", None)
|
||||
if ffn_dim is None:
|
||||
mlp_ratio = getattr(arch_config, "mlp_ratio", None)
|
||||
if mlp_ratio and arch_config.hidden_size:
|
||||
ffn_dim = int(arch_config.hidden_size * mlp_ratio)
|
||||
if ffn_dim is not None:
|
||||
metrics["ffn_dim"] = ffn_dim
|
||||
|
||||
self.tracker.log(metrics, step)
|
||||
if step % self.training_args.training_state_checkpointing_steps == 0:
|
||||
@@ -838,8 +956,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
step_captions.append(batch.prompt)
|
||||
|
||||
# Run validation inference
|
||||
# Use the validation pipeline args to match its loaded modules.
|
||||
validation_args = self.validation_pipeline.fastvideo_args
|
||||
if validation_args is None:
|
||||
validation_args = training_args
|
||||
output_batch = self.validation_pipeline.forward(
|
||||
batch, training_args)
|
||||
batch, validation_args)
|
||||
samples = output_batch.output
|
||||
|
||||
if self.rank_in_sp_group != 0:
|
||||
|
||||
@@ -849,7 +849,14 @@ def load_distillation_checkpoint(
|
||||
|
||||
def normalize_dit_input(model_type, latents, vae) -> torch.Tensor:
|
||||
if model_type == "hunyuan_hf" or model_type == "hunyuan":
|
||||
return latents * 0.476986
|
||||
scaling_factor = vae.scaling_factor
|
||||
if scaling_factor is None:
|
||||
raise AttributeError(
|
||||
"Hunyuan VAE must define scaling_factor for training.")
|
||||
if isinstance(scaling_factor, torch.Tensor):
|
||||
scaling_factor = scaling_factor.to(device=latents.device,
|
||||
dtype=latents.dtype)
|
||||
return latents * scaling_factor
|
||||
elif model_type == "wan":
|
||||
latents_mean = torch.tensor(vae.latents_mean)
|
||||
latents_std = 1.0 / torch.tensor(vae.latents_std)
|
||||
|
||||
@@ -200,8 +200,14 @@ class ParquetDatasetSaver:
|
||||
value = getattr(batch, key.name)
|
||||
if isinstance(value, list):
|
||||
for idx in range(len(value)):
|
||||
if isinstance(value[idx], torch.Tensor):
|
||||
value[idx] = value[idx].cpu().numpy()
|
||||
# if isinstance(value[idx], torch.Tensor):
|
||||
# value[idx] = value[idx].cpu().numpy()
|
||||
t = value[idx]
|
||||
if isinstance(t, torch.Tensor):
|
||||
if t.dtype == torch.bfloat16:
|
||||
t = t.float() # or t.to(torch.float16) if you want smaller
|
||||
value[idx] = t.cpu().numpy()
|
||||
|
||||
elif isinstance(value, torch.Tensor):
|
||||
value = value.cpu().numpy()
|
||||
setattr(batch, key.name, value)
|
||||
@@ -288,10 +294,17 @@ def build_dataset(preprocess_config: PreprocessConfig, split: str,
|
||||
split=split)
|
||||
column_names = dataset.column_names
|
||||
# rename columns to match the schema
|
||||
if "cap" in column_names:
|
||||
if "cap" in column_names and "caption" not in column_names:
|
||||
dataset = dataset.rename_column("cap", "caption")
|
||||
if "path" in column_names:
|
||||
dataset = dataset.rename_column("path", "name")
|
||||
|
||||
# Add num_frames column if it doesn't exist but fps and duration do
|
||||
if 'num_frames' not in column_names and 'fps' in column_names and 'duration' in column_names:
|
||||
def add_num_frames(item: dict[str, Any]) -> dict[str, Any]:
|
||||
item['num_frames'] = int(item['fps'] * item['duration'])
|
||||
return item
|
||||
dataset = dataset.map(add_num_frames)
|
||||
|
||||
dataset = dataset.filter(validator)
|
||||
dataset = dataset.shard(num_shards=get_world_size(),
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
#!/bin/bash
|
||||
# Kill all GPU processes shown by nvidia-smi
|
||||
|
||||
pids=$(nvidia-smi --query-compute-apps=pid --format=csv,noheader 2>/dev/null | tr -d ' ')
|
||||
|
||||
if [ -z "$pids" ]; then
|
||||
echo "No GPU processes found."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
echo "Killing GPU processes: $pids"
|
||||
for pid in $pids; do
|
||||
kill -9 "$pid" 2>/dev/null && echo "Killed PID $pid" || echo "Failed to kill PID $pid"
|
||||
done
|
||||
Reference in New Issue
Block a user