已合并
[feat] Add disable_tensor_unsafe_check option to restart_device #36164
一只食醉狗创建于 5月20日
[feat] Add disable_tensor_unsafe_check option to restart_device #36164
已合并
共 3 个文件变更+46-6
| @@ -1,7 +1,12 @@ | |||
| 1 | import torch | 1 | import torch |
| 2 | import torch_npu | 2 | import torch_npu |
| 3 | 3 | ||
| 4 | -from torch_npu.npu._recovery import check_npu_tensor_is_safe, mark_all_npu_tensor_unsafe, set_npu_tensor_unsafe_check_flag | 4 | +from torch_npu.npu._recovery import ( |
| 5 | + check_npu_tensor_is_safe, | ||
| 6 | + mark_all_npu_tensor_unsafe, | ||
| 7 | + set_npu_tensor_unsafe_check_flag, | ||
| 8 | + get_npu_tensor_unsafe_check_flag, | ||
| 9 | +) | ||
| 5 | from torch_npu.testing.testcase import TestCase, run_tests | 10 | from torch_npu.testing.testcase import TestCase, run_tests |
| 6 | from torch_npu.npu._recovery import restart_device | 11 | from torch_npu.npu._recovery import restart_device |
| 7 | 12 | ||
| @@ -46,6 +51,30 @@ class TestNpu(TestCase): | |||
| 46 | with self.assertRaises(RuntimeError): | 51 | with self.assertRaises(RuntimeError): |
| 47 | check_npu_tensor_is_safe("invalid_tensor") | 52 | check_npu_tensor_is_safe("invalid_tensor") |
| 48 | 53 | ||
| 54 | + def test_restart_device_with_disable_tensor_unsafe_check(self): | ||
| 55 | + torch.npu.set_device(0) | ||
| 56 | + tensor = torch.randn(2, 3, device="npu:0") | ||
| 57 | + | ||
| 58 | + set_npu_tensor_unsafe_check_flag(False) | ||
| 59 | + self.assertTrue(check_npu_tensor_is_safe(tensor)) | ||
| 60 | + self.assertFalse(get_npu_tensor_unsafe_check_flag()) | ||
| 61 | + | ||
| 62 | + restart_device(0, rebuild_all_resources=True, disable_tensor_unsafe_check=True) | ||
| 63 | + self.assertTrue(check_npu_tensor_is_safe(tensor)) | ||
| 64 | + self.assertFalse(get_npu_tensor_unsafe_check_flag()) | ||
| 65 | + | ||
| 66 | + set_npu_tensor_unsafe_check_flag(False) | ||
| 67 | + restart_device(0, rebuild_all_resources=True, disable_tensor_unsafe_check=False) | ||
| 68 | + self.assertFalse(check_npu_tensor_is_safe(tensor)) | ||
| 69 | + self.assertTrue(get_npu_tensor_unsafe_check_flag()) | ||
| 70 | + | ||
| 71 | + set_npu_tensor_unsafe_check_flag(False) | ||
| 72 | + restart_device(0, rebuild_all_resources=True) | ||
| 73 | + self.assertFalse(check_npu_tensor_is_safe(tensor)) | ||
| 74 | + self.assertTrue(get_npu_tensor_unsafe_check_flag()) | ||
| 75 | + | ||
| 76 | + set_npu_tensor_unsafe_check_flag(False) | ||
| 77 | + | ||
| 49 | 78 | ||
| 50 | if __name__ == '__main__': | 79 | if __name__ == '__main__': |
| 51 | run_tests() | 80 | run_tests() |
| @@ -1152,7 +1152,7 @@ | |||
| 1152 | "signature": "(device=None)" | 1152 | "signature": "(device=None)" |
| 1153 | }, | 1153 | }, |
| 1154 | "torch_npu.npu.restart_device": { | 1154 | "torch_npu.npu.restart_device": { |
| 1155 | - "signature": "(device_id: int, rebuild_all_resources: int = False)" | 1155 | + "signature": "(device_id: int, rebuild_all_resources: bool = False, disable_tensor_unsafe_check: bool = False)" |
| 1156 | }, | 1156 | }, |
| 1157 | "torch_npu.npu.stop_device": { | 1157 | "torch_npu.npu.stop_device": { |
| 1158 | "signature": "(device_id)" | 1158 | "signature": "(device_id)" |
| @@ -56,12 +56,23 @@ def _recovery_all_npu_stream(device: int) -> None: | |||
| 56 | return torch_npu._C._recovery_all_npu_stream(device) | 56 | return torch_npu._C._recovery_all_npu_stream(device) |
| 57 | 57 | ||
| 58 | 58 | ||
| 59 | -def restart_device(device_id: int, rebuild_all_resources: int = False): | 59 | +def restart_device( |
| 60 | - logger.info(f"restart device start, device_id={device_id}, rebuild_all_resources={rebuild_all_resources}") | 60 | + device_id: int, |
| 61 | + rebuild_all_resources: bool = False, | ||
| 62 | + disable_tensor_unsafe_check: bool = False, | ||
| 63 | + ): | ||
| 64 | + logger.info( | ||
| 65 | + "restart device start, device_id=%s, rebuild_all_resources=%s, " | ||
| 66 | + "disable_tensor_unsafe_check=%s", | ||
| 67 | + device_id, | ||
| 68 | + rebuild_all_resources, | ||
| 69 | + disable_tensor_unsafe_check, | ||
| 70 | + ) | ||
| 61 | torch_npu.npu._lazy_init() | 71 | torch_npu.npu._lazy_init() |
| 62 | if rebuild_all_resources: | 72 | if rebuild_all_resources: |
| 63 | - mark_all_npu_tensor_unsafe(device_id) | 73 | + if not disable_tensor_unsafe_check: |
| 64 | - set_npu_tensor_unsafe_check_flag(True) | 74 | + mark_all_npu_tensor_unsafe(device_id) |
| 75 | + set_npu_tensor_unsafe_check_flag(True) | ||
| 65 | _recovery_all_npu_stream(device_id) | 76 | _recovery_all_npu_stream(device_id) |
| 66 | torch_npu._C._npu_restart_device(device_id) | 77 | torch_npu._C._npu_restart_device(device_id) |
| 67 | _except_handler.set_force_stop_exception(False) | 78 | _except_handler.set_force_stop_exception(False) |