更新 'pretrain_for_c2net.py'

This commit is contained in:
liuzx 2023-04-14 09:19:25 +08:00
parent 5e391cb913
commit 93dd98b67a
1 changed files with 1 additions and 13 deletions

View File

@ -174,15 +174,6 @@ def DownloadModelFromQizhi(model_url, model_dir):
if local_rank%8==0:
C2netModelToEnv(model_url,model_dir)
return
def UploadToQizhi(train_dir, obs_train_url):
device_num = int(os.getenv('RANK_SIZE'))
local_rank=int(os.getenv('RANK_ID'))
if device_num == 1:
EnvToObs(train_dir, obs_train_url)
if device_num > 1:
if local_rank%8==0:
EnvToObs(train_dir, obs_train_url)
return
parser = argparse.ArgumentParser(description='MindSpore Lenet Example')
### --multi_data_url,--train_url,--device_target,These 3 parameters must be defined first in a multi-dataset,
@ -273,7 +264,4 @@ if __name__ == "__main__":
model.train(epoch_size,
ds_train,
callbacks=[time_cb, ckpoint_cb,
LossMonitor()])
###Copy the trained output data from the local running environment back to obs,
###and download it in the training task corresponding to the Qizhi platform
UploadToQizhi(train_dir,args.train_url)
LossMonitor()])