自定义Callback重载函数调用顺序错误
收藏回复举报
自定义Callback重载函数调用顺序错误
发表于2023-09-14 17:27:51
0 查看

1 系统环境

硬件环境(Ascend/GPU/CPU): Ascend/GPU/CPU

MindSpore版本: mindspore=1.10.1

执行模式(PyNative/ Graph):不限

Python版本: Python=3.9.7

操作系统平台: 不限

2 报错信息

2.1 问题描述

自定义Callback的函数调用顺序违背方法的名称。现有调用顺序为on_train_epoch_begin -> on_train_step_* -> on_eval_epoch_begin -> on_eval_step_* -> on_eval_epoch_end -> on_train_epoch_end
这导致train_epoch内含了eval_epoch。这样的运行流程导致设计callback的回显内容时增加了不必要的麻烦,并且逻辑上也不直观。

当前结构:

  • on_train_epoch_begin
    • on_train_step_*
    • on_eval_epoch_begin
      • on_eval_step_*
    • on-eval_epoch_end
  • on_train_epoch_end

期望结构:

  • on_train_epoch_begin
    • on_train_step_*
  • on_train_epoch_end
  • on_eval_epoch_begin
    • on_eval_step_*
  • on-eval_epoch_end

2.2 脚本代码(代码格式,可上传附件)

import mindspore as ms
import mindspore.nn as nn
import mindspore.ops as ops

import numpy as np

from typing import *

# ms.set_context(mode=ms.GRAPH_MODE, device_target="CPU")
ms.set_context(mode=ms.PYNATIVE_MODE, device_target="CPU")

class Net(nn.Cell):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 1, 3, 1)

    def construct(self, x):
        return self.conv1(x)



class Callback1(ms.train.callback.Callback):
    def __init__(self):
        super().__init__()

    def on_train_epoch_begin(self, run_context):
        print("on_train_epoch_begin")

    def on_train_step_end(self, run_context):
        print("🍓", end='')

    def on_train_epoch_end(self, run_context, *args):
        print()
        print("on_train_epoch_end")

    def on_eval_epoch_begin(self, run_context):
        print("on_eval_epoch_begin")

    def on_eval_step_end(self, run_context):
        print("🧃", end='')

    def on_eval_epoch_end(self, run_context, *args):
        print()
        print("on_eval_epoch_end")



if __name__ == "__main__":
    dataset = ms.dataset.NumpySlicesDataset(
        (np.random.randn(100, 1, 16, 16).astype(np.float32), np.random.randn(100, 1, 16, 16).astype(np.float32)),
        column_names=["image", "label"], shuffle=False)
    train_ds, val_ds = dataset.split([0.5, 0.5], randomize=False)

    train_ds = train_ds.batch(4)
    val_ds = val_ds.batch(4)

    net = Net()
    loss = nn.MSELoss()
    opt = nn.optim.Adam(net.trainable_params(), 0.001)
    callback1 = Callback1()

    model = ms.Model(
        network=net,
        loss_fn=loss,
        optimizer=opt,
        metrics={"MSE": nn.metrics.MSE()},
    )

    model.fit(
        epoch=4,
        train_dataset=train_ds,
        valid_dataset=val_ds,
        callbacks=[callback1],
        dataset_sink_mode=False,
        valid_dataset_sink_mode=False,
        sink_size=-1,
        initial_epoch=0
    )

3 根因分析

cke_7139.png

4 解决方案

******此处由用户填写******

包含文字方案和最终脚本代码

请将正确的脚本打包并上传附件

我要发帖子