mindspore/docs/api/api_python/nn/mindspore.nn.Loss.rst

50 lines
1.1 KiB
ReStructuredText
Raw Normal View History

2021-12-04 13:55:42 +08:00
mindspore.nn.Loss
=================
.. py:class:: mindspore.nn.Loss
计算loss的平均值。如果每 :math:`n` 次迭代调用一次 `update` 方法,则评估结果为:
.. math::
loss = \frac{\sum_{k=1}^{n}loss_k}{n}
**样例:**
>>> import numpy as np
>>> from mindspore import nn, Tensor
>>>
>>> x = Tensor(np.array(0.2), mindspore.float32)
>>> loss = nn.Loss()
>>> loss.clear()
>>> loss.update(x)
>>> result = loss.eval()
.. py:method:: clear()
内部评估结果清零。
.. py:method:: eval()
计算loss的平均值。
**返回:**
2021-12-04 20:36:47 +08:00
2021-12-04 13:55:42 +08:00
Floatloss的平均值。
2021-12-04 18:37:47 +08:00
**异常:**
RuntimeError样本总数为0。
2021-12-04 13:55:42 +08:00
.. py:method:: update(*inputs)
更新内部评估结果。
2021-12-04 18:37:47 +08:00
**参数:**
2021-12-04 20:36:47 +08:00
- **inputs** - 输入只包含一个元素且该元素为loss。loss的维度必须为0或1。
2021-12-04 18:37:47 +08:00
**异常:**
2021-12-04 13:55:42 +08:00
2021-12-04 20:36:47 +08:00
- **ValueError** - `inputs` 的长度不为1。
- **ValueError** - `inputs` 的维度不为0或1。