已合并
【bugfix】修复mindspore低版本兼容性问题 #947
curry808创建于 9 天前
【bugfix】修复mindspore低版本兼容性问题 #947
已合并
curry808创建于 9 天前
1 个文件变更+3-2
Mpython/msprobe/mindspore/dump/dump_processor/jit_dump.py+3-2
@@ -112,10 +112,11 @@ class JitDump(_MindsporeFunctionExecutor):
112 def grad(self, obj, grad, weights, grad_position, *args, **kwargs):112 def grad(self, obj, grad, weights, grad_position, *args, **kwargs):
113 if JitDump.jit_dump_switch and JitDump.jit_enable:113 if JitDump.jit_dump_switch and JitDump.jit_enable:
114 _api_register.restore_all_api()114 _api_register.restore_all_api()
115- if version_parse(mindspore.__version__) >= version_parse("2.5"):115+ if version_parse(mindspore.__version__) >= version_parse("2.9.0"):
116- # mindspore>=2.5 的调用方将 has_aux 作为参数传入(位于 *args 首位),需要原样透传给底层 executor
117 has_aux = args[0] if args else False116 has_aux = args[0] if args else False
118 output = self._executor.grad(grad, obj, weights, grad_position, has_aux, *args[1:], *(kwargs.values()))117 output = self._executor.grad(grad, obj, weights, grad_position, has_aux, *args[1:], *(kwargs.values()))
118+ elif version_parse(mindspore.__version__) >= version_parse("2.5"):
119+ output = self._executor.grad(grad, obj, weights, grad_position, False, *args, *(kwargs.values()))
119 else:120 else:
120 output = self._executor.grad(grad, obj, weights, grad_position, *args, *(kwargs.values()))121 output = self._executor.grad(grad, obj, weights, grad_position, *args, *(kwargs.values()))
121 if JitDump.jit_dump_switch and JitDump.jit_enable:122 if JitDump.jit_dump_switch and JitDump.jit_enable: