Skip to content

Commit ce82cc7

Browse files
committed
feat: add initial support for Ascend devices with unified device abstraction
This commit introduces native support for Ascend NPUs in the ROLL project, while preserving compatibility with existing CUDA-based infrastructure. Key changes include: - Introduced a unified device abstraction interface to encapsulate device initialization, memory management, and synchronization, enabling extensibility for both CUDA and Ascend devices. - Replaced direct usage of and Ray CUDA resource APIs with the new abstraction layer to support multi-device environments. - Integrated Ascend inference backend via vLLM + vLLM-ascend. - Added experimental support for training with MindSpeed on Ascend hardware. This enhancement lays the groundwork for seamless switching across CUDA and Ascend devices. Signed-off-by: noemotiovon <757486878@qq.com>
1 parent 30ec292 commit ce82cc7

43 files changed

Lines changed: 442 additions & 176 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

mcore_adapter/src/mcore_adapter/initialize.py

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,8 @@
88
from .training_args import TrainingArguments
99
from .utils import get_logger
1010

11+
from roll.platforms import current_platform
12+
1113

1214
logger = get_logger(__name__)
1315

@@ -28,7 +30,7 @@ def _set_random_seed(seed_):
2830
random.seed(seed)
2931
np.random.seed(seed)
3032
torch.manual_seed(seed)
31-
if torch.cuda.device_count() > 0:
33+
if current_platform.device_count() > 0:
3234
tensor_parallel.model_parallel_cuda_manual_seed(seed)
3335
else:
3436
raise ValueError("Seed ({}) should be a positive integer.".format(seed))
@@ -45,10 +47,10 @@ def _initialize_distributed(args: "TrainingArguments"):
4547
logger.info(f"Initializing mpu on device {args.device}")
4648
if not torch.distributed.is_initialized():
4749
# Manually set the device ids.
48-
torch.cuda.set_device(args.device)
50+
current_platform.set_device(args.device)
4951
# Call the init process
5052
torch.distributed.init_process_group(
51-
backend=args.ddp_backend or "nccl",
53+
backend=args.ddp_backend or current_platform.communication_backend,
5254
rank=int(os.getenv("RANK", "0")),
5355
world_size=int(os.getenv("WORLD_SIZE", "1")),
5456
timeout=args.ddp_timeout_delta,

mcore_adapter/src/mcore_adapter/models/converter/convert_utils.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
import torch.distributed as dist
88
from megatron.core import mpu
99
from packaging.version import Version as PkgVersion
10+
from roll.platforms import current_platform
1011

1112

1213
if TYPE_CHECKING:
@@ -232,7 +233,7 @@ class StackedTensors:
232233

233234

234235
class TensorBucket:
235-
def __init__(self, bucket_size, device="cuda"):
236+
def __init__(self, bucket_size, device=current_platform.device_type):
236237
self.buffer = torch.empty(bucket_size, dtype=torch.int8, device=device)
237238
self.device = device
238239
self.bucket_size = bucket_size

mcore_adapter/src/mcore_adapter/models/model_factory.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@
2525
from .converter.model_converter import ModelConverter
2626
from .model_config import McaModelConfig
2727
from .model_utils import ModuleUtilsMixin, RMSNorm, exists_hf_config, exists_mca_config
28+
from roll.platforms import current_platform
2829

2930

3031
if TYPE_CHECKING:
@@ -279,7 +280,7 @@ def __init__(self, config: "McaModelConfig", **kwargs):
279280
for param in self.parameters():
280281
tensor_parallel.set_defaults_if_not_set_tensor_model_parallel_attributes(param)
281282
if not config.use_cpu_initialization:
282-
self.cuda(torch.cuda.current_device())
283+
self.cuda(current_platform.current_device())
283284

284285
def _get_transformer_layer_spec(self, config: Optional["McaModelConfig"]=None):
285286
config = config or self.config

mcore_adapter/src/mcore_adapter/models/model_utils.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77

88
from ..constants import MCA_CONFIG_NAME
99
from ..utils import get_logger
10-
10+
from roll.platforms import current_platform
1111

1212
if TYPE_CHECKING:
1313
from megatron.core.transformer import TransformerConfig
@@ -91,7 +91,7 @@ def floating_point_ops(
9191
class RMSNorm(nn.Module):
9292
def __init__(self, config: "TransformerConfig", hidden_size, eps=1e-6, **kwargs):
9393
super().__init__()
94-
device = torch.cuda.current_device() if not config.use_cpu_initialization else None
94+
device = current_platform.current_device() if not config.use_cpu_initialization else None
9595
self.weight = torch.nn.Parameter(torch.ones(hidden_size, dtype=config.params_dtype, device=device))
9696
self.variance_epsilon = eps
9797

mcore_adapter/src/mcore_adapter/models/qwen2_5_vl/modeling_qwen2_5_vl.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
from megatron.core import mpu
55
from megatron.core.transformer.attention import SelfAttention
66
from torch import nn
7+
from roll.platforms import current_platform
78

89
from ..auto.modeling_auto import register_model
910
from ..model_factory import McaGPTModel
@@ -58,7 +59,7 @@ def __init__(
5859
if rotary_percent < 1.0:
5960
dim = int(dim * rotary_percent)
6061

61-
device = "cpu" if use_cpu_initialization else torch.cuda.current_device()
62+
device = "cpu" if use_cpu_initialization else current_platform.current_device()
6263
self.inv_freq = 1.0 / (rotary_base ** (torch.arange(0, dim, 2, dtype=torch.float32, device=device) / dim))
6364

6465
@torch.no_grad()
@@ -199,7 +200,7 @@ def __init__(self, config: "Qwen2_5_VLConfig", **kwargs):
199200
Qwen2_5_VLVisionConfig(**config.vision_config),
200201
attn_implementation="flash_attention_2",
201202
torch_dtype=self.config.params_dtype,
202-
).to(torch.cuda.current_device())
203+
).to(current_platform.current_device())
203204
for param in self.vision_model.parameters():
204205
setattr(param, "sequence_parallel", config.sequence_parallel)
205206

mcore_adapter/src/mcore_adapter/models/qwen2_vl/modeling_qwen2_vl.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
from megatron.core import mpu
55
from megatron.core.transformer.attention import SelfAttention
66
from torch import nn
7+
from roll.platforms import current_platform
78

89
from ..auto.modeling_auto import register_model
910
from ..model_factory import McaGPTModel
@@ -61,7 +62,7 @@ def __init__(
6162
self.rotary_interleaved = rotary_interleaved
6263

6364
self.seq_len_interpolation_factor = seq_len_interpolation_factor
64-
device = "cpu" if use_cpu_initialization else torch.cuda.current_device()
65+
device = "cpu" if use_cpu_initialization else current_platform.current_device()
6566
self.inv_freq = 1.0 / (rotary_base ** (torch.arange(0, dim, 2, dtype=torch.float32, device=device) / dim))
6667

6768
@torch.no_grad()
@@ -204,7 +205,7 @@ def __init__(self, config: "Qwen2VLConfig", **kwargs):
204205
Qwen2VLVisionConfig(**config.vision_config),
205206
attn_implementation="sdpa",
206207
torch_dtype=self.config.params_dtype,
207-
).to(torch.cuda.current_device())
208+
).to(current_platform.current_device())
208209
for param in self.vision_model.parameters():
209210
setattr(param, "sequence_parallel", config.sequence_parallel)
210211

