From db9b2954084faca6c0eb62213b89c8f0ca9e83b1 Mon Sep 17 00:00:00 2001 From: Aryan Date: Mon, 21 Oct 2024 02:40:05 +0530 Subject: [PATCH] Windows support for T2V scripts (#48) --- training/cogvideox_image_to_video_lora.py | 8 +++--- training/cogvideox_text_to_video_lora.py | 34 ++++++++++++++--------- training/cogvideox_text_to_video_sft.py | 34 ++++++++++++++--------- 3 files changed, 46 insertions(+), 30 deletions(-) diff --git a/training/cogvideox_image_to_video_lora.py b/training/cogvideox_image_to_video_lora.py index 305aac4..ff40a54 100644 --- a/training/cogvideox_image_to_video_lora.py +++ b/training/cogvideox_image_to_video_lora.py @@ -202,11 +202,11 @@ def log_validation( class CollateFunction: - def __init__(self, weight_dtype, load_tensors): + def __init__(self, weight_dtype: torch.dtype, load_tensors: bool) -> None: self.weight_dtype = weight_dtype self.load_tensors = load_tensors - def __call__(self, data): + def __call__(self, data: Dict[str, Any]) -> Dict[str, torch.Tensor]: prompts = [x["prompt"] for x in data[0]] if self.load_tensors: @@ -519,13 +519,13 @@ def main(args): video_reshape_mode=args.video_reshape_mode, **dataset_init_kwargs ) - collate_fn_instance = CollateFunction(weight_dtype, args.load_tensors) + collate_fn = CollateFunction(weight_dtype, args.load_tensors) train_dataloader = DataLoader( train_dataset, batch_size=1, sampler=BucketSampler(train_dataset, batch_size=args.train_batch_size, shuffle=True), - collate_fn=collate_fn_instance, + collate_fn=collate_fn, num_workers=args.dataloader_num_workers, pin_memory=args.pin_memory, ) diff --git a/training/cogvideox_text_to_video_lora.py b/training/cogvideox_text_to_video_lora.py index cb94d76..de505f2 100644 --- a/training/cogvideox_text_to_video_lora.py +++ b/training/cogvideox_text_to_video_lora.py @@ -198,6 +198,26 @@ def log_validation( return videos +class CollateFunction: + def __init__(self, weight_dtype: torch.dtype, load_tensors: bool) -> None: + self.weight_dtype = weight_dtype + self.load_tensors = load_tensors + + def __call__(self, data: Dict[str, Any]) -> Dict[str, torch.Tensor]: + prompts = [x["prompt"] for x in data[0]] + + if self.load_tensors: + prompts = torch.stack(prompts).to(dtype=self.weight_dtype, non_blocking=True) + + videos = [x["video"] for x in data[0]] + videos = torch.stack(videos).to(dtype=self.weight_dtype, non_blocking=True) + + return { + "videos": videos, + "prompts": prompts, + } + + def main(args): if args.report_to == "wandb" and args.hub_token is not None: raise ValueError( @@ -491,19 +511,7 @@ def main(args): video_reshape_mode=args.video_reshape_mode, **dataset_init_kwargs ) - def collate_fn(data): - prompts = [x["prompt"] for x in data[0]] - - if args.load_tensors: - prompts = torch.stack(prompts).to(dtype=weight_dtype, non_blocking=True) - - videos = [x["video"] for x in data[0]] - videos = torch.stack(videos).to(dtype=weight_dtype, non_blocking=True) - - return { - "videos": videos, - "prompts": prompts, - } + collate_fn = CollateFunction(weight_dtype, args.load_tensors) train_dataloader = DataLoader( train_dataset, diff --git a/training/cogvideox_text_to_video_sft.py b/training/cogvideox_text_to_video_sft.py index 72592d6..e65a62a 100644 --- a/training/cogvideox_text_to_video_sft.py +++ b/training/cogvideox_text_to_video_sft.py @@ -188,6 +188,26 @@ def log_validation( return videos +class CollateFunction: + def __init__(self, weight_dtype: torch.dtype, load_tensors: bool) -> None: + self.weight_dtype = weight_dtype + self.load_tensors = load_tensors + + def __call__(self, data: Dict[str, Any]) -> Dict[str, torch.Tensor]: + prompts = [x["prompt"] for x in data[0]] + + if self.load_tensors: + prompts = torch.stack(prompts).to(dtype=self.weight_dtype, non_blocking=True) + + videos = [x["video"] for x in data[0]] + videos = torch.stack(videos).to(dtype=self.weight_dtype, non_blocking=True) + + return { + "videos": videos, + "prompts": prompts, + } + + def main(args): if args.report_to == "wandb" and args.hub_token is not None: raise ValueError( @@ -457,19 +477,7 @@ def main(args): video_reshape_mode=args.video_reshape_mode, **dataset_init_kwargs ) - def collate_fn(data): - prompts = [x["prompt"] for x in data[0]] - - if args.load_tensors: - prompts = torch.stack(prompts).to(dtype=weight_dtype, non_blocking=True) - - videos = [x["video"] for x in data[0]] - videos = torch.stack(videos).to(dtype=weight_dtype, non_blocking=True) - - return { - "videos": videos, - "prompts": prompts, - } + collate_fn = CollateFunction(weight_dtype, args.load_tensors) train_dataloader = DataLoader( train_dataset,