已合并
refactor: remove obsolete sanitizer autograd workaround #43468
refactor: remove obsolete sanitizer autograd workaround #43468
已合并
No_neck创建于 7月31日
共 2 个文件变更+6-31
@@ -1,12 +1,6 @@
1-import sys
W
Wwanglijun558月1日
已过期

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

likedislike
No_neck
8月3日 评论:
2-import logging
3from unittest import mock1from unittest import mock
4-import unittest
5 2 
6import torch3import torch
7-import torch.cuda._sanitizer as csan
8-from torch.utils._python_dispatch import TorchDispatchMode
9-import torch_npu
10import torch_npu.npu._stream_check as stream_check4import torch_npu.npu._stream_check as stream_check
11from torch_npu.testing.testcase import TestCase, run_tests5from 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 pass42 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 
65if __name__ == "__main__":51if __name__ == "__main__":
66- run_tests()52+ run_tests()
W
Wwanglijun557月31日
已过期

这些是代码风格的改动,是否考虑不在PR中带入

likedislike
@@ -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_write364 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_handler380 self.event_handler = event_handler
382 self.args_handler = None381 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 kwargs384 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 operator394 # execute operator
406 outputs = func(*args, **kwargs)395 outputs = func(*args, **kwargs)
@@ -417,10 +406,10 @@ class NPUSanitizerDispatchMode(TorchDispatchMode):
417 npu_stream = 0406 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
Wwanglijun557月31日
已过期

为啥要去掉as err,去掉后会导致异常信息无法打印

likedislike
No_neck
7月31日 评论:
421 logger.info(410 logger.info(
N
NNo_neck8月1日
	except RuntimeError as err:
        logger.info(
            "Failed to get current stream, ignore this kernel launch record. error info is: %s",
            err
        )

上面注释的代码为旧版本代码,旧版本代码无法通过ci门禁,会触发flake8-logging-format 插件的一条日志规范(G200)——Logging statements should not include the exception in logged string; use exception() or exc_info=True. 故对该代码做以下修改:通过增加日志简要信息,方便开发者定位问题,若此处取消error的显示,可能会增加开发调试的难度。

likedislike
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- err412+ exc_info=True,
424 )413 )
425 return outputs414 return outputs
426 self.check_errors(func, npu_stream)415 self.check_errors(func, npu_stream)