已合并
refactor(npu): route RNG generator lookup through the accelerator entry point as upstream does #44830
dwoai22创建于 8月18日
refactor(npu): route RNG generator lookup through the accelerator entry point as upstream does #44830
已合并
共 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 | |||
| 155 | from torch.storage import _LegacyStorage, _warn_typed_storage_removal | 155 | from torch.storage import _LegacyStorage, _warn_typed_storage_removal |
| 156 | from torch._utils import classproperty | 156 | from torch._utils import classproperty |
| 157 | from torch_npu._init.common.warning_utils import _should_print_warning | 157 | from torch_npu._init.common.warning_utils import _should_print_warning |
| 158 | +from torch_npu._compat.accelerator import get_default_generator | ||
| 158 | 159 | ||
| 159 | import torch_npu | 160 | import torch_npu |
| 160 | from torch_npu.utils._error_code import ErrCode, pta_error, prof_error | 161 | from 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.index | 367 | 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 | ||
| 372 | def _set_rng_state_offset(offset: int, device: Union[int, str, torch.device] = 'npu') -> None: | 373 | def _set_rng_state_offset(offset: int, device: Union[int, str, torch.device] = 'npu') -> None: |
| @@ -1,8 +1,9 @@ | |||
| 1 | -from typing import Iterable, Union | 1 | +from typing import Union |
| 2 | import torch | 2 | import 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_npu | 6 | +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 | |||
| 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.index | 54 | + 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 = 0 | 136 | random_seed = 0 |
| 149 | seeded = False | 137 | 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() |
兼容性有问题,这里修改可能会导致
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