mindspore2022/docs/api/api_python/dataset/mindspore.dataset.WaitedDSC...

37 lines
2.1 KiB
ReStructuredText
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

mindspore.dataset.WaitedDSCallback
==================================
.. py:class:: mindspore.dataset.WaitedDSCallback(step_size=1)
阻塞式数据处理回调类的抽象基类,用于与训练回调类 `mindspore.train.callback <https://mindspore.cn/docs/api/zh-CN/master/api_python/mindspore.train.html#mindspore.train.callback.Callback>`_ 的同步。
可用于在step或epoch开始前执行自定义的回调方法例如在自动数据增强中根据上一个epoch的loss值来更新增强算子参数配置。
用户可通过 `train_run_context` 获取网络训练相关信息,如 `network` 、 `train_network` 、 `epoch_num` 、 `batch_num` 、 `loss_fn` 、 `optimizer` 、 `parallel_mode` 、 `device_number` 、 `list_callback` 、 `cur_epoch_num` 、 `cur_step_num` 、 `dataset_sink_mode` 、 `net_outputs` 等,详见 `mindspore.train.callback <https://mindspore.cn/docs/api/zh-CN/master/api_python/mindspore.train.html#mindspore.train.callback.Callback>`_
用户可通过 `ds_run_context` 获取数据处理管道相关信息,包括 `cur_epoch_num` (当前epoch数)、 `cur_step_num_in_epoch` (当前epoch的step数)、 `cur_step_num` (当前step数)。
.. note:: 注意第2个step或epoch开始时才会触发该调用。
**参数:**
- **step_size** (int, optional) - 每个step包含的数据行数。通常step_size与batch_size一致默认值1。
.. py:method:: sync_epoch_begin(train_run_context, ds_run_context)
用于定义在数据epoch开始前训练epoch结束后执行的回调方法。
**参数:**
- **train_run_context**包含前一个epoch的反馈信息的网络训练运行信息。
- **ds_run_context**:数据处理管道运行信息。
.. py:method:: sync_step_begin(train_run_context, ds_run_context)
用于定义在数据step开始前训练step结束后执行的回调方法。
**参数:**
- **train_run_context**包含前一个step的反馈信息的网络训练运行信息。
- **ds_run_context**:数据处理管道运行信息。