|
4 | 4 |
|
5 | 5 | import torch |
6 | 6 | import torch.distributed as dist |
| 7 | +from vime.utils.common import is_npu |
| 8 | +if is_npu(): |
| 9 | + import mindspeed.megatron_adaptor |
7 | 10 | from megatron.core.enums import ModelType |
8 | 11 | from megatron.training.arguments import parse_args, validate_args |
9 | 12 | from megatron.training.checkpointing import get_checkpoint_name, get_checkpoint_tracker_filename, save_checkpoint |
|
14 | 17 | from vime.backends.megatron_utils.arguments import set_default_megatron_args |
15 | 18 | from vime.backends.megatron_utils.initialize import init |
16 | 19 | from vime.backends.megatron_utils.model_provider import get_model_provider_func |
17 | | -from vime.utils.common import is_npu |
18 | 20 | from vime.utils.logging_utils import configure_logger |
19 | 21 | from vime.utils.memory_utils import print_memory |
20 | 22 |
|
@@ -92,15 +94,19 @@ def main(): |
92 | 94 | os.environ.setdefault("LOCAL_RANK", str(local_rank)) |
93 | 95 | os.environ.setdefault("MASTER_ADDR", "localhost") |
94 | 96 | os.environ.setdefault("MASTER_PORT", "12355") |
95 | | - backend = "nccl" |
96 | 97 | if is_npu(): |
97 | | - backend = "hccl" |
98 | | - dist.init_process_group( |
99 | | - backend=backend, |
100 | | - world_size=world_size, |
101 | | - rank=global_rank, |
102 | | - device_id=torch.device(f"cuda:{local_rank}"), |
103 | | - ) |
| 98 | + dist.init_process_group( |
| 99 | + backend="hccl", |
| 100 | + world_size=world_size, |
| 101 | + rank=global_rank, |
| 102 | + ) |
| 103 | + else: |
| 104 | + dist.init_process_group( |
| 105 | + backend="nccl", |
| 106 | + world_size=world_size, |
| 107 | + rank=global_rank, |
| 108 | + device_id=torch.device(f"cuda:{local_rank}"), |
| 109 | + ) |
104 | 110 | args = get_args() |
105 | 111 | init(args) |
106 | 112 |
|
|
0 commit comments