[Fix] fix local-rank in pytorch2.0 (#728)
parent
448b4fae0d
commit
6c5e83608b
|
@ -38,7 +38,10 @@ def parse_args():
|
|||
choices=['none', 'pytorch', 'slurm', 'mpi'],
|
||||
default='none',
|
||||
help='job launcher')
|
||||
parser.add_argument('--local_rank', type=int, default=0)
|
||||
# When using PyTorch version >= 2.0.0, the `torch.distributed.launch`
|
||||
# will pass the `--local-rank` parameter to `tools/train.py` instead
|
||||
# of `--local_rank`.
|
||||
parser.add_argument('--local_rank', '--local-rank', type=int, default=0)
|
||||
args = parser.parse_args()
|
||||
if 'LOCAL_RANK' not in os.environ:
|
||||
os.environ['LOCAL_RANK'] = str(args.local_rank)
|
||||
|
|
Loading…
Reference in New Issue