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

35 lines
2.2 KiB
ReStructuredText
Raw Normal View History

2022-10-17 16:27:44 +08:00
mindspore.nn.CTCLoss
====================
.. py:class:: mindspore.nn.CTCLoss(blank=0, reduction='mean', zero_infinity=False)
CTCLoss损失函数。
关于CTCLoss算法详细介绍请参考 `Connectionist Temporal Classification: Labeling Unsegmented Sequence Data withRecurrent Neural Networks <http://www.cs.toronto.edu/~graves/icml_2006.pdf>`_
参数:
- **blank** (int) - 空白标签。默认值0。
2023-02-06 16:55:54 +08:00
- **reduction** (str) - 对输出应用归约方法。可选值为"none"、"mean"或"sum"。默认值:"mean"。
2022-10-25 15:33:13 +08:00
- **zero_infinity** (bool) - 是否设置无限损失和相关梯度为零。默认值False。
2022-10-17 16:27:44 +08:00
输入:
2022-10-25 15:33:13 +08:00
- **log_probs** (Tensor) - 输入Tensorshape :math:`(T, N, C)`:math:`(T, C)` 。其中T表示输入长度N表示批次大小C是分类数。TNC均为正整数。
2022-10-26 16:33:58 +08:00
- **targets** (Tensor) - 目标Tensorshape :math:`(N, S)` 或 (sum( `target_lengths` ))。其中S表示最大目标长度。
2022-12-13 15:58:00 +08:00
- **input_lengths** (Union[tuple, Tensor]) - shape为N的Tensor或tuple。表示输入长度。
- **target_lengths** (Union[tuple, Tensor]) - shape为N的Tensor或tuple。表示目标长度。
2022-10-17 16:27:44 +08:00
输出:
- **neg_log_likelihood** (Tensor) - 对每一个输入节点可微调的损失值。
异常:
2022-12-13 15:58:00 +08:00
- **TypeError** - `log_probs``targets` 不是Tensor。
2022-10-17 16:27:44 +08:00
- **TypeError** - `zero_infinity` 不是布尔值, `reduction` 不是字符串。
2022-10-25 15:33:13 +08:00
- **TypeError** - `log_probs` 的数据类型不是float或double。
2022-10-17 16:27:44 +08:00
- **TypeError** - `targets` `input_lengths``target_lengths` 数据类型不是int32或int64。
- **ValueError** - `reduction` 不为"none""mean"或"sum"。
- **ValueError** - `targets` `input_lengths``target_lengths` 的数据类型是不同的。
2022-10-25 15:33:13 +08:00
- **ValueError** - `blank` 值不介于0到C之间。C是 `log_probs` 的分类数。
2023-02-10 16:27:40 +08:00
- **ValueError** - 当 `log_prob` 的shape是 :math:`(T, C)` 时, `target` 的维度不等于1。
2022-12-24 10:09:59 +08:00
- **RuntimeError** - `input_lengths` 的值大于T。T是 `log_probs` 的长度。
- **RuntimeError** - `target_lengths[i]` 的值不介于0到 `input_length[i]` 之间。