mcore_adapter/src/mcore_adapter/trainer/trainer.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,7 @@
3636
seed_worker,
3737
speed_metrics,
3838
)
39+
from roll.platforms import current_platform
3940

4041
from ..checkpointing import get_checkpoint_dir, load_state_dict_from_checkpoint
4142
from ..constants import DIST_OPTIMIZER_DIR, IGNORE_INDEX
@@ -501,7 +502,7 @@ def _save_rng_state(self, output_dir):
501502
"random_rng_state": random.getstate(),
502503
"np_rng_state": np.random.get_state(),
503504
"torch_rng_state": torch.get_rng_state(),
504-
"cuda_rng_state": torch.cuda.get_rng_state(),
505+
"cuda_rng_state": current_platform.get_rng_state(),
505506
"rng_tracker_states": tensor_parallel.get_cuda_rng_tracker().get_states(),
506507
}
507508
if self.args.world_size <= 1:
@@ -537,7 +538,7 @@ def _load_rng_state(self, checkpoint):
537538
random.setstate(checkpoint_rng_state["random_rng_state"])
538539
np.random.set_state(checkpoint_rng_state["np_rng_state"])
539540
torch.set_rng_state(checkpoint_rng_state["torch_rng_state"])
540-
torch.cuda.set_rng_state(checkpoint_rng_state["cuda_rng_state"])
541+
current_platform.set_rng_state(checkpoint_rng_state["cuda_rng_state"])
541542
# Check for empty states array
542543
if not checkpoint_rng_state["rng_tracker_states"]:
543544
raise KeyError

