已合并
[feat] Add disable_tensor_unsafe_check option to restart_device #36164
一只食醉狗创建于 5月20日
[feat] Add disable_tensor_unsafe_check option to restart_device #36164
已合并
一只食醉狗创建于 5月20日
3 个文件变更+46-6
Mtest/npu/test_recovery.py+30-1
@@ -1,7 +1,12 @@
1import torch1import torch
2import torch_npu2import 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_flag4+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+)
5from torch_npu.testing.testcase import TestCase, run_tests10from torch_npu.testing.testcase import TestCase, run_tests
6from torch_npu.npu._recovery import restart_device11from 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 
50if __name__ == '__main__':79if __name__ == '__main__':
51 run_tests()80 run_tests()
Mtest/torch_npu_schema.json+1-1
@@ -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)"
Mtorch_npu/npu/_recovery.py+15-4
@@ -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)