Compare commits

...
Author SHA1 Message Date
SolitaryThinker 0c6e1be568 fix 2025-07-05 22:11:10 +00:00
SolitaryThinker ff4e183ada run lint 3.10 2025-07-05 13:36:20 -07:00
“BrianChen1129” f5245f1ea5 dataset bug 2025-07-05 13:32:17 -07:00
“BrianChen1129” b0a104fef6 update 2025-07-05 13:32:16 -07:00
“BrianChen1129” a01b6d386d update 2025-07-05 13:32:16 -07:00
3 changed files with 2 additions and 3 deletions
-1
View File
@@ -140,7 +140,6 @@ def collate_rows_from_parquet_schema(rows,
if shape_key in row and bytes_key in row:
shape = row[shape_key]
bytes_data = row[bytes_key]
if len(bytes_data) == 0:
tensor = torch.zeros(0, dtype=torch.bfloat16)
else:
+1 -1
View File
@@ -925,7 +925,6 @@ def maybe_init_distributed_environment_and_model_parallel(
world_size = int(os.environ.get("WORLD_SIZE", 1))
rank = int(os.environ.get("RANK", 0))
device = torch.device(f"cuda:{local_rank}")
torch.cuda.set_device(device)
init_distributed_environment(
world_size=world_size,
@@ -935,6 +934,7 @@ def maybe_init_distributed_environment_and_model_parallel(
device_id=device)
initialize_model_parallel(tensor_model_parallel_size=tp_size,
sequence_model_parallel_size=sp_size)
torch.cuda.set_device(device)
def model_parallel_is_initialized() -> bool:
@@ -229,4 +229,4 @@ if __name__ == "__main__":
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.use_cpu_offload = False
main(args)
main(args)