mindspore/docs/api/api_python/nn/mindspore.nn.TrainOneStepCe...

24 lines
904 B
ReStructuredText
Raw Normal View History

2021-12-04 15:18:50 +08:00
mindspore.nn.TrainOneStepCell
=============================
.. py:class:: mindspore.nn.TrainOneStepCell(network, optimizer, sens=1.0)
训练网络封装类。
2022-08-24 17:45:11 +08:00
封装 `network``optimizer` 。构建一个输入'\*inputs'的用于训练的Cell。
2021-12-04 15:18:50 +08:00
执行函数 `construct` 中会构建反向图以更新网络参数。支持不同的并行训练模式。
参数:
- **network** (Cell) - 训练网络。只支持单输出网络。
- **optimizer** (Union[Cell]) - 用于更新网络参数的优化器。
- **sens** (numbers.Number) - 反向传播的输入缩放系数。默认值为1.0。
2021-12-04 15:18:50 +08:00
输入:
2023-02-23 10:04:43 +08:00
- **\*inputs** (Tuple(Tensor)) - shape为 :math:`(N, \ldots)` 的Tensor组成的元组。
2021-12-04 15:18:50 +08:00
输出:
Tensor损失函数值其shape通常为 :math:`()`
2021-12-04 15:18:50 +08:00
异常:
- **TypeError** - `sens` 不是numbers.Number。