更新 'pretrain_for_c2net.py'
This commit is contained in:
parent
5e391cb913
commit
93dd98b67a
|
|
@ -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()])
|
||||
Loading…
Reference in New Issue