mindspore/docs/api/api_python/mint/mindspore.mint.optim.SGD.rst

46 lines
2.2 KiB
ReStructuredText
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

mindspore.mint.optim.SGD
=========================
.. py:class:: mindspore.mint.optim.SGD(params, lr, momentum=0, dampening=0, weight_decay=0, nesterov=False, *, maximize=False)
随机梯度下降算法。
.. math::
v_{t+1} = u \ast v_{t} + gradient \ast (1-dampening)
如果nesterov为True
.. math::
p_{t+1} = p_{t} - lr \ast (gradient + u \ast v_{t+1})
如果nesterov为False
.. math::
p_{t+1} = p_{t} - lr \ast v_{t+1}
需要注意的是,对于训练的第一步 :math:`v_{t+1} = gradient`。其中p、v和u分别表示 `parameters``accum``momentum`
.. warning::
这是一个实验性的优化器接口,后续可能修改或删除。需要和 `LRScheduler <https://www.mindspore.cn/docs/zh-CN/master/api_python/mindspore.experimental.html#lrscheduler%E7%B1%BB>`_ 下的动态学习率接口配合使用。
参数:
- **params** (Union[list(Parameter), list(dict)]) - 网络参数的列表或指定了参数组的列表。
- **lr** (Union[bool, int, float, Tensor]) - 学习率。
- **momentum** (Union[bool, int, float], 可选) - 动量值。默认值:``0``
- **weight_decay** (Union[bool, int, float], 可选) - 权重衰减L2 penalty必须大于等于0。默认值``0.``
- **dampening** (Union[bool, int, float], 可选) - 动量的阻尼值。默认值:``0``
- **nesterov** (bool, 可选) - 启用Nesterov动量。如果使用Nesterov动量必须为正阻尼必须等于0。默认值``False``
关键字参数:
- **maximize** (bool, 可选) - 是否根据目标函数最大化网络参数。默认值:``False``
输入:
- **gradients** (tuple[Tensor]) - 网络权重的梯度。
异常:
- **ValueError** - 学习率不是bool、int、float或Tensor。
- **ValueError** - 学习率小于0。
- **ValueError** - `momentum``weight_decay` 值小于0。
- **ValueError** - `momentum``dampening``weight_decay` 不是bool、int或float。
- **ValueError** - `nesterov``maximize` 不是bool类型。
- **ValueError** - `nesterov` 为True时 `momentum` 不为正或 `dampening` 不为0。