Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0c6e1be568 | ||
|
|
ff4e183ada | ||
|
|
f5245f1ea5 | ||
|
|
b0a104fef6 | ||
|
|
a01b6d386d |
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user