Windows support for T2V scripts (#48)

This commit is contained in:
Aryan
2024-10-21 02:40:05 +05:30
committed by GitHub
parent 6c00cf094b
commit db9b295408
3 changed files with 46 additions and 30 deletions
+4 -4
View File
@@ -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,
)
+21 -13
View File
@@ -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,
+21 -13
View File
@@ -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,