Skip to content

Commit 445b688

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

2 files changed

Lines changed: 16 additions & 15 deletions

File tree

tools/convert_hf_to_torch_dist.py

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

55
import torch
66
import torch.distributed as dist
7+
from vime.utils.common import is_npu
8+
if is_npu():
9+
import mindspeed.megatron_adaptor
710
from megatron.core.enums import ModelType
811
from megatron.training.arguments import parse_args, validate_args
912
from megatron.training.checkpointing import get_checkpoint_name, get_checkpoint_tracker_filename, save_checkpoint
@@ -14,7 +17,6 @@
1417
from vime.backends.megatron_utils.arguments import set_default_megatron_args
1518
from vime.backends.megatron_utils.initialize import init
1619
from vime.backends.megatron_utils.model_provider import get_model_provider_func
17-
from vime.utils.common import is_npu
1820
from vime.utils.logging_utils import configure_logger
1921
from vime.utils.memory_utils import print_memory
2022

@@ -92,15 +94,19 @@ def main():
9294
os.environ.setdefault("LOCAL_RANK", str(local_rank))
9395
os.environ.setdefault("MASTER_ADDR", "localhost")
9496
os.environ.setdefault("MASTER_PORT", "12355")
95-
backend = "nccl"
9697
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+
)
104110
args = get_args()
105111
init(args)
106112

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)