已合并
refactor(npu): route RNG generator lookup through the accelerator entry point as upstream does #44830
refactor(npu): route RNG generator lookup through the accelerator entry point as upstream does #44830
已合并
dwoai22创建于 8月18日
共 3 个文件变更+33-29
@@ -0,0 +1,15 @@
1+import torch
2+ 
3+from torch_npu._compat.version import CURRENT_VERSION
4+ 
5+__all__ = ["get_default_generator"]
6+ 
7+ 
8+# COMPAT(>= 2.14): upstream added
9+# torch._C._accelerator_getDefaultGenerator, the unified entry point for a
10+# backend's default generator. 2.13 and earlier have no accelerator-level
11+# equivalent, so fall back to the NPU-specific default_generators tuple.
12+def get_default_generator(device_index: int):
13+ if CURRENT_VERSION >= (2, 14):
14+ return torch._C._accelerator_getDefaultGenerator(device_index)
15+ return torch.npu.default_generators[device_index]
@@ -155,6 +155,7 @@ import torch
155from torch.storage import _LegacyStorage, _warn_typed_storage_removal155from torch.storage import _LegacyStorage, _warn_typed_storage_removal
156from torch._utils import classproperty156from torch._utils import classproperty
157from torch_npu._init.common.warning_utils import _should_print_warning157from torch_npu._init.common.warning_utils import _should_print_warning
158+from torch_npu._compat.accelerator import get_default_generator
158 159 
159import torch_npu160import torch_npu
160from torch_npu.utils._error_code import ErrCode, pta_error, prof_error161from torch_npu.utils._error_code import ErrCode, pta_error, prof_error
@@ -366,7 +367,7 @@ def _get_generator(device: torch.device) -> torch._C.Generator:
366 idx = device.index367 idx = device.index
367 if idx is None:368 if idx is None:
368 idx = current_device()369 idx = current_device()
369- return torch.npu.default_generators[idx]370+ return get_default_generator(idx)
370 371 
371 372 
372def _set_rng_state_offset(offset: int, device: Union[int, str, torch.device] = 'npu') -> None:373def _set_rng_state_offset(offset: int, device: Union[int, str, torch.device] = 'npu') -> None:
@@ -1,8 +1,9 @@
1-from typing import Iterable, Union1+from typing import Union
2import torch2import torch
3+from torch.accelerator._utils import _get_device_index
4+from torch_npu._compat.accelerator import get_default_generator
3 5 
4-import torch_npu6+from . import _lazy_init, _lazy_call, device_count, is_initialized
5-from . import _lazy_init, _lazy_call, device_count, current_device, is_initialized
6 7 
7__all__ = ['get_rng_state', 'set_rng_state',8__all__ = ['get_rng_state', 'set_rng_state',
8 'get_rng_state_all', 'set_rng_state_all',9 'get_rng_state_all', 'set_rng_state_all',
@@ -21,14 +22,8 @@ def get_rng_state(device: Union[int, str, torch.device] = 'npu') -> torch.Tensor
21 This function eagerly initializes NPU.22 This function eagerly initializes NPU.
22 """23 """
23 _lazy_init()24 _lazy_init()
24- if isinstance(device, str):25+ idx = _get_device_index(device, optional=True)
X
Xxuyun1519 天前

兼容性有问题,这里修改可能会导致

megatron/core/tensor_parallel/random.py", line 53, in _get_cuda_rng_state [rank4]: return torch.cuda.random.get_rng_state(device=device) [rank4]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ [rank4]: File "/usr/local/lib/python3.11/site-packages/torch_npu/npu/random.py", line 25, in get_rng_state [rank4]: idx = _get_device_index(device, optional=True) [rank4]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ [rank4]: File "/usr/local/lib/python3.11/site-packages/torch/accelerator/_utils.py", line 16, in _get_device_index [rank4]: raise ValueError( [rank4]: ValueError: cuda doesn't match the current accelerator npu. when instantiating RowParallelLinear

likedislike
25- device = torch.device(device)26+ default_generator = get_default_generator(idx)
26- elif isinstance(device, int):
27- device = torch.device('npu', device)
28- idx = device.index
29- if idx is None:
30- idx = current_device()
31- default_generator = torch_npu.npu.default_generators[idx]
32 return default_generator.get_state()27 return default_generator.get_state()
33 28 
34 29 
@@ -55,16 +50,9 @@ def set_rng_state(new_state: torch.Tensor, device: Union[int, str, torch.device]
55 # later when NPU is lazy initialized.50 # later when NPU is lazy initialized.
56 new_state = new_state.clone(memory_format=torch.contiguous_format)51 new_state = new_state.clone(memory_format=torch.contiguous_format)
57 52 
58- if isinstance(device, str):
59- device = torch.device(device)
60- elif isinstance(device, int):
61- device = torch.device('npu', device)
62- 
63 def cb():53 def cb():
64- idx = device.index54+ idx = _get_device_index(device, optional=True)
65- if idx is None:55+ default_generator = get_default_generator(idx)
66- idx = current_device()
67- default_generator = torch_npu.npu.default_generators[idx]
68 default_generator.set_state(new_state)56 default_generator.set_state(new_state)
69 57 
70 _lazy_call(cb)58 _lazy_call(cb)
@@ -95,8 +83,8 @@ def manual_seed(seed):
95 seed = int(seed)83 seed = int(seed)
96 84 
97 def cb():85 def cb():
98- idx = current_device()86+ idx = torch.accelerator.current_device_index()
99- default_generator = torch_npu.npu.default_generators[idx]87+ default_generator = get_default_generator(idx)
100 default_generator.manual_seed(seed)88 default_generator.manual_seed(seed)
101 89 
102 _lazy_call(cb)90 _lazy_call(cb)
@@ -114,7 +102,7 @@ def manual_seed_all(seed):
114 102 
115 def cb():103 def cb():
116 for i in range(device_count()):104 for i in range(device_count()):
117- default_generator = torch_npu.npu.default_generators[i]105+ default_generator = get_default_generator(i)
118 default_generator.manual_seed(seed)106 default_generator.manual_seed(seed)
119 107 
120 _lazy_call(cb)108 _lazy_call(cb)
@@ -131,8 +119,8 @@ def seed():
131 """119 """
132 120 
133 def cb():121 def cb():
134- idx = current_device()122+ idx = torch.accelerator.current_device_index()
135- default_generator = torch_npu.npu.default_generators[idx]123+ default_generator = get_default_generator(idx)
136 default_generator.seed()124 default_generator.seed()
137 125 
138 _lazy_call(cb)126 _lazy_call(cb)
@@ -148,7 +136,7 @@ def seed_all():
148 random_seed = 0136 random_seed = 0
149 seeded = False137 seeded = False
150 for i in range(device_count()):138 for i in range(device_count()):
151- default_generator = torch_npu.npu.default_generators[i]139+ default_generator = get_default_generator(i)
152 if not seeded:140 if not seeded:
153 default_generator.seed()141 default_generator.seed()
154 random_seed = default_generator.initial_seed()142 random_seed = default_generator.initial_seed()
@@ -166,6 +154,6 @@ def initial_seed():
166 This function eagerly initializes NPU.154 This function eagerly initializes NPU.
167 """155 """
168 _lazy_init()156 _lazy_init()
169- idx = current_device()157+ idx = torch.accelerator.current_device_index()
170- default_generator = torch_npu.npu.default_generators[idx]158+ default_generator = get_default_generator(idx)
171 return default_generator.initial_seed()159 return default_generator.initial_seed()