Skip to content

Commit b93ff56

Browse files
committed
fix ref-load torch_dist error
Signed-off-by: flb_ <floatlibai@gmail.com>
1 parent e5aa5e7 commit b93ff56

2 files changed

Lines changed: 17 additions & 15 deletions

File tree

tools/convert_hf_to_torch_dist.py

Lines changed: 16 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,10 @@
44

55
import torch
66
import torch.distributed as dist
7+
from vime.utils.common import is_npu
8+
9+
if is_npu():
10+
import mindspeed.megatron_adaptor # noqa: F401
711
from megatron.core.enums import ModelType
812
from megatron.training.arguments import parse_args, validate_args
913
from megatron.training.checkpointing import get_checkpoint_name, get_checkpoint_tracker_filename, save_checkpoint
@@ -14,7 +18,6 @@
1418
from vime.backends.megatron_utils.arguments import set_default_megatron_args
1519
from vime.backends.megatron_utils.initialize import init
1620
from vime.backends.megatron_utils.model_provider import get_model_provider_func
17-
from vime.utils.common import is_npu
1821
from vime.utils.logging_utils import configure_logger
1922
from vime.utils.memory_utils import print_memory
2023

@@ -92,15 +95,19 @@ def main():
9295
os.environ.setdefault("LOCAL_RANK", str(local_rank))
9396
os.environ.setdefault("MASTER_ADDR", "localhost")
9497
os.environ.setdefault("MASTER_PORT", "12355")
95-
backend = "nccl"
9698
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-
)
99+
dist.init_process_group(
100+
backend="hccl",
101+
world_size=world_size,
102+
rank=global_rank,
103+
)
104+
else:
105+
dist.init_process_group(
106+
backend="nccl",
107+
world_size=world_size,
108+
rank=global_rank,
109+
device_id=torch.device(f"cuda:{local_rank}"),
110+
)
104111
args = get_args()
105112
init(args)
106113

vime/utils/arguments.py

Lines changed: 1 addition & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,6 @@
1010

1111
from vime.backends.vllm_utils.arguments import validate_args as vllm_validate_args
1212
from vime.backends.vllm_utils.arguments import vllm_parse_args
13-
from vime.utils.common import is_npu
1413
from vime.utils.eval_config import EvalDatasetConfig, build_eval_dataset_configs, ensure_dataset_list
1514
from vime.utils.logging_utils import configure_logger
1615

@@ -132,14 +131,10 @@ def add_train_arguments(parser):
132131
default=1024**3,
133132
help="Add margin for train memory allocation. By default we will reserve 1GB as margin.",
134133
)
135-
try:
136-
default_megatron_to_hf_mode = "bridge" if is_npu() else "raw"
137-
except RuntimeError:
138-
default_megatron_to_hf_mode = "raw"
139134
parser.add_argument(
140135
"--megatron-to-hf-mode",
141136
choices=["raw", "bridge"],
142-
default=default_megatron_to_hf_mode,
137+
default="raw",
143138
help="The method to convert megatron weights to hugging face weights for vLLM.",
144139
)
145140
parser.add_argument(

0 commit comments

Comments
 (0)