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