mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
updates.
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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))
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
"
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user