mindspore/docs/api/api_python/ops/mindspore.ops.ApplyMomentum...

37 lines
2.4 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.ops.ApplyMomentum
============================
.. py:class:: mindspore.ops.ApplyMomentum(use_nesterov=False, use_locking=False, gradient_scale=1.0)
使用动量算法的优化器。
更多详细信息,请参阅论文 `On the importance of initialization and momentum in deep learning <https://dl.acm.org/doi/10.5555/3042817.3043064>`_
输入的 `variable``accumulation``gradient` 的输入遵循隐式类型转换规则,使数据类型一致。如果它们具有不同的数据类型,则低精度数据类型将转换为相对最高精度的数据类型。
有关公式和用法的更多详细信息,请参阅 :class:`mindspore.nn.Momentum`
参数:
- **use_locking** (bool, 可选) - 是否对参数更新加锁保护。默认值: ``False``
- **use_nesterov** (bool, 可选) - 是否使用nesterov动量。默认值 ``False``
- **gradient_scale** (float, 可选) - 梯度的缩放比例。默认值: ``1.0``
输入:
- **variable** (Union[Parameter, Tensor]) - 要更新的权重数据类型必须为float64、int64、float、
float16、int16、int32、int8、uint16、uint32、uint64、uint8、complex64、complex128。
- **accumulation** (Union[Parameter, Tensor]) - 按动量权重计算的累加梯度值,数据类型与 `variable` 相同。
- **learning_rate** (Union[Number, Tensor]) - 学习率必须是float64、int64、float、
float16、int16、int32、int8、uint16、uint32、uint64、uint8、complex64、complex128或为float64、int64、float、float16、int16、int32、int8、uint16、uint32、uint64、uint8、
complex64、complex128数据类型的Scalar的Tensor。
- **gradient** (Tensor) - 梯度,数据类型与 `variable` 相同。
- **momentum** (Union[Number, Tensor]) - 动量必须是float64、int64、float、float16、int16、int32、
int8、uint16、uint32、uint64、uint8、complex64、complex128类型的数值或是具有float64、int64、float、float16、int16、int32、int8、uint16、uint32、uint64、uint8、
complex64、complex128数据类型的Scalar的Tensor。
输出:
Tensor更新后的参数。
异常:
- **TypeError** - 如果 `use_locking``use_nesterov` 不是bool`gradient_scale` 不是float。
- **TypeError** - 如果 `var``accum``grad` 不支持数据类型转换。