Compare commits

...
Author SHA1 Message Date
JerryZhou54 1757d3dba0 Pass pre-commit tests 2025-05-30 17:59:34 +00:00
JerryZhou54 a9d0c29ed9 fix distributed datasets issue 2025-05-30 17:50:10 +00:00
2 changed files with 5 additions and 1 deletions
+4 -1
View File
@@ -43,7 +43,8 @@ class ParquetVideoTextDataset(Dataset):
self.cfg_rate = cfg_rate
self.num_latent_t = num_latent_t
self.local_indices = None
self.plan_output_dir = os.path.join(self.path, "data_plan.json")
self.plan_output_dir = os.path.join(
self.path, f"data_plan_{self.world_size}_{self.sp_world_size}.json")
ranks = get_sp_group().ranks
group_ranks: List[List] = [[] for _ in range(self.world_size)]
@@ -54,6 +55,7 @@ class ParquetVideoTextDataset(Dataset):
# This will be useful when resume training
if os.path.exists(self.plan_output_dir):
print(f"Using existing plan from {self.plan_output_dir}")
dist.barrier()
return
# Find all parquet files recursively, and record num_rows for each file
@@ -87,6 +89,7 @@ class ParquetVideoTextDataset(Dataset):
with open(self.plan_output_dir, "w") as f:
json.dump(plan, f)
dist.barrier()
def __len__(self):
if self.local_indices is None:
@@ -97,6 +97,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
prefetch_factor=2,
shuffle=False,
pin_memory=True,
pin_memory_device=f"cuda:{torch.cuda.current_device()}",
drop_last=True)
self.noise_scheduler = noise_scheduler