Skip to content

Commit 2ee7599

Browse files
committed
add int64 to int32
Signed-off-by: noemotiovon <757486878@qq.com>
1 parent a958bf4 commit 2ee7599

5 files changed

Lines changed: 22 additions & 6 deletions

File tree

roll/distributed/scheduler/protocol.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,9 @@
1919

2020
from roll.utils.functionals import union_two_dict, divide_by_chunk_size
2121
from roll.platforms import current_platform
22+
from roll.utils.logging import get_logger
23+
24+
logger = get_logger()
2225

2326
try:
2427
tensordict.set_lazy_legacy(False).set()
@@ -248,6 +251,9 @@ def from_single_dict(cls, data: Dict[str, Union[torch.Tensor, np.ndarray]], meta
248251

249252
for key, val in data.items():
250253
if isinstance(val, torch.Tensor):
254+
if current_platform.is_npu and val.dtype == torch.int64:
255+
logger.debug(f"[NPU] Converting Tensor {key} from int64 -> int32, shape={val.shape}")
256+
val = val.to(torch.int32)
251257
tensors[key] = val
252258
elif isinstance(val, np.ndarray):
253259
non_tensors[key] = val

roll/pipeline/distill/distill_worker.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020
from roll.utils.cuda_ipc_utils import MultiprocessingSerializer
2121
from roll.utils.offload_states import OffloadStateType
2222
from roll.pipeline.distill.various_divergence import VariousDivergence, GPTLMLoss
23+
from roll.platforms import current_platform
2324

2425

2526

@@ -71,7 +72,7 @@ def train_step(self, data: DataProto):
7172
is_offload_states=is_offload_states,
7273
load_kwargs={"include": None},
7374
):
74-
data = data.to("cuda")
75+
data = data.to(current_platform.device_type)
7576
data = self.strategy.get_data_input(data)
7677
self.logger.info(f"global_step: {data.meta_info.get('global_step',0)}")
7778
per_device_train_batch_size = self.worker_config.training_args.per_device_train_batch_size
@@ -176,7 +177,7 @@ def forward(self, data: DataProto):
176177
is_offload_states=is_offload_states,
177178
load_kwargs={"include": None},
178179
):
179-
data = data.to("cuda")
180+
data = data.to(current_platform.device_type)
180181
data.meta_info["micro_batch_size"] = self.pipeline_config.student.training_args.per_device_train_batch_size
181182
data.meta_info["output_on_all_tp_ranks"] = True
182183
self.logger.info(f"global_step: {data.meta_info.get('global_step', 0)}")

roll/pipeline/dpo/actor_worker.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
from roll.distributed.scheduler.decorator import Dispatch, register
77
from roll.distributed.scheduler.protocol import DataProto
88
from roll.pipeline.base_worker import ActorWorker as BaseActorWorker
9+
from roll.platforms import current_platform
910
from roll.utils.context_managers import state_offload_manger
1011
from roll.utils.functionals import append_to_dict
1112
from roll.utils.offload_states import OffloadStateType
@@ -75,7 +76,7 @@ def train_step(self, data: DataProto):
7576
metric_infix=f"{self.cluster_name}/train_step",
7677
is_offload_states=is_offload_states,
7778
):
78-
data = data.to("cuda")
79+
data = data.to(current_platform.device_type)
7980
data = self.strategy.get_data_input(data)
8081
per_device_train_batch_size = self.worker_config.training_args.per_device_train_batch_size
8182
backward_batch_size = (
@@ -124,7 +125,7 @@ def compute_log_probs(self, data: DataProto):
124125
metric_infix=f"{self.cluster_name}/compute_log_probs",
125126
is_offload_states=is_offload_states,
126127
):
127-
data = data.to("cuda")
128+
data = data.to(current_platform.device_type)
128129
data.meta_info["micro_batch_size"] = self.worker_config.infer_batch_size
129130

130131
with torch.no_grad():

roll/platforms/npu.py

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -13,9 +13,13 @@ class NpuPlatform(Platform):
1313
ray_experimental_noset: str = "RAY_EXPERIMENTAL_NOSET_ASCEND_RT_VISIBLE_DEVICES"
1414
communication_backend: str = "hccl"
1515

16+
@classmethod
17+
def is_npu(cls) -> bool:
18+
return True
19+
1620
@classmethod
1721
def clear_cublas_workspaces(cls) -> None:
18-
pass
22+
return
1923

2024
@classmethod
2125
def get_vllm_worker_class(clas):
@@ -38,7 +42,7 @@ def get_vllm_worker_class(clas):
3842

3943
@classmethod
4044
def set_allocator_settings(cls) -> None:
41-
pass
45+
return
4246

4347
@classmethod
4448
def get_custom_env_vars(cls) -> dict:

roll/platforms/platform.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -87,6 +87,10 @@ def __getattr__(self, key: str):
8787
else:
8888
logger.warning("Current platform %s does not have '%s'" " attribute.", self.device_type, key)
8989
return None
90+
91+
@classmethod
92+
def is_npu(cls) -> bool:
93+
return False
9094

9195
@classmethod
9296
def clear_cublas_workspaces(cls) -> None:

0 commit comments

Comments
 (0)