mindspore/docs/api/api_python/ops/mindspore.ops.func_grad.rst

26 lines
1.7 KiB
ReStructuredText
Raw Normal View History

mindspore.ops.grad
==================
.. py:function:: mindspore.ops.grad(fn, grad_position=0, weights=None, has_aux=False)
生成求导函数,用于计算给定函数的梯度。
函数求导包含以下三种场景:
1. 对输入求导,此时 `grad_position` 非None`weights` 是None;
2. 对网络变量求导,此时 `grad_position` 是None`weights` 非None;
2022-08-16 15:30:10 +08:00
3. 同时对输入和网络变量求导,此时 `grad_position``weights` 都非None。
参数:
- **fn** (Union[Cell, Function]) - 待求导的函数或网络。
- **grad_position** (Union[NoneType, int, tuple[int]]) - 指定求导输入位置的索引。若为int类型表示对单个输入求导若为tuple类型表示对tuple内索引的位置求导其中索引从0开始若是None表示不对输入求导这种场景下 `weights` 非None。默认值0。
- **weights** (Union[ParameterTuple, Parameter, list[Parameter]]) - 训练网络中需要返回梯度的网络变量。一般可通过 `weights = net.trainable_params()` 获取。默认值None。
- **has_aux** (bool) - 是否返回辅助参数的标志。若为True `fn` 输出数量必须超过一个,其中只有 `fn` 第一个输出参与求导其他输出值将直接返回。默认值False。
返回:
Function用于计算给定函数的梯度的求导函数。例如 `out1, out2 = fn(*args)` ,若 `has_aux` 为True梯度函数将返回 `(gradient, out2)` 形式的结果,其中 `out2` 不参与求导若为False将直接返回 `gradient`
异常:
- **ValueError** - 入参 `grad_position``weights` 同时为None。
- **TypeError** - 入参类型不符合要求。