This commit is contained in:
sayakpaul
2024-11-19 10:21:28 +05:30
parent ac83c78a5a
commit 1409d478d5
5 changed files with 22 additions and 16 deletions
+7
View File
@@ -151,6 +151,13 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None:
default=64,
help="The lora_alpha to compute scaling factor (lora_alpha / rank) for LoRA matrices.",
)
parser.add_argument(
"--target_modules",
nargs="+",
type=str,
default=["to_k", "to_q", "to_v", "to_out.0"],
help="Target modules to train LoRA for."
)
parser.add_argument(
"--mixed_precision",
type=str,
+6 -9
View File
@@ -128,9 +128,7 @@ class VideoDataset(Dataset):
# temporal compression factor is 6. Initially, the VAE encodings will have
# 24 latent number of frames. So, if we were to train with a
# max frame size of 84 and frame bucket of [84], we need to have the following logic.
# print(f"{video_latents.shape=}")
latent_num_frames = video_latents.size(0)
# print(f"{latent_num_frames=}")
num_frames = (latent_num_frames // 2) * (VAE_TEMPORAL_SCALE_FACTOR + 1)
height = video_latents.size(2) * VAE_SPATIAL_SCALE_FACTOR
@@ -288,11 +286,10 @@ class VideoDatasetWithResizing(VideoDataset):
video_num_frames = len(video_reader)
nearest_frame_bucket = min(
self.frame_buckets, key=lambda x: abs(x - min(video_num_frames, self.max_num_frames))
[bucket for bucket in self.frame_buckets if bucket <= video_num_frames],
key=lambda x: abs(x - min(video_num_frames, self.max_num_frames)),
default=1,
)
if video_num_frames < nearest_frame_bucket:
# TODO: we could handle this by padding zero frames or duplicating the existing frames?
return None, None, None
frame_indices = list(range(0, video_num_frames, video_num_frames // nearest_frame_bucket))
frames = video_reader.get_batch(frame_indices)
@@ -355,10 +352,10 @@ class VideoDatasetWithResizeAndRectangleCrop(VideoDataset):
video_reader = decord.VideoReader(uri=path.as_posix())
video_num_frames = len(video_reader)
nearest_frame_bucket = min(
self.frame_buckets, key=lambda x: abs(x - min(video_num_frames, self.max_num_frames))
[bucket for bucket in self.frame_buckets if bucket <= video_num_frames],
key=lambda x: abs(x - min(video_num_frames, self.max_num_frames)),
default=1,
)
if video_num_frames < nearest_frame_bucket:
return None, None, None
frame_indices = list(range(0, video_num_frames, video_num_frames // nearest_frame_bucket))
+5 -4
View File
@@ -321,8 +321,11 @@ def serialize_artifacts(
prompt_embeds: Optional[torch.Tensor] = None,
prompt_attention_mask: Optional[torch.Tensor] = None
) -> None:
num_frames, height, width = videos.size(1), videos.size(3), videos.size(4)
metadata = [{"num_frames": num_frames, "height": height, "width": width}]
metadata = []
for i in range(videos.size(0)):
video = videos[i:i+1]
metadata_dict = {"num_frames": video.size(1), "height": video.size(3), "width": video.size(4)}
metadata.append(metadata_dict)
data_folder_mapper_list = [
(images, images_dir, lambda img, path: save_image(img[0], path), "png"),
@@ -542,7 +545,6 @@ def main():
video_latents = vae._encode(videos)
video_latents = video_latents.to(memory_format=torch.contiguous_format, dtype=weight_dtype)
print(f"{video_latents.shape=}")
# Encode prompts
prompt_embeds, prompt_attention_mask = compute_prompt_embeddings(
@@ -554,7 +556,6 @@ def main():
weight_dtype,
requires_grad=False,
)
print(f"{prompt_attention_mask.shape=}")
if images is not None:
images = (images.permute(0, 2, 1, 3, 4) + 1) / 2
+3 -2
View File
@@ -11,11 +11,11 @@ VIDEO_COLUMN="videos.txt"
OUTPUT_DIR="/home/sayak/cogvideox-factory/video-dataset-disney/mochi-1/preprocessed-dataset"
HEIGHT_BUCKETS="480"
WIDTH_BUCKETS="848"
FRAME_BUCKETS="84"
FRAME_BUCKETS="1 84"
MAX_NUM_FRAMES="84"
MAX_SEQUENCE_LENGTH=256
TARGET_FPS=30
BATCH_SIZE=1
BATCH_SIZE=4
DTYPE=fp32
# To create a folder-style dataset structure without pre-encoding videos and captions
@@ -35,6 +35,7 @@ CMD_WITHOUT_PRE_ENCODING="\
--max_sequence_length $MAX_SEQUENCE_LENGTH \
--target_fps $TARGET_FPS \
--batch_size $BATCH_SIZE \
--use_slicing \
--dtype $DTYPE
"
+1 -1
View File
@@ -354,7 +354,7 @@ def main(args):
r=args.rank,
lora_alpha=args.lora_alpha,
init_lora_weights=True,
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
target_modules=args.target_modules,
)
transformer.add_adapter(transformer_lora_config)