From 7aab8b05c69f29f39caebf3d903feeefe3eec28e Mon Sep 17 00:00:00 2001 From: liuzx Date: Mon, 12 Jun 2023 10:55:50 +0800 Subject: [PATCH] =?UTF-8?q?=E6=9B=B4=E6=96=B0=20'train=5Ffor=5Fc2net.py'?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- train_for_c2net.py | 13 ++++++++----- 1 file changed, 8 insertions(+), 5 deletions(-) diff --git a/train_for_c2net.py b/train_for_c2net.py index a7590e5..fe5cf58 100644 --- a/train_for_c2net.py +++ b/train_for_c2net.py @@ -112,10 +112,13 @@ if __name__ == "__main__": if (args.epoch_size): epoch_size = args.epoch_size print('epoch_size is: ', epoch_size) + model.train(epoch_size, ds_train,callbacks=[time_cb, ckpoint_cb,LossMonitor()]) # set callback functions - callback =[time_cb,LossMonitor()] - local_rank=int(os.getenv('RANK_ID')) + # callback =[time_cb,LossMonitor()] + # local_rank=int(os.getenv('RANK_ID')) # for data parallel, only save checkpoint on rank 0 - if local_rank==0 : - callback.append(ckpoint_cb) - model.train(epoch_size,ds_train,callbacks=callback) \ No newline at end of file + # if local_rank==0 : + # callback.append(ckpoint_cb) + # model.train(epoch_size, + # ds_train, + # callbacks=callback) \ No newline at end of file