roll/distributed/executor/cluster.py

Lines changed: 20 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
collect_all_to_all,
1818
dispatch_one_to_all,
1919
)
20+
from roll.platforms import current_platform
2021
from roll.utils.constants import RAY_NAMESPACE
2122
from roll.distributed.scheduler.resource_manager import ResourceManager
2223
from roll.utils.import_utils import safe_import_class
@@ -122,14 +123,25 @@ def _create_workers(self):
122123

123124
runtime_env = RuntimeEnv(env_vars=env_vars)
124125
self.worker_config.resource_placement_groups = pgs
125-
worker = self.worker_cls.options(
126-
scheduling_strategy=PlacementGroupSchedulingStrategy(placement_group=deploy_pg["placement_group"]),
127-
name=worker_name,
128-
namespace=RAY_NAMESPACE,
129-
runtime_env=runtime_env,
130-
num_cpus=0.01,
131-
num_gpus=0.01 if self.worker_config.device_mapping else 0,
132-
).remote(worker_config=self.worker_config)
126+
if current_platform.ray_device_key == "GPU":
127+
worker = self.worker_cls.options(
128+
scheduling_strategy=PlacementGroupSchedulingStrategy(placement_group=deploy_pg["placement_group"]),
129+
name=worker_name,
130+
namespace=RAY_NAMESPACE,
131+
runtime_env=runtime_env,
132+
num_cpus=0.01,
133+
num_gpus=0.01 if self.worker_config.device_mapping else 0,
134+
).remote(worker_config=self.worker_config)
135+
else:
136+
worker = self.worker_cls.options(
137+
scheduling_strategy=PlacementGroupSchedulingStrategy(placement_group=deploy_pg["placement_group"]),
138+
name=worker_name,
139+
namespace=RAY_NAMESPACE,
140+
runtime_env=runtime_env,
141+
num_cpus=0.01,
142+
num_gpus=0,
143+
resources={current_platform.ray_device_key: 0.01 if self.worker_config.device_mapping else 0},
144+
).remote(worker_config=self.worker_config)
133145
self.workers.append(worker)
134146
if rank == 0:
135147
self.master_addr, self.master_port = ray.get(worker.get_master_addr_and_port.remote())

roll/distributed/executor/worker.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
from roll.utils.context_managers import state_offload_manger
1717
from roll.utils.logging import get_logger
1818
from roll.utils.offload_states import OffloadStateType
19+
from roll.platforms import current_platform
1920
from roll.utils.ray_utils import RayUtils
2021

2122

roll/distributed/scheduler/decorator.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616

1717
from roll.distributed.scheduler.protocol import DataProto, ObjectRefWrap
1818
from roll.utils.logging import get_logger
19+
from roll.platforms import current_platform
1920

2021
logger = get_logger()
2122

@@ -274,9 +275,9 @@ async def inner_async(*args, **kwargs):
274275
result = await func(*args, **kwargs)
275276
if clear_cache:
276277
try:
277-
torch._C._cuda_clearCublasWorkspaces()
278+
current_platform._cuda_clearCublasWorkspaces()
278279
gc.collect()
279-
torch.cuda.empty_cache()
280+
current_platform.empty_cache()
280281
except Exception as oe:
281282
pass
282283

@@ -295,9 +296,9 @@ def inner(*args, **kwargs):
295296
result = func(*args, **kwargs)
296297
if clear_cache:
297298
try:
298-
torch._C._cuda_clearCublasWorkspaces()
299+
current_platform._cuda_clearCublasWorkspaces()
299300
gc.collect()
300-
torch.cuda.empty_cache()
301+
current_platform.empty_cache()
301302
except Exception as oe:
302303
pass
303304

0 commit comments

Comments
 (0)