mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
updates
This commit is contained in:
+30
-24
@@ -22,8 +22,8 @@ decord.bridge.set_bridge("torch")
|
||||
logger = get_logger(__name__)
|
||||
|
||||
HEIGHT_BUCKETS = [256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536]
|
||||
WIDTH_BUCKETS = [256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536]
|
||||
FRAME_BUCKETS = [16, 24, 32, 48, 64, 80]
|
||||
WIDTH_BUCKETS = [256, 320, 384, 480, 512, 576, 720, 768, 848, 960, 1024, 1280, 1536]
|
||||
FRAME_BUCKETS = [16, 24, 32, 48, 64, 80, 84]
|
||||
|
||||
|
||||
class VideoDataset(Dataset):
|
||||
@@ -144,17 +144,17 @@ class VideoDataset(Dataset):
|
||||
}
|
||||
else:
|
||||
image, video, _ = self._preprocess_video(self.video_paths[index])
|
||||
|
||||
return {
|
||||
"prompt": self.id_token + self.prompts[index],
|
||||
"image": image,
|
||||
"video": video,
|
||||
"video_metadata": {
|
||||
"num_frames": video.shape[0],
|
||||
"height": video.shape[2],
|
||||
"width": video.shape[3],
|
||||
},
|
||||
}
|
||||
if video is not None:
|
||||
return {
|
||||
"prompt": self.id_token + self.prompts[index],
|
||||
"image": image,
|
||||
"video": video,
|
||||
"video_metadata": {
|
||||
"num_frames": video.shape[0],
|
||||
"height": video.shape[2],
|
||||
"width": video.shape[3],
|
||||
},
|
||||
}
|
||||
|
||||
def _load_dataset_from_local_path(self) -> Tuple[List[str], List[str]]:
|
||||
if not self.data_root.exists():
|
||||
@@ -276,12 +276,15 @@ class VideoDatasetWithResizing(VideoDataset):
|
||||
else:
|
||||
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))
|
||||
)
|
||||
|
||||
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)
|
||||
frames = frames[:nearest_frame_bucket].float()
|
||||
frames = frames.permute(0, 3, 1, 2).contiguous()
|
||||
@@ -344,6 +347,8 @@ class VideoDatasetWithResizeAndRectangleCrop(VideoDataset):
|
||||
nearest_frame_bucket = min(
|
||||
self.frame_buckets, key=lambda x: abs(x - min(video_num_frames, self.max_num_frames))
|
||||
)
|
||||
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))
|
||||
|
||||
@@ -404,16 +409,17 @@ class BucketSampler(Sampler):
|
||||
|
||||
def __iter__(self):
|
||||
for index, data in enumerate(self.data_source):
|
||||
video_metadata = data["video_metadata"]
|
||||
f, h, w = video_metadata["num_frames"], video_metadata["height"], video_metadata["width"]
|
||||
if data is not None:
|
||||
video_metadata = data["video_metadata"]
|
||||
f, h, w = video_metadata["num_frames"], video_metadata["height"], video_metadata["width"]
|
||||
|
||||
self.buckets[(f, h, w)].append(data)
|
||||
if len(self.buckets[(f, h, w)]) == self.batch_size:
|
||||
if self.shuffle:
|
||||
random.shuffle(self.buckets[(f, h, w)])
|
||||
yield self.buckets[(f, h, w)]
|
||||
del self.buckets[(f, h, w)]
|
||||
self.buckets[(f, h, w)] = []
|
||||
self.buckets[(f, h, w)].append(data)
|
||||
if len(self.buckets[(f, h, w)]) == self.batch_size:
|
||||
if self.shuffle:
|
||||
random.shuffle(self.buckets[(f, h, w)])
|
||||
yield self.buckets[(f, h, w)]
|
||||
del self.buckets[(f, h, w)]
|
||||
self.buckets[(f, h, w)] = []
|
||||
|
||||
if self.drop_last:
|
||||
return
|
||||
|
||||
@@ -13,7 +13,7 @@ from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from diffusers import AutoencoderKL
|
||||
from diffusers import AutoencoderKLMochi
|
||||
from diffusers.training_utils import set_seed
|
||||
from diffusers.utils import export_to_video, get_logger
|
||||
from torch.utils.data import DataLoader
|
||||
@@ -24,7 +24,9 @@ from transformers import T5EncoderModel, T5Tokenizer
|
||||
|
||||
import decord # isort:skip
|
||||
|
||||
from ..dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip
|
||||
import sys
|
||||
sys.path.append("..")
|
||||
from dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip
|
||||
|
||||
|
||||
decord.bridge.set_bridge("torch")
|
||||
@@ -485,12 +487,12 @@ def main():
|
||||
tokenizer = T5Tokenizer.from_pretrained(args.model_id, subfolder="tokenizer")
|
||||
text_encoder = T5EncoderModel.from_pretrained(
|
||||
args.model_id, subfolder="text_encoder", torch_dtype=weight_dtype
|
||||
)
|
||||
text_encoder = text_encoder.to(device)
|
||||
|
||||
vae = AutoencoderKL.from_pretrained(args.model_id, subfolder="vae", torch_dtype=weight_dtype)
|
||||
vae = vae.to(device)
|
||||
|
||||
).to(device)
|
||||
|
||||
vae = AutoencoderKLMochi.from_pretrained(
|
||||
args.model_id, subfolder="vae", torch_dtype=weight_dtype
|
||||
).to(device)
|
||||
|
||||
if args.use_slicing:
|
||||
vae.enable_slicing()
|
||||
if args.use_tiling:
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
#!/bin/bash
|
||||
|
||||
MODEL_ID="genmo/mochi-1-preview"
|
||||
|
||||
NUM_GPUS=1
|
||||
|
||||
# For more details on the expected data format, please refer to the README.
|
||||
DATA_ROOT="/home/sayak/cogvideox-factory/video-dataset-disney" # This needs to be the path to the base directory where your videos are located.
|
||||
CAPTION_COLUMN="prompt.txt"
|
||||
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"
|
||||
MAX_NUM_FRAMES="84"
|
||||
MAX_SEQUENCE_LENGTH=256
|
||||
TARGET_FPS=30
|
||||
BATCH_SIZE=1
|
||||
DTYPE=fp32
|
||||
|
||||
# To create a folder-style dataset structure without pre-encoding videos and captions
|
||||
# For Image-to-Video finetuning, make sure to pass `--save_image_latents`
|
||||
CMD_WITHOUT_PRE_ENCODING="\
|
||||
torchrun --nproc_per_node=$NUM_GPUS \
|
||||
prepare_dataset.py \
|
||||
--model_id $MODEL_ID \
|
||||
--data_root $DATA_ROOT \
|
||||
--caption_column $CAPTION_COLUMN \
|
||||
--video_column $VIDEO_COLUMN \
|
||||
--output_dir $OUTPUT_DIR \
|
||||
--height_buckets $HEIGHT_BUCKETS \
|
||||
--width_buckets $WIDTH_BUCKETS \
|
||||
--frame_buckets $FRAME_BUCKETS \
|
||||
--max_num_frames $MAX_NUM_FRAMES \
|
||||
--max_sequence_length $MAX_SEQUENCE_LENGTH \
|
||||
--target_fps $TARGET_FPS \
|
||||
--batch_size $BATCH_SIZE \
|
||||
--dtype $DTYPE
|
||||
"
|
||||
|
||||
CMD_WITH_PRE_ENCODING="$CMD_WITHOUT_PRE_ENCODING --save_latents_and_embeddings"
|
||||
|
||||
# Select which you'd like to run
|
||||
CMD=$CMD_WITH_PRE_ENCODING
|
||||
|
||||
echo "===== Running \`$CMD\` ====="
|
||||
eval $CMD
|
||||
echo -ne "===== Finished running script =====\n"
|
||||
Reference in New Issue
Block a user