已合并
use torch_npu._C instead of option #43706
huangyunlong创建于 8月4日
use torch_npu._C instead of option #43706
已合并
共 7 个文件变更+59-10
| @@ -122,7 +122,7 @@ class TorchFillUninitializedMemoryTestCase(TestCase): | |||
| 122 | self.assertTrue(res.eq(max_val).all()) | 122 | self.assertTrue(res.eq(max_val).all()) |
| 123 | 123 | ||
| 124 | 124 | ||
| 125 | -instantiate_device_type_tests(TorchFillUninitializedMemoryTestCase, globals(), only_for='privateuse1') | 125 | +instantiate_device_type_tests(TorchFillUninitializedMemoryTestCase, globals(), only_for=('privateuse1',)) |
| 126 | 126 | ||
| 127 | 127 | ||
| 128 | if __name__ == '__main__': | 128 | if __name__ == '__main__': |
| @@ -3,6 +3,7 @@ | |||
| 3 | 3 | ||
| 4 | 4 | ||
| 5 | 5 | ||
| 6 | + | ||
| 6 | 7 | ||
| 7 | 8 | ||
| 8 | 9 | ||
| @@ -62,7 +63,7 @@ const at::Tensor& NPUNativeFunctions::resize_( | |||
| 62 | // fill newly added storage section with NaN/MAX_INT when requested. | 63 | // fill newly added storage section with NaN/MAX_INT when requested. |
| 63 | if (C10_UNLIKELY(at::globalContext().deterministicAlgorithms() && | 64 | if (C10_UNLIKELY(at::globalContext().deterministicAlgorithms() && |
| 64 | at::globalContext().deterministicFillUninitializedMemory() && | 65 | at::globalContext().deterministicFillUninitializedMemory() && |
| 65 | - at_npu::native::env::CheckFillUninitializedMemory())) { | 66 | + at_npu::native::env::globalNpuContext().npuFillUninitializedMemory())) { |
| 66 | at::native::fill_resize_deterministic_(self, static_cast<int64_t>(old_storage_nbytes)); | 67 | at::native::fill_resize_deterministic_(self, static_cast<int64_t>(old_storage_nbytes)); |
| 67 | } | 68 | } |
| 68 | return self; | 69 | return self; |
| @@ -120,7 +120,7 @@ inline bool should_fill_empty_deterministic() | |||
| 120 | { | 120 | { |
| 121 | return at::globalContext().deterministicAlgorithms() && | 121 | return at::globalContext().deterministicAlgorithms() && |
| 122 | at::globalContext().deterministicFillUninitializedMemory() && | 122 | at::globalContext().deterministicFillUninitializedMemory() && |
| 123 | - at_npu::native::env::CheckFillUninitializedMemory(); | 123 | + at_npu::native::env::globalNpuContext().npuFillUninitializedMemory(); |
| 124 | } | 124 | } |
| 125 | 125 | ||
| 126 | void window_function_checks( | 126 | void window_function_checks( |
| @@ -321,9 +321,6 @@ bool CheckCompatibleImplFor(CompatibleKey compatible_key) | |||
| 321 | return CheckCompatibleImpl(); | 321 | return CheckCompatibleImpl(); |
| 322 | } | 322 | } |
| 323 | 323 | ||
| 324 | -TORCH_NPU_REGISTER_OPTION(TORCH_NPU_FILL_UNINITIALIZED_MEMORY) | ||
| 325 | -REGISTER_OPTION_BOOL_FUNCTION(CheckFillUninitializedMemory, TORCH_NPU_FILL_UNINITIALIZED_MEMORY, "0", "1") | ||
| 326 | - | ||
| 327 | REGISTER_OPTION_HOOK(ALLOW_CONV_HF32, [](const std::string &val) { | 324 | REGISTER_OPTION_HOOK(ALLOW_CONV_HF32, [](const std::string &val) { |
| 328 | static const std::string mm_hf32_option_name = "ALLOW_MATMUL_HF32"; | 325 | static const std::string mm_hf32_option_name = "ALLOW_MATMUL_HF32"; |
| 329 | auto mm_hf32_val = c10_npu::option::GetOption(mm_hf32_option_name); | 326 | auto mm_hf32_val = c10_npu::option::GetOption(mm_hf32_option_name); |
| @@ -364,6 +361,19 @@ REGISTER_OPTION_HOOK(ACL_OP_DEBUG_OPTION, [](const std::string &val) { | |||
| 364 | 361 | ||
| 365 | TORCH_NPU_REGISTER_OPTION(CUBE_MATH_TYPE) | 362 | TORCH_NPU_REGISTER_OPTION(CUBE_MATH_TYPE) |
| 366 | 363 | ||
| 364 | +NpuContext& globalNpuContext() { | ||
| 365 | + static NpuContext globalNpuContext_; | ||
| 366 | + return globalNpuContext_; | ||
| 367 | +} | ||
| 368 | + | ||
| 369 | +bool NpuContext::npuFillUninitializedMemory() const { | ||
| 370 | + return _npu_fill_uninitialized_memory; | ||
| 371 | +} | ||
| 372 | + | ||
| 373 | +void NpuContext::setNpuFillUninitializedMemory(bool b) { | ||
| 374 | + _npu_fill_uninitialized_memory = b; | ||
| 375 | +} | ||
| 376 | + | ||
| 367 | } // namespace env | 377 | } // namespace env |
| 368 | } // namespace native | 378 | } // namespace native |
| 369 | } // namespace at_npu | 379 | } // namespace at_npu |
| @@ -32,12 +32,22 @@ enum class CompatibleKey { | |||
| 32 | bool CheckCompatibleImpl(); | 32 | bool CheckCompatibleImpl(); |
| 33 | bool CheckCompatibleImplFor(CompatibleKey compatible_key); | 33 | bool CheckCompatibleImplFor(CompatibleKey compatible_key); |
| 34 | bool CheckCompatibleImplBlackListFor(CompatibleKey compatible_key); | 34 | bool CheckCompatibleImplBlackListFor(CompatibleKey compatible_key); |
| 35 | -bool CheckFillUninitializedMemory(); | ||
| 36 | 35 | ||
| 37 | bool IsAllowFP32ToFP16(); | 36 | bool IsAllowFP32ToFP16(); |
| 38 | bool IsAllowConvHF32(); | 37 | bool IsAllowConvHF32(); |
| 39 | bool IsAllowMatmulHF32(); | 38 | bool IsAllowMatmulHF32(); |
| 40 | 39 | ||
| 40 | +class NpuContext { | ||
| 41 | + public: | ||
| 42 | + bool npuFillUninitializedMemory() const; | ||
| 43 | + void setNpuFillUninitializedMemory(bool /*b*/); | ||
| 44 | + | ||
| 45 | + private: | ||
| 46 | + bool _npu_fill_uninitialized_memory = false; | ||
| 47 | +}; | ||
| 48 | + | ||
| 49 | +NpuContext& globalNpuContext(); | ||
| 50 | + | ||
| 41 | } // namespace env | 51 | } // namespace env |
| 42 | } // namespace native | 52 | } // namespace native |
| 43 | } // namespace at_npu | 53 | } // namespace at_npu |
| @@ -2530,6 +2530,27 @@ PyObject* THNPModule_get_deterministic_level(PyObject* self, PyObject* noargs) { | |||
| 2530 | END_HANDLE_TH_ERRORS | 2530 | END_HANDLE_TH_ERRORS |
| 2531 | } | 2531 | } |
| 2532 | 2532 | ||
| 2533 | +static PyObject* THNPModule_setNpuFillUninitializedMemory( | ||
| 2534 | + PyObject* _unused, | ||
| 2535 | + PyObject* arg) { | ||
| 2536 | + HANDLE_TH_ERRORS | ||
| 2537 | + TORCH_CHECK( | ||
| 2538 | + PyBool_Check(arg), "expected a bool, but got ", THPUtils_typename(arg)); | ||
| 2539 | + at_npu::native::env::globalNpuContext().setNpuFillUninitializedMemory( | ||
| 2540 | + arg == Py_True); | ||
| 2541 | + Py_RETURN_NONE; | ||
| 2542 | + END_HANDLE_TH_ERRORS | ||
| 2543 | +} | ||
| 2544 | + | ||
| 2545 | +static PyObject* THNPModule_npuFillUninitializedMemory( | ||
| 2546 | + PyObject* _unused, | ||
| 2547 | + PyObject* noargs) { | ||
| 2548 | + if (at_npu::native::env::globalNpuContext().npuFillUninitializedMemory()) | ||
| 2549 | + Py_RETURN_TRUE; | ||
| 2550 | + else | ||
| 2551 | + Py_RETURN_FALSE; | ||
| 2552 | +} | ||
| 2553 | + | ||
| 2533 | PyObject* THNPModule_hasPrimaryContext_wrap(PyObject* self, PyObject* arg) { | 2554 | PyObject* THNPModule_hasPrimaryContext_wrap(PyObject* self, PyObject* arg) { |
| 2534 | HANDLE_TH_ERRORS | 2555 | HANDLE_TH_ERRORS |
| 2535 | TORCH_CHECK( | 2556 | TORCH_CHECK( |
| @@ -2860,6 +2881,14 @@ static struct PyMethodDef THNPModule_methods[] = { | |||
| 2860 | (PyCFunction)THNPModule_get_deterministic_level, | 2881 | (PyCFunction)THNPModule_get_deterministic_level, |
| 2861 | METH_NOARGS, | 2882 | METH_NOARGS, |
| 2862 | nullptr}, | 2883 | nullptr}, |
| 2884 | + {"_npu_set_fill_uninitialized_memory", | ||
| 2885 | + (PyCFunction)THNPModule_setNpuFillUninitializedMemory, | ||
| 2886 | + METH_O, | ||
| 2887 | + nullptr}, | ||
| 2888 | + {"_npu_get_fill_uninitialized_memory", | ||
| 2889 | + (PyCFunction)THNPModule_npuFillUninitializedMemory, | ||
| 2890 | + METH_NOARGS, | ||
| 2891 | + nullptr}, | ||
| 2863 | {"_npu_sleep", | 2892 | {"_npu_sleep", |
| 2864 | (PyCFunction)THNPModule_npuSleep, | 2893 | (PyCFunction)THNPModule_npuSleep, |
| 2865 | METH_O, | 2894 | METH_O, |
| @@ -2,7 +2,7 @@ def _add_deterministic_patch(): | |||
| 2 | """ | 2 | """ |
| 3 | Patch torch.utils.deterministic.fill_uninitialized_memory setter so that | 3 | Patch torch.utils.deterministic.fill_uninitialized_memory setter so that |
| 4 | setting it also synchronizes the NPU-side option | 4 | setting it also synchronizes the NPU-side option |
| 5 | - TORCH_NPU_FILL_UNINITIALIZED_MEMORY via torch_npu._C._npu_setOption. | 5 | + _npu_fill_uninitialized_memory via torch_npu._C._npu_set_fill_uninitialized_memory. |
| 6 | """ | 6 | """ |
| 7 | import torch_npu | 7 | import torch_npu |
| 8 | import torch.utils.deterministic as det_mod | 8 | import torch.utils.deterministic as det_mod |
| @@ -19,7 +19,6 @@ def _add_deterministic_patch(): | |||
| 19 | # 1. call original setter: torch._C._set_deterministic_fill_uninitialized_memory(mode) | 19 | # 1. call original setter: torch._C._set_deterministic_fill_uninitialized_memory(mode) |
| 20 | _orig_fset(self, mode) | 20 | _orig_fset(self, mode) |
| 21 | # 2. sync NPU-side option | 21 | # 2. sync NPU-side option |
| 22 | - option = {"TORCH_NPU_FILL_UNINITIALIZED_MEMORY": "1" if mode else "0"} | 22 | + torch_npu._C._npu_set_fill_uninitialized_memory(mode) |
| 23 | - torch_npu._C._npu_setOption(option) | ||
| 24 | 23 | ||
| 25 | _Deterministic.fill_uninitialized_memory = property(_orig_fget, _new_fset) | 24 | _Deterministic.fill_uninitialized_memory = property(_orig_fget, _new_fset) |