Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bf5726cd06 | ||
|
|
57cfd16136 |
@@ -1,21 +1,29 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from collections import defaultdict
|
||||
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
import tqdm
|
||||
from einops import rearrange
|
||||
from torch import distributed as dist
|
||||
from torch.utils.data import IterableDataset, get_worker_info
|
||||
from torch.utils.data import Dataset
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
|
||||
get_sp_group)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
# Path to your dataset
|
||||
dataset_path = "/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/train/"
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ParquetVideoTextDataset(IterableDataset):
|
||||
class ParquetVideoTextDataset(Dataset):
|
||||
"""Efficient loader for video-text data from a directory of Parquet files."""
|
||||
|
||||
def __init__(self,
|
||||
@@ -24,237 +32,185 @@ class ParquetVideoTextDataset(IterableDataset):
|
||||
rank: int = 0,
|
||||
world_size: int = 1,
|
||||
cfg_rate: float = 0.0,
|
||||
num_latent_t: int = 2):
|
||||
num_latent_t: int = 2,
|
||||
seed: int = 0):
|
||||
super().__init__()
|
||||
self.path = str(path)
|
||||
self.batch_size = batch_size
|
||||
self.rank = rank
|
||||
self.world_size = world_size
|
||||
self.local_rank = get_sequence_model_parallel_rank()
|
||||
self.sp_world_size = world_size
|
||||
self.world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
self.cfg_rate = cfg_rate
|
||||
self.num_latent_t = num_latent_t
|
||||
self.local_indices = None
|
||||
self.plan_output_dir = os.path.join(self.path, "data_plan.json")
|
||||
|
||||
# Find all parquet files recursively
|
||||
print(f"Scanning for parquet files in {self.path}")
|
||||
self.parquet_files = []
|
||||
for root, _, files in os.walk(self.path):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
self.parquet_files.append(os.path.join(root, file))
|
||||
# Sort files for consistent ordering
|
||||
self.parquet_files.sort()
|
||||
ranks = get_sp_group().ranks
|
||||
group_ranks = [None for _ in range(self.world_size)]
|
||||
torch.distributed.all_gather_object(group_ranks, ranks)
|
||||
|
||||
# Distribute files among workers
|
||||
# drop last unenven files
|
||||
print(f"Total files: {len(self.parquet_files)}")
|
||||
total_files = len(self.parquet_files)
|
||||
base_count = total_files // world_size
|
||||
extra_files = total_files % world_size
|
||||
if rank == 0:
|
||||
# If a plan already exists, then skip creating a new plan
|
||||
# This will be useful when resume training
|
||||
if os.path.exists(self.plan_output_dir):
|
||||
print(f"Using existing plan from {self.plan_output_dir}")
|
||||
return
|
||||
|
||||
if rank < extra_files:
|
||||
start_idx = rank * (base_count + 1)
|
||||
end_idx = start_idx + base_count + 1
|
||||
else:
|
||||
start_idx = rank * base_count + extra_files
|
||||
end_idx = start_idx + base_count
|
||||
# Find all parquet files recursively, and record num_rows for each file
|
||||
print(f"Scanning for parquet files in {self.path}")
|
||||
metadatas = []
|
||||
for root, _, files in os.walk(self.path):
|
||||
for file in sorted(files):
|
||||
if file.endswith('.parquet'):
|
||||
file_path = os.path.join(root, file)
|
||||
num_rows = pq.ParquetFile(file_path).metadata.num_rows
|
||||
for row_idx in range(num_rows):
|
||||
metadatas.append((file_path, row_idx))
|
||||
|
||||
self.parquet_files = self.parquet_files[start_idx:end_idx]
|
||||
# Generate the plan that distribute rows among workers
|
||||
random.seed(seed)
|
||||
random.shuffle(metadatas)
|
||||
|
||||
print(f"Files assigned to rank {rank}: {len(self.parquet_files)}")
|
||||
if len(self.parquet_files) > 0:
|
||||
print(f"First file: {self.parquet_files[0]}")
|
||||
print(f"Last file: {self.parquet_files[-1]}")
|
||||
# Get all sp groups
|
||||
# e.g. if num_gpus = 4, sp_size = 2
|
||||
# group_ranks = [(0, 1), (2, 3)]
|
||||
# We will assign the same batches of data to ranks in the same sp group, and we'll assign different batches to ranks in different sp groups
|
||||
# e.g. plan = {0: [row 1, row 4], 1: [row 1, row 4], 2: [row 2, row 3], 3: [row 2, row 3]}
|
||||
group_ranks = list(set(tuple(r) for r in group_ranks))
|
||||
num_sp_groups = len(group_ranks)
|
||||
plan = defaultdict(list)
|
||||
for idx, metadata in enumerate(metadatas):
|
||||
sp_group_idx = idx % num_sp_groups
|
||||
for global_rank in group_ranks[sp_group_idx]:
|
||||
plan[global_rank].append(metadata)
|
||||
|
||||
# Initialize current file index
|
||||
self.current_file_idx = 0
|
||||
self.current_reader = None
|
||||
self.current_batches = None
|
||||
self.total_samples = 0
|
||||
with open(self.plan_output_dir, "w") as f:
|
||||
json.dump(plan, f)
|
||||
|
||||
def _open_next_file(self):
|
||||
"""Open the next parquet file for reading."""
|
||||
num_workers = get_worker_info().num_workers
|
||||
worker_id = get_worker_info().id
|
||||
total_files = len(self.parquet_files)
|
||||
base_count = total_files // num_workers
|
||||
extra_files = total_files % num_workers
|
||||
|
||||
if worker_id < extra_files:
|
||||
start_idx = worker_id * (base_count + 1)
|
||||
end_idx = start_idx + base_count + 1
|
||||
else:
|
||||
start_idx = worker_id * base_count + extra_files
|
||||
end_idx = start_idx + base_count
|
||||
|
||||
worker_parquet_files = self.parquet_files[start_idx:end_idx]
|
||||
if self.current_file_idx >= len(worker_parquet_files):
|
||||
print(
|
||||
f"Rank {self.rank}, Worker {worker_id}: No more files to open (current_idx={self.current_file_idx}, total_files={len(worker_parquet_files)})"
|
||||
)
|
||||
return False
|
||||
|
||||
if self.current_reader is not None:
|
||||
self.current_reader.close()
|
||||
|
||||
file_path = worker_parquet_files[self.current_file_idx]
|
||||
print(
|
||||
f"Rank {self.rank}, Worker {worker_id}: Opening file {self.current_file_idx + 1}/{len(worker_parquet_files)}: {file_path}"
|
||||
)
|
||||
|
||||
try:
|
||||
self.current_reader = pq.ParquetFile(file_path)
|
||||
self.current_batches = self.current_reader.iter_batches(
|
||||
batch_size=self.batch_size)
|
||||
self.current_file_idx += 1
|
||||
return True
|
||||
except Exception as e:
|
||||
print(f"Error opening file {file_path}: {str(e)}")
|
||||
return False
|
||||
|
||||
def __iter__(self):
|
||||
"""Iterate over the dataset in a streaming fashion."""
|
||||
print(f"Rank {self.rank}: Starting iteration")
|
||||
|
||||
# First try to open a file
|
||||
if not self._open_next_file():
|
||||
print(f"Rank {self.rank}: Failed to open first file")
|
||||
return
|
||||
|
||||
while True:
|
||||
def __len__(self):
|
||||
if self.local_indices is None:
|
||||
try:
|
||||
# Get next batch from current file
|
||||
batch = next(self.current_batches)
|
||||
batch_dict = batch.to_pydict()
|
||||
processed = self._process_batch(batch_dict)
|
||||
with open(self.plan_output_dir) as f:
|
||||
plan = json.load(f)
|
||||
self.local_indices = plan[str(self.rank)]
|
||||
except:
|
||||
raise Exception("The data plan hasn't been created yet")
|
||||
return len(self.local_indices)
|
||||
|
||||
# Update sample count
|
||||
batch_size = len(processed["latents"])
|
||||
self.total_samples += batch_size
|
||||
def __getitem__(self, idx):
|
||||
if self.local_indices is None:
|
||||
try:
|
||||
with open(self.plan_output_dir) as f:
|
||||
plan = json.load(f)
|
||||
self.local_indices = plan[self.rank]
|
||||
except:
|
||||
raise Exception("The data plan hasn't been created yet")
|
||||
file_path, row_idx = self.local_indices[idx]
|
||||
parquet_file = pq.ParquetFile(file_path)
|
||||
|
||||
# Print progress
|
||||
if self.total_samples % 1000 == 0:
|
||||
print(
|
||||
f"Rank {self.rank}: Processed {self.total_samples} samples"
|
||||
)
|
||||
# Calculate the row group to read into memory and the local idx
|
||||
# This way we can avoid reading in the entire parquet file
|
||||
cumulative = 0
|
||||
for i in range(parquet_file.num_row_groups):
|
||||
num_rows = parquet_file.metadata.row_group(i).num_rows
|
||||
if cumulative + num_rows > idx:
|
||||
row_group_index = i
|
||||
local_index = idx - cumulative
|
||||
break
|
||||
cumulative += num_rows
|
||||
|
||||
# Yield each item in the batch
|
||||
for lat, emb, mask, info in zip(processed["latents"],
|
||||
processed["embeddings"],
|
||||
processed["masks"],
|
||||
processed["info"]):
|
||||
if lat.numel() == 0: # Split is validation
|
||||
yield lat, emb, mask, info
|
||||
else:
|
||||
yield lat[:, -self.num_latent_t:], emb, mask, info
|
||||
row_group = parquet_file.read_row_group(row_group_index).to_pydict()
|
||||
row_dict = {k: v[local_index] for k, v in row_group.items()}
|
||||
del row_group
|
||||
|
||||
except StopIteration:
|
||||
# Current file is exhausted, try next file
|
||||
print(
|
||||
f"Rank {self.rank}: Current file exhausted, trying next file"
|
||||
)
|
||||
self.current_batches = None
|
||||
if not self._open_next_file():
|
||||
print(
|
||||
f"Rank {self.rank}: No more files to process. Total samples: {self.total_samples}"
|
||||
)
|
||||
break
|
||||
except Exception as e:
|
||||
print(f"Error processing batch: {str(e)}")
|
||||
self.current_batches = None
|
||||
if not self._open_next_file():
|
||||
print(
|
||||
f"Rank {self.rank}: Failed to open next file after error"
|
||||
)
|
||||
break
|
||||
processed = self._process_row(row_dict)
|
||||
lat, emb, mask, info = processed["latents"], processed[
|
||||
"embeddings"], processed["masks"], processed["info"]
|
||||
if lat.numel() == 0: # Validation parquet
|
||||
return lat, emb, mask, info
|
||||
else:
|
||||
lat = lat[:, -self.num_latent_t:]
|
||||
if self.sp_world_size > 1:
|
||||
lat = rearrange(lat,
|
||||
"t (n s) h w -> t n s h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
lat = lat[:, self.local_rank, :, :, :]
|
||||
return lat, emb, mask, info
|
||||
|
||||
# Clean up
|
||||
if self.current_reader is not None:
|
||||
self.current_reader.close()
|
||||
|
||||
def _process_batch(self, batch):
|
||||
def _process_row(self, row):
|
||||
"""Process a PyArrow batch into tensors."""
|
||||
out = {"lat": [], "emb": [], "msk": [], "info": []}
|
||||
out = {"lat": None, "emb": None, "msk": None, "info": None}
|
||||
|
||||
for i in range(len(batch["vae_latent_bytes"])):
|
||||
vae_latent_bytes = batch["vae_latent_bytes"][i]
|
||||
vae_latent_shape = batch["vae_latent_shape"][i]
|
||||
text_embedding_bytes = batch["text_embedding_bytes"][i]
|
||||
text_embedding_shape = batch["text_embedding_shape"][i]
|
||||
text_attention_mask_bytes = batch["text_attention_mask_bytes"][i]
|
||||
text_attention_mask_shape = batch["text_attention_mask_shape"][i]
|
||||
vae_latent_bytes = row["vae_latent_bytes"]
|
||||
vae_latent_shape = row["vae_latent_shape"]
|
||||
text_embedding_bytes = row["text_embedding_bytes"]
|
||||
text_embedding_shape = row["text_embedding_shape"]
|
||||
text_attention_mask_bytes = row["text_attention_mask_bytes"]
|
||||
text_attention_mask_shape = row["text_attention_mask_shape"]
|
||||
|
||||
# Process latent
|
||||
if not vae_latent_shape: # No VAE latent is stored. Split is validation
|
||||
lat = np.array([])
|
||||
else:
|
||||
lat = np.frombuffer(vae_latent_bytes,
|
||||
dtype=np.float32).reshape(vae_latent_shape)
|
||||
# Make array writable
|
||||
lat = np.copy(lat)
|
||||
# Process latent
|
||||
if not vae_latent_shape: # No VAE latent is stored. Split is validation
|
||||
lat = np.array([])
|
||||
else:
|
||||
lat = np.frombuffer(vae_latent_bytes,
|
||||
dtype=np.float32).reshape(vae_latent_shape)
|
||||
# Make array writable
|
||||
lat = np.copy(lat)
|
||||
|
||||
if random.random() < self.cfg_rate:
|
||||
emb = np.zeros((512, 4096), dtype=np.float32)
|
||||
else:
|
||||
emb = np.frombuffer(
|
||||
text_embedding_bytes,
|
||||
dtype=np.float32).reshape(text_embedding_shape)
|
||||
# Make array writable
|
||||
emb = np.copy(emb)
|
||||
if emb.shape[0] < 512:
|
||||
padded_emb = np.zeros((512, emb.shape[1]), dtype=np.float32)
|
||||
padded_emb[:emb.shape[0], :] = emb
|
||||
emb = padded_emb
|
||||
elif emb.shape[0] > 512:
|
||||
emb = emb[:512, :]
|
||||
if random.random() < self.cfg_rate:
|
||||
emb = np.zeros((512, 4096), dtype=np.float32)
|
||||
else:
|
||||
emb = np.frombuffer(text_embedding_bytes,
|
||||
dtype=np.float32).reshape(text_embedding_shape)
|
||||
# Make array writable
|
||||
emb = np.copy(emb)
|
||||
if emb.shape[0] < 512:
|
||||
padded_emb = np.zeros((512, emb.shape[1]), dtype=np.float32)
|
||||
padded_emb[:emb.shape[0], :] = emb
|
||||
emb = padded_emb
|
||||
elif emb.shape[0] > 512:
|
||||
emb = emb[:512, :]
|
||||
|
||||
# Process mask
|
||||
if len(text_attention_mask_bytes) > 0 and len(
|
||||
text_attention_mask_shape) > 0:
|
||||
msk = np.frombuffer(text_attention_mask_bytes,
|
||||
dtype=np.uint8).astype(np.bool_)
|
||||
msk = msk.reshape(1, -1)
|
||||
# Make array writable
|
||||
msk = np.copy(msk)
|
||||
if msk.shape[1] < 512:
|
||||
padded_msk = np.zeros((1, 512), dtype=np.bool_)
|
||||
padded_msk[:, :msk.shape[1]] = msk
|
||||
msk = padded_msk
|
||||
elif msk.shape[1] > 512:
|
||||
msk = msk[:, :512]
|
||||
else:
|
||||
msk = np.ones((1, 512), dtype=np.bool_)
|
||||
# to string
|
||||
file_name = str(batch["file_name"][i])
|
||||
# Collect metadata
|
||||
info = {
|
||||
"width": batch["width"][i],
|
||||
"height": batch["height"][i],
|
||||
"num_frames": batch["num_frames"][i],
|
||||
"duration_sec": batch["duration_sec"][i],
|
||||
"fps": batch["fps"][i],
|
||||
"file_name": batch["file_name"][i],
|
||||
"caption": batch["caption"][i],
|
||||
}
|
||||
# Process mask
|
||||
if len(text_attention_mask_bytes) > 0 and len(
|
||||
text_attention_mask_shape) > 0:
|
||||
msk = np.frombuffer(text_attention_mask_bytes,
|
||||
dtype=np.uint8).astype(np.bool_)
|
||||
msk = msk.reshape(1, -1)
|
||||
# Make array writable
|
||||
msk = np.copy(msk)
|
||||
if msk.shape[1] < 512:
|
||||
padded_msk = np.zeros((1, 512), dtype=np.bool_)
|
||||
padded_msk[:, :msk.shape[1]] = msk
|
||||
msk = padded_msk
|
||||
elif msk.shape[1] > 512:
|
||||
msk = msk[:, :512]
|
||||
else:
|
||||
msk = np.ones((1, 512), dtype=np.bool_)
|
||||
|
||||
out["lat"].append(torch.from_numpy(lat))
|
||||
out["emb"].append(torch.from_numpy(emb))
|
||||
out["msk"].append(torch.from_numpy(msk))
|
||||
out["info"].append(info)
|
||||
|
||||
return {
|
||||
"latents": torch.stack(out["lat"]) if out["lat"] else None,
|
||||
"embeddings": torch.stack(out["emb"]) if out["emb"] else None,
|
||||
"masks": torch.stack(out["msk"]) if out["msk"] else None,
|
||||
"info": out["info"]
|
||||
# Collect metadata
|
||||
info = {
|
||||
"width": row["width"],
|
||||
"height": row["height"],
|
||||
"num_frames": row["num_frames"],
|
||||
"duration_sec": row["duration_sec"],
|
||||
"fps": row["fps"],
|
||||
"file_name": row["file_name"],
|
||||
"caption": row["caption"],
|
||||
}
|
||||
|
||||
out["lat"] = torch.from_numpy(lat)
|
||||
out["emb"] = torch.from_numpy(emb)
|
||||
out["msk"] = torch.from_numpy(msk)
|
||||
out["info"] = info
|
||||
|
||||
def bind_cpu_cores(local_rank, cpu_per_process=16):
|
||||
"""根据local_rank绑定固定cpu核。"""
|
||||
start = local_rank * cpu_per_process
|
||||
end = start + cpu_per_process
|
||||
cores = list(range(start, end))
|
||||
print(f"[Rank {local_rank}] Binding to CPU cores: {cores}")
|
||||
os.sched_setaffinity(0, cores)
|
||||
return {
|
||||
"latents": out["lat"],
|
||||
"embeddings": out["emb"],
|
||||
"masks": out["msk"],
|
||||
"info": out["info"]
|
||||
}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
@@ -297,9 +253,6 @@ if __name__ == "__main__":
|
||||
f"Initialized process: rank={rank}, local_rank={local_rank}, world_size={world_size}, device={device}"
|
||||
)
|
||||
|
||||
# Bind CPU cores after distributed initialization
|
||||
# bind_cpu_cores(local_rank, cpu_per_process=16)
|
||||
|
||||
# Create dataset
|
||||
dataset = ParquetVideoTextDataset(
|
||||
args.path,
|
||||
|
||||
@@ -394,16 +394,7 @@ class TransformerLoader(ComponentLoader):
|
||||
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
|
||||
# Load the model using FSDP loader
|
||||
logger.info("Loading model from %s, default_dtype: %s", cls_name, default_dtype)
|
||||
# model = load_fsdp_model(model_cls=model_cls,
|
||||
# init_params={
|
||||
# "config": dit_config,
|
||||
# "hf_config": hf_config
|
||||
# },
|
||||
# weight_dir_list=safetensors_list,
|
||||
# device=fastvideo_args.device,
|
||||
# cpu_offload=fastvideo_args.use_cpu_offload,
|
||||
# default_dtype=default_dtype)
|
||||
logger.info("Loading model from %s", cls_name)
|
||||
model = load_fsdp_model(model_cls=model_cls,
|
||||
init_params={
|
||||
"config": dit_config,
|
||||
@@ -412,12 +403,7 @@ class TransformerLoader(ComponentLoader):
|
||||
weight_dir_list=safetensors_list,
|
||||
device=fastvideo_args.device,
|
||||
cpu_offload=fastvideo_args.use_cpu_offload,
|
||||
default_dtype=default_dtype,
|
||||
# TODO(will): make these configurable
|
||||
param_dtype=torch.bfloat16,
|
||||
reduce_dtype=torch.float32,
|
||||
output_dtype=None,
|
||||
)
|
||||
default_dtype=default_dtype)
|
||||
if fastvideo_args.enable_torch_compile:
|
||||
logger.info("Torch Compile enabled for DiT")
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
|
||||
@@ -14,16 +14,13 @@ from typing import (Any, Callable, DefaultDict, Dict, Generator, Hashable, List,
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.distributed import DeviceMesh, init_device_mesh
|
||||
from torch.distributed.fsdp import CPUOffloadPolicy, fully_shard, MixedPrecisionPolicy
|
||||
from torch.distributed._composable.fsdp import CPUOffloadPolicy, fully_shard
|
||||
from torch.distributed._tensor import distribute_tensor
|
||||
from torch.nn.modules.module import _IncompatibleKeys
|
||||
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.models.loader.weight_utils import safetensors_weights_iterator
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
# TODO(PY): move this to utils elsewhere
|
||||
@@ -89,29 +86,16 @@ def get_param_names_mapping(
|
||||
|
||||
|
||||
# TODO(PY): add compile option
|
||||
# param_dtype: torch.dtype,
|
||||
# reduce_dtype: torch.dtype,
|
||||
# output_dtype: torch.dtype,
|
||||
# pp_enabled: bool = False,
|
||||
# cpu_offload: bool = False,
|
||||
def load_fsdp_model(
|
||||
model_cls: Type[nn.Module],
|
||||
init_params: Dict[str, Any],
|
||||
weight_dir_list: List[str],
|
||||
device: torch.device,
|
||||
default_dtype: torch.dtype,
|
||||
param_dtype: torch.dtype,
|
||||
reduce_dtype: torch.dtype,
|
||||
cpu_offload: bool = False,
|
||||
output_dtype: Optional[torch.dtype] = None,
|
||||
default_dtype: Optional[torch.dtype] = torch.bfloat16,
|
||||
) -> torch.nn.Module:
|
||||
|
||||
mp_policy = MixedPrecisionPolicy(param_dtype, reduce_dtype, output_dtype, cast_forward_inputs=True)
|
||||
|
||||
# with set_default_dtype(default_dtype), torch.device("meta"):
|
||||
with set_default_dtype(default_dtype), torch.device("meta"):
|
||||
model = model_cls(**init_params)
|
||||
|
||||
device_mesh = init_device_mesh(
|
||||
"cuda",
|
||||
mesh_shape=(get_sequence_model_parallel_world_size(), ),
|
||||
@@ -120,7 +104,6 @@ def load_fsdp_model(
|
||||
shard_model(model,
|
||||
cpu_offload=cpu_offload,
|
||||
reshard_after_forward=True,
|
||||
mp_policy=mp_policy,
|
||||
dp_mesh=device_mesh["dp"])
|
||||
weight_iterator = safetensors_weights_iterator(weight_dir_list)
|
||||
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
|
||||
@@ -147,7 +130,6 @@ def shard_model(
|
||||
*,
|
||||
cpu_offload: bool,
|
||||
reshard_after_forward: bool = True,
|
||||
mp_policy: Optional[MixedPrecisionPolicy] = None,
|
||||
dp_mesh: Optional[DeviceMesh] = None,
|
||||
) -> None:
|
||||
"""
|
||||
@@ -175,17 +157,14 @@ def shard_model(
|
||||
"""
|
||||
fsdp_kwargs = {
|
||||
"reshard_after_forward": reshard_after_forward,
|
||||
"mesh": dp_mesh,
|
||||
"mp_policy": mp_policy,
|
||||
"mesh": dp_mesh
|
||||
}
|
||||
if cpu_offload:
|
||||
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy()
|
||||
|
||||
# iterating in reverse to start with
|
||||
# Shard the model with FSDP, iterating in reverse to start with
|
||||
# lowest-level modules first
|
||||
num_layers_sharded = 0
|
||||
# TODO(will): don't reshard after forward for the last layer to save on the
|
||||
# all-gather that will immediately happen Shard the model with FSDP,
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
if any([
|
||||
shard_condition(n, m)
|
||||
@@ -232,10 +211,6 @@ def load_fsdp_model_from_full_model_state_dict(
|
||||
NotImplementedError: If got FSDP with more than 1D.
|
||||
"""
|
||||
meta_sharded_sd = model.state_dict()
|
||||
# s = fully_shard.state(model)
|
||||
# logger.info(f"type(s): {type(s)}")
|
||||
# logger.info(f"s: {s}")
|
||||
# import pdb; pdb.set_trace()
|
||||
|
||||
sharded_sd = {}
|
||||
to_merge_params: DefaultDict[Hashable, Dict[Any, Any]] = defaultdict(dict)
|
||||
|
||||
@@ -155,21 +155,17 @@ class ComposedPipelineBase(ABC):
|
||||
for key, value in config_args.items():
|
||||
setattr(fastvideo_args, key, value)
|
||||
|
||||
# we use cpu offload for training
|
||||
fastvideo_args.use_cpu_offload = False
|
||||
# make sure we are in training mode
|
||||
fastvideo_args.inference_mode = False
|
||||
# we hijack the precision to be the master weight type so that the
|
||||
# model is loaded with the correct precision. Subsequently we will
|
||||
# use FSDP2's MixedPrecisionPolicy to set the precision for the
|
||||
# fwd, bwd, and other operations' precision.
|
||||
fastvideo_args.precision = fastvideo_args.master_weight_type
|
||||
assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
|
||||
|
||||
fastvideo_args.check_fastvideo_args()
|
||||
|
||||
logger.info(f"fastvideo_args in from_pretrained: {fastvideo_args}")
|
||||
|
||||
# fastvideo_args = FastVideoArgs(
|
||||
# model_path=model_path,
|
||||
# device_str=device or "cuda" if torch.cuda.is_available() else "cpu",
|
||||
# **config_args)
|
||||
fastvideo_args.check_fastvideo_args()
|
||||
|
||||
return cls(model_path,
|
||||
fastvideo_args,
|
||||
required_config_modules=required_config_modules)
|
||||
|
||||
@@ -140,7 +140,6 @@ class PreprocessPipeline(ComposedPipelineBase):
|
||||
0], result_batch.prompt_attention_mask[0]
|
||||
assert prompt_embeds.shape[0] == prompt_attention_mask.shape[0]
|
||||
|
||||
# Remove padding from prompt_embeds using attention mask for all batches
|
||||
# Get sequence lengths from attention masks (number of 1s)
|
||||
seq_lens = prompt_attention_mask.sum(dim=1)
|
||||
# Create a list to store non-padded embeddings and masks
|
||||
@@ -354,9 +353,6 @@ class PreprocessPipeline(ComposedPipelineBase):
|
||||
"validation_parquet_dataset")
|
||||
os.makedirs(validation_parquet_dir, exist_ok=True)
|
||||
|
||||
# Initialize Parquet dataset
|
||||
validation_parquet_path = os.path.join(validation_parquet_dir,
|
||||
"data.parquet")
|
||||
|
||||
with open(args.validation_prompt_txt, encoding="utf-8") as file:
|
||||
lines = file.readlines()
|
||||
|
||||
@@ -74,8 +74,7 @@ class DenoisingStage(PipelineStage):
|
||||
)
|
||||
|
||||
# Setup precision and autocast settings
|
||||
# target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
target_dtype = torch.bfloat16
|
||||
target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
@@ -190,7 +189,7 @@ class DenoisingStage(PipelineStage):
|
||||
|
||||
# Predict noise residual
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=torch.bfloat16,
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
|
||||
# TODO(will-refactor): all of this should be in the stage's init
|
||||
|
||||
@@ -191,15 +191,14 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
logger.info(f"infos: {infos}")
|
||||
caption = infos['caption']
|
||||
captions.append(caption)
|
||||
prompt_embeds = embeddings.to(fastvideo_args.device).to(torch.bfloat16)
|
||||
prompt_attention_mask = masks.to(fastvideo_args.device).to(torch.bfloat16)
|
||||
prompt_embeds = embeddings.to(fastvideo_args.device)
|
||||
prompt_attention_mask = masks.to(fastvideo_args.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]
|
||||
logger.info('embed dtype', prompt_embeds.dtype)
|
||||
|
||||
# Prepare batch for validation
|
||||
# print('shape of embeddings', prompt_embeds.shape)
|
||||
@@ -216,7 +215,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
width=args.num_width,
|
||||
num_frames=args.num_frames,
|
||||
# num_inference_steps=fastvideo_args.validation_sampling_steps,
|
||||
num_inference_steps=50,
|
||||
num_inference_steps=10,
|
||||
# guidance_scale=fastvideo_args.validation_guidance_scale,
|
||||
guidance_scale=1,
|
||||
n_tokens=n_tokens,
|
||||
@@ -226,11 +225,10 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
)
|
||||
|
||||
# Run validation inference
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
with torch.inference_mode():
|
||||
output_batch = self.validation_pipeline.forward(
|
||||
batch, fastvideo_args)
|
||||
samples = output_batch.output
|
||||
with torch.inference_mode():
|
||||
output_batch = self.validation_pipeline.forward(
|
||||
batch, fastvideo_args)
|
||||
samples = output_batch.output
|
||||
|
||||
# Process outputs
|
||||
video = rearrange(samples, "b c t h w -> t b c h w")
|
||||
@@ -531,8 +529,6 @@ class WanTrainingPipeline(TrainingPipeline):
|
||||
|
||||
args_copy.inference_mode = True
|
||||
args_copy.vae_config.load_encoder = False
|
||||
# TODO(will): clean this up
|
||||
args_copy.precision = "bf16"
|
||||
validation_pipeline = WanValidationPipeline.from_pretrained(
|
||||
args.model_path, args=args_copy)
|
||||
|
||||
|
||||
@@ -15,6 +15,7 @@ from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage)
|
||||
|
||||
# TODO(will): move PRECISION_TO_TYPE to better place
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
+2
-2
@@ -49,5 +49,5 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--weight_decay 0.01 \
|
||||
--not_apply_cfg_solver \
|
||||
--master_weight_type "fp32" \
|
||||
--max_grad_norm 1.0
|
||||
--master_weight_type "bf16" \
|
||||
--max_grad_norm 1.0
|
||||
Reference in New Issue
Block a user