mindspore/docs/api/api_python/amp/mindspore.amp.LossScaleMana...

26 lines
1.1 KiB
ReStructuredText
Raw Normal View History

2022-08-30 16:52:05 +08:00
mindspore.amp.LossScaleManager
==============================
.. py:class:: mindspore.amp.LossScaleManager
使用混合精度时用于管理损失缩放系数loss scale的抽象类。
2022-09-14 15:17:50 +08:00
派生类需要实现该类的所有方法。 `get_loss_scale` 用于获取当前的梯度放大系数。 `update_loss_scale` 用于更新梯度放大系数,该方法将在训练过程中被调用。 `get_update_cell` 用于获取更新梯度放大系数的 :class:`mindspore.nn.Cell` 实例,该实例将在训练过程中被调用。当前多使用 `get_update_cell` 方式。
2022-08-30 16:52:05 +08:00
例如::class:`mindspore.amp.FixedLossScaleManager`:class:`mindspore.amp.DynamicLossScaleManager`
.. py:method:: get_loss_scale()
获取梯度放大系数loss scale的值。
.. py:method:: get_update_cell()
2022-09-02 16:46:41 +08:00
获取用于更新梯度放大系数的 :class:`mindspore.nn.Cell` 实例。
2022-08-30 16:52:05 +08:00
.. py:method:: update_loss_scale(overflow)
2022-09-06 17:48:14 +08:00
根据 `overflow` 状态更新梯度放大系数loss scale
2022-08-30 16:52:05 +08:00
参数:
- **overflow** (bool) - 表示训练过程是否溢出。