diff --git a/docs/api/api_python/nn/mindspore.nn.Cell.rst b/docs/api/api_python/nn/mindspore.nn.Cell.rst index 30f19e8a26e..0a4e88875a2 100644 --- a/docs/api/api_python/nn/mindspore.nn.Cell.rst +++ b/docs/api/api_python/nn/mindspore.nn.Cell.rst @@ -317,12 +317,14 @@ .. py:method:: register_forward_pre_hook(hook_fn) - 设置Cell对象的正向pre_hook函数。此函数仅在PyNative模式下支持。 + 设置Cell对象的正向pre_hook函数。 .. note:: + - `register_forward_pre_hook(hook_fn)` 在图模式下,或者在PyNative模式下使用 `ms_function` 功能时不起作用。 - hook_fn必须有如下代码定义。 `cell_id` 是已注册Cell对象的信息,包括名称和ID。 `inputs` 是网络正向传播时Cell对象的输入数据。用户可以在hook_fn中打印输入数据或者返回新的输入数据。 - hook_fn返回新的输入数据或者None:hook_fn(cell_id, inputs) -> New inputs or None。 - - 为了避免脚本在切换到图模式时运行失败,不建议在Cell对象的 `construct` 函数中调用 `register_forward_pre_hook(hook_fn)`。 + - 为了避免脚本在切换到图模式时运行失败,不建议在Cell对象的 `construct` 函数中调用 `register_forward_pre_hook(hook_fn)` 。 + - PyNative模式下,如果在Cell对象的 `construct` 函数中调用 `register_forward_pre_hook(hook_fn)` ,那么Cell对象每次运行都将增加一个 `hook_fn` 。 **参数:** @@ -338,12 +340,14 @@ .. py:method:: register_forward_hook(hook_fn) - 设置Cell对象的正向hook函数。此函数仅在PyNative模式下支持。 + 设置Cell对象的正向hook函数。 .. note:: + - `register_forward_hook(hook_fn)` 在图模式下,或者在PyNative模式下使用 `ms_function` 功能时不起作用。 - hook_fn必须有如下代码定义。 `cell_id` 是已注册Cell对象的信息,包括名称和ID。 `inputs` 是网络正向传播时Cell对象的输入数据。 `outputs` 是网络正向传播时Cell对象的输出数据。用户可以在hook_fn中打印数据或者返回新的输出数据。 - hook_fn返回新的输出数据或者None:hook_fn(cell_id, inputs, outputs) -> New outputs or None。 - - 为了避免脚本在切换到图模式时运行失败,不建议在Cell对象的 `construct` 函数中调用 `register_forward_hook(hook_fn)`。 + - 为了避免脚本在切换到图模式时运行失败,不建议在Cell对象的 `construct` 函数中调用 `register_forward_hook(hook_fn)` 。 + - PyNative模式下,如果在Cell对象的 `construct` 函数中调用 `register_forward_hook(hook_fn)` ,那么Cell对象每次运行都将增加一个 `hook_fn` 。 **参数:** @@ -359,12 +363,14 @@ .. py:method:: register_backward_hook(hook_fn) - 设置Cell对象的反向hook函数。此函数仅在PyNative模式下支持。 + 设置Cell对象的反向hook函数。 .. note:: + - `register_backward_hook(hook_fn)` 在图模式下,或者在PyNative模式下使用 `ms_function` 功能时不起作用。 - hook_fn必须有如下代码定义。 `cell_id` 是已注册Cell对象的信息,包括名称和ID。 `grad_input` 是反向传递给Cell对象的梯度。 `grad_output` 是Cell对象的反向输出梯度。用户可以在hook_fn中打印梯度数据或者返回新的输出梯度。 - hook_fn返回新的输出梯度或者None:hook_fn(cell_id, grad_input, grad_output) -> New grad_output or None。 - - 为了避免脚本在切换到图模式时运行失败,不建议在Cell对象的 `construct` 函数中调用 `register_backward_hook(hook_fn)`。 + - 为了避免脚本在切换到图模式时运行失败,不建议在Cell对象的 `construct` 函数中调用 `register_backward_hook(hook_fn)` 。 + - PyNative模式下,如果在Cell对象的 `construct` 函数中调用 `register_backward_hook(hook_fn)` ,那么Cell对象每次运行都将增加一个 `hook_fn` 。 **参数:** @@ -384,16 +390,16 @@ 这个接口通常不需要显式调用。 - .. py:method:: run_construct(self, cast_inputs, kwargs) + .. py:method:: run_construct(cast_inputs, kwargs) 运行construct方法。 .. note:: - 此函数将会在未来版本中弃用,不推荐使用此函数。 + - 该函数已经弃用,将会在未来版本中删除,不推荐使用此函数。 **参数:** - - **cast_inputs** (tuple) – 输入的Cell对象。 + - **cast_inputs** (tuple) – Cell的输入。 - **kwargs** (dict) – 关键字参数。 **返回:** diff --git a/mindspore/python/mindspore/nn/cell.py b/mindspore/python/mindspore/nn/cell.py index caa0b747de1..57273fe4dcd 100755 --- a/mindspore/python/mindspore/nn/cell.py +++ b/mindspore/python/mindspore/nn/cell.py @@ -1554,7 +1554,7 @@ class Cell(Cell_): def _run_forward_pre_hook(self, inputs): """ - Running forward pre hook function registered on cell object. + Running forward pre hook function registered on Cell object. Args: inputs: The input objects of cell object. @@ -1577,17 +1577,20 @@ class Cell(Cell_): def register_forward_pre_hook(self, hook_fn): """ - Register forward pre hook function for cell object. Note that this function is only supported in pynative mode. + Register forward pre hook function for Cell object. Note: + - The `register_forward_pre_hook(hook_fn)` does not work in graph mode or ms_function. - 'hook_fn' must be defined as the following code. - `cell_id` is the information of registered cell object, including name and ID. `inputs` is the forward - input objects passed to the cell. The 'hook_fn' can modify the forward input objects by returning new + `cell_id` is the information of registered Cell object, including name and ID. `inputs` is the forward + input objects passed to the Cell. The 'hook_fn' can modify the forward input objects by returning new forward input objects. - It should have the following signature: hook_fn(cell_id, inputs) -> new input objects or none. - - In order to prevent running failed when switching to graph mode, it is not recommended to write in the - construct. + - In order to prevent running failed when switching to graph mode, it is not recommended to write it in the + `construct` function of Cell object. In the pynative mode, if the `register_forward_pre_hook` function is + called in the `construct` function of the Cell object, a hook function will be added at each run time of + Cell object. Args: hook_fn (function): Python function. Forward pre hook function. @@ -1653,11 +1656,11 @@ class Cell(Cell_): def _run_forward_hook(self, inputs, output): """ - Running forward hook function registered on cell object. + Running forward hook function registered on Cell object. Args: - inputs: The input objects of cell object. - output: The output object of cell object. + inputs: The input objects of Cell object. + output: The output object of Cell object. Returns: - **output** - New output object or none. @@ -1674,17 +1677,20 @@ class Cell(Cell_): def register_forward_hook(self, hook_fn): """ - Set the cell forward hook function. Note that this function is only supported in pynative mode. + Set the Cell forward hook function. Note: + - The `register_forward_hook(hook_fn)` does not work in graph mode or ms_function. - 'hook_fn' must be defined as the following code. - `cell_id` is the information of registered cell object, including name and ID. `inputs` is the forward - input objects passed to the cell. `output` is the forward output object of the cell. The 'hook_fn' can + `cell_id` is the information of registered Cell object, including name and ID. `inputs` is the forward + input objects passed to the Cell. `output` is the forward output object of the Cell. The 'hook_fn' can modify the forward output object by returning new forward output object. - It should have the following signature: hook_fn(cell_id, inputs, output) -> new output object or none. - - In order to prevent running failed when switching to graph mode, it is not recommended to write in the - construct. + - In order to prevent running failed when switching to graph mode, it is not recommended to write it in the + `construct` function of Cell object. In the pynative mode, if the `register_forward_hook` function is + called in the `construct` function of the Cell object, a hook function will be added at each run time of + Cell object. Args: hook_fn (function): Python function. Forward hook function. @@ -1755,10 +1761,10 @@ class Cell(Cell_): Backward hook construct method to replace original construct method. Args: - inputs: The input objects of cell object. + inputs: The input objects of Cell object. Returns: - - **outputs** - The output objects of cell object. + - **outputs** - The output objects of Cell object. Supported Platforms: ``Ascend`` ``GPU`` ``CPU`` @@ -1776,17 +1782,20 @@ class Cell(Cell_): def register_backward_hook(self, hook_fn): """ - Register the backward hook function. Note that this function is only supported in pynative mode. + Register the backward hook function. Note: + - The `register_backward_hook(hook_fn)` does not work in graph mode or ms_function. - The 'hook_fn' must be defined as the following code. - `cell_id` is the information of registered cell, including name and ID. `grad_input` is the gradient - passed to the cell. `grad_output` is the gradient computed and passed to the next cell or primitive, - which may be modified by returning a new output gradient. + `cell_id` is the information of registered Cell object, including name and ID. `grad_input` is the + gradient passed to the Cell. `grad_output` is the gradient computed and passed to the next Cell or + primitive, which may be modified by returning a new output gradient. - The 'hook_fn' should have the following signature: hook_fn(cell_id, grad_input, grad_output) -> New output gradient or none. - The 'hook_fn' is executed in the python environment. In order to prevent running failed when switching to - graph mode, it is not recommended to write in the construct. + graph mode, it is not recommended to write it in the `construct` function of Cell object. In the pynative + mode, if the `register_backward_hook` function is called in the `construct` function of the Cell object, + a hook function will be added at each run time of Cell object. Args: hook_fn (function): Python function. Backward hook function.