已合并
refactor: remove obsolete sanitizer autograd workaround #43468
No_neck创建于 7月31日
refactor: remove obsolete sanitizer autograd workaround #43468
已合并
共 2 个文件变更+6-31
| @@ -1,12 +1,6 @@ | |||
| 1 | -import sys | ||
W | |||
| 2 | -import logging | ||
| 3 | from unittest import mock | 1 | from unittest import mock |
| 4 | -import unittest | ||
| 5 | 2 | ||
| 6 | import torch | 3 | import torch |
| 7 | -import torch.cuda._sanitizer as csan | ||
| 8 | -from torch.utils._python_dispatch import TorchDispatchMode | ||
| 9 | -import torch_npu | ||
| 10 | import torch_npu.npu._stream_check as stream_check | 4 | import torch_npu.npu._stream_check as stream_check |
| 11 | from torch_npu.testing.testcase import TestCase, run_tests | 5 | from torch_npu.testing.testcase import TestCase, run_tests |
| 12 | 6 | ||
| @@ -43,24 +37,16 @@ class TestStreamCheck(TestCase): | |||
| 43 | mock_outputs = [torch.tensor([4.0])] | 37 | mock_outputs = [torch.tensor([4.0])] |
| 44 | 38 | ||
| 45 | with mock.patch.object(mode, 'parse_inputs') as mock_parse_inputs, \ | 39 | with mock.patch.object(mode, 'parse_inputs') as mock_parse_inputs, \ |
| 46 | - mock.patch.object(mode, 'parse_outputs') as mock_parse_outputs, \ | 40 | + mock.patch.object(mode, 'parse_outputs') as mock_parse_outputs, \ |
| 47 | - mock.patch.object(mode, 'check_errors') as mock_check_errors: | 41 | + mock.patch.object(mode, 'check_errors') as mock_check_errors: |
| 48 | pass | 42 | pass |
| 49 | 43 | ||
| 50 | - def test_enable_autograd_with_matching_api(self): | ||
| 51 | - mock_event_handler = mock.MagicMock() | ||
| 52 | - mode = stream_check.NPUSanitizerDispatchMode(mock_event_handler) | ||
| 53 | - with mock.patch('torch._C._dispatch_tls_set_dispatch_key_excluded') as mock_set_dispatch: | ||
| 54 | - mode.enable_autograd("adaptive_avg_pool2d") | ||
| 55 | - mock_set_dispatch.assert_called_once_with(torch._C.DispatchKey.AutogradFunctionality, False) | ||
| 56 | - | ||
| 57 | def test_init_with_event_handler(self): | 44 | def test_init_with_event_handler(self): |
| 58 | mock_event_handler = mock.MagicMock() | 45 | mock_event_handler = mock.MagicMock() |
| 59 | mode = stream_check.NPUSanitizerDispatchMode(mock_event_handler) | 46 | mode = stream_check.NPUSanitizerDispatchMode(mock_event_handler) |
| 60 | self.assertEqual(mode.event_handler, mock_event_handler) | 47 | self.assertEqual(mode.event_handler, mock_event_handler) |
| 61 | self.assertIsNone(mode.args_handler) | 48 | self.assertIsNone(mode.args_handler) |
| 62 | - self.assertEqual(mode.npu_adjust_autograd, ["adaptive_avg_pool2d", "batch_norm", "log_softmax", "nll_loss", "to"]) | ||
| 63 | 49 | ||
| 64 | 50 | ||
| 65 | if __name__ == "__main__": | 51 | if __name__ == "__main__": |
| 66 | - run_tests() | 52 | + run_tests() |
W 已过期 这些是代码风格的改动,是否考虑不在PR中带入 ![]() ![]() | |||
| @@ -360,7 +360,6 @@ class NPUArgumentHandler: | |||
| 360 | ) | 360 | ) |
| 361 | 361 | ||
| 362 | def parse_outputs(self, schema, outputs, *, is_factory: bool = False) -> None: | 362 | def parse_outputs(self, schema, outputs, *, is_factory: bool = False) -> None: |
| 363 | - from torch.cuda._sanitizer import zip_arguments | ||
| 364 | for res, value in zip(schema.returns, (outputs,)): | 363 | for res, value in zip(schema.returns, (outputs,)): |
| 365 | metadata_only = res.alias_info is not None and not res.alias_info.is_write | 364 | metadata_only = res.alias_info is not None and not res.alias_info.is_write |
| 366 | pytree.tree_map_( | 365 | pytree.tree_map_( |
| @@ -380,14 +379,6 @@ class NPUSanitizerDispatchMode(TorchDispatchMode): | |||
| 380 | super().__init__() | 379 | super().__init__() |
| 381 | self.event_handler = event_handler | 380 | self.event_handler = event_handler |
| 382 | self.args_handler = None | 381 | self.args_handler = None |
| 383 | - self.npu_adjust_autograd = [ | ||
| 384 | - "adaptive_avg_pool2d", "batch_norm", | ||
| 385 | - "log_softmax", "nll_loss", "to" | ||
| 386 | - ] | ||
| 387 | - | ||
| 388 | - def enable_autograd(self, aten_api): | ||
| 389 | - if aten_api in self.npu_adjust_autograd: | ||
| 390 | - torch._C._dispatch_tls_set_dispatch_key_excluded(torch._C.DispatchKey.AutogradFunctionality, False) | ||
| 391 | 382 | ||
| 392 | def __torch_dispatch__(self, func, types, args=(), kwargs=None): | 383 | def __torch_dispatch__(self, func, types, args=(), kwargs=None): |
| 393 | kwargs = {} if kwargs is None else kwargs | 384 | kwargs = {} if kwargs is None else kwargs |
| @@ -399,8 +390,6 @@ class NPUSanitizerDispatchMode(TorchDispatchMode): | |||
| 399 | is_factory = bool(FACTORY_FUNCTION_REGEX.match(func._schema.name)) | 390 | is_factory = bool(FACTORY_FUNCTION_REGEX.match(func._schema.name)) |
| 400 | 391 | ||
| 401 | self.args_handler = NPUArgumentHandler() | 392 | self.args_handler = NPUArgumentHandler() |
| 402 | - aten_api = func.__name__.split(".")[0] | ||
| 403 | - self.enable_autograd(aten_api) | ||
| 404 | self.parse_inputs(func._schema, args, kwargs, is_factory=is_factory) | 393 | self.parse_inputs(func._schema, args, kwargs, is_factory=is_factory) |
| 405 | # execute operator | 394 | # execute operator |
| 406 | outputs = func(*args, **kwargs) | 395 | outputs = func(*args, **kwargs) |
| @@ -417,10 +406,10 @@ class NPUSanitizerDispatchMode(TorchDispatchMode): | |||
| 417 | npu_stream = 0 | 406 | npu_stream = 0 |
| 418 | try: | 407 | try: |
| 419 | npu_stream = int(torch_npu.npu.current_stream().npu_stream) | 408 | npu_stream = int(torch_npu.npu.current_stream().npu_stream) |
| 420 | - except RuntimeError as err: | 409 | + except RuntimeError: |
W 已过期 为啥要去掉as err,去掉后会导致异常信息无法打印 ![]() ![]() | |||
| 421 | logger.info( | 410 | logger.info( |
N
上面注释的代码为旧版本代码,旧版本代码无法通过ci门禁,会触发flake8-logging-format 插件的一条日志规范(G200)——Logging statements should not include the exception in logged string; use exception() or exc_info=True. 故对该代码做以下修改:通过增加日志简要信息,方便开发者定位问题,若此处取消error的显示,可能会增加开发调试的难度。 ![]() ![]() | |||
| 422 | - "Failed to get current stream, ignore this kernel launch record. error info is: %s", | 411 | + "Failed to get current stream, ignore this kernel launch record.", |
| 423 | - err | 412 | + exc_info=True, |
| 424 | ) | 413 | ) |
| 425 | return outputs | 414 | return outputs |
| 426 | self.check_errors(func, npu_stream) | 415 | self.check_errors(func, npu_stream) |


用理执行结果的截图可以放在PR描述的功能验证章节