mindspore/docs/api/api_python/train/mindspore.train.RunContext.rst

82 lines
5.8 KiB
ReStructuredText
Raw Permalink Normal View History

mindspore.train.RunContext
==========================
2022-05-10 11:59:45 +08:00
.. py:class:: mindspore.train.RunContext(original_args)
2021-12-02 09:29:36 +08:00
2022-05-13 17:23:34 +08:00
保存和管理模型的相关信息。
`RunContext` 主要用于收集训练或推理过程中模型的上下文相关信息并作为入参传入callback对象中来实现信息的共享。
Callback的类方法中调用 `RunContext.original_args()` 可以获取模型当前的上下文信息,用户也可以为此信息添加额外的自定义属性,同时 `request_stop()` 方法可以控制训练过程的停止。具体用法请查看 `Callback <https://www.mindspore.cn/tutorials/experts/zh-CN/master/debug/custom_debug.html>`_
`RunContext.original_args()` 存储的模型信息为一个字典型变量,在训练和推理过程会存储不同的属性。详情如下:
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| 训练过程支持的属性 | 推理过程支持的属性 | 说明 |
2022-05-13 17:23:34 +08:00
+==========================+========================+=======================================+
2022-05-16 16:55:57 +08:00
| train_network | | 包含了优化器和损失的训练网络 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| epoch_num | | 训练的epoch数 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| train_dataset | | 训练集 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| loss_fn | | 损失函数 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| optimizer | | 优化器 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| parallel_mode | | 并行模式 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| device_number | | 设备编号 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| train_dataset_element | | 当前step的训练数据 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| last_save_ckpt_step | | 最后一次存储ckpt的step |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| latest_ckpt_file | | ckpt文件名 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| cur_epoch_num | | 当前的epoch |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| | eval_network | 评估网络 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| | valid_dataset | 验证集 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| | metrics | 评估指标 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| mode | mode | "train"或"eval"模式 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| batch_num | batch_num | 训练或推理的batch数 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| list_callback | list_callback | 回调列表 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| network | network | 基础的网络结构 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| cur_step_num | cur_step_num | 当前的训练或推理的step |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| dataset_sink_mode | dataset_sink_mode | 训练或推理的数据是否下沉 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2022-05-16 16:55:57 +08:00
| net_outputs | net_outputs | 训练或推理的网络输出 |
2022-05-13 17:23:34 +08:00
+--------------------------+------------------------+---------------------------------------+
2021-12-02 09:29:36 +08:00
2022-07-22 10:22:12 +08:00
参数:
- **original_args** (dict) - 模型的相关信息。
2021-12-04 20:36:47 +08:00
2021-12-02 09:29:36 +08:00
.. py:method:: get_stop_requested()
2021-12-28 20:07:42 +08:00
获取是否停止训练的标志。
2021-12-02 09:29:36 +08:00
2022-07-22 10:22:12 +08:00
返回:
bool如果为True`Model.train()` 停止迭代。
2021-12-04 20:36:47 +08:00
2021-12-02 09:29:36 +08:00
.. py:method:: original_args()
2021-12-28 20:07:42 +08:00
获取模型相关信息的对象。
2021-12-02 09:29:36 +08:00
2022-07-22 10:22:12 +08:00
返回:
dict含有模型的相关信息的对象。
2021-12-04 20:36:47 +08:00
2021-12-02 09:29:36 +08:00
.. py:method:: request_stop()
在训练期间设置停止请求。
可以使用此函数请求停止训练。 `Model.train()` 会检查是否调用此函数。