update
This commit is contained in:
parent
5c8d47300b
commit
dcd316f7cb
10
openi.py
10
openi.py
|
|
@ -47,7 +47,7 @@ def openi_multidataset_to_env(multi_data_url, data_dir):
|
|||
print("download_multidataset_input failed")
|
||||
return
|
||||
|
||||
def openi_ckpt_to_env(obs_ckpt_url, ckpt_url):
|
||||
def uploadfile_to_obs(obs_ckpt_url, ckpt_url):
|
||||
"""
|
||||
copy ckpt file to training image or copy ckpt file to obs
|
||||
"""
|
||||
|
|
@ -80,13 +80,13 @@ def env_to_openi(train_dir, train_url):
|
|||
device_num = int(os.getenv('RANK_SIZE'))
|
||||
local_rank=int(os.getenv('RANK_ID'))
|
||||
if device_num == 1:
|
||||
upload_to_obs(train_dir, train_url)
|
||||
uploadfolder_to_obs(train_dir, train_url)
|
||||
if device_num > 1:
|
||||
if local_rank%8==0:
|
||||
upload_to_obs(train_dir, train_url)
|
||||
uploadfolder_to_obs(train_dir, train_url)
|
||||
return
|
||||
|
||||
def upload_to_obs(train_dir, obs_train_url):
|
||||
def uploadfolder_to_obs(train_dir, obs_train_url):
|
||||
"""
|
||||
upload to obs
|
||||
"""
|
||||
|
|
@ -136,4 +136,4 @@ class EnvToOpenIEpochEnd(Callback):
|
|||
self.train_dir = train_dir
|
||||
self.obs_train_url = obs_train_url
|
||||
def epoch_end(self,run_context):
|
||||
upload_to_obs(self.train_dir,self.obs_train_url)
|
||||
uploadfolder_to_obs(self.train_dir,self.obs_train_url)
|
||||
|
|
@ -21,14 +21,12 @@ from mindspore import load_checkpoint, load_param_into_net
|
|||
from mindspore.train import Model
|
||||
from mindspore.nn.metrics import Accuracy
|
||||
from mindspore.communication.management import get_rank
|
||||
import mindspore.ops as ops
|
||||
import time
|
||||
|
||||
from openi import openi_ckpt_to_env
|
||||
from openi import uploadfile_to_obs
|
||||
from openi import oponi_dataset_to_Env
|
||||
|
||||
parser = argparse.ArgumentParser(description='MindSpore Lenet Example')
|
||||
parser.add_argument('--data_url',
|
||||
parser.add_argument('--multi_data_url',
|
||||
help='path to training/inference dataset folder',
|
||||
default= '/cache/data/')
|
||||
|
||||
|
|
@ -71,13 +69,13 @@ if __name__ == "__main__":
|
|||
except Exception as e:
|
||||
print("path already exists")
|
||||
|
||||
oponi_dataset_to_Env(args.data_url, data_dir)
|
||||
oponi_dataset_to_Env(args.multi_data_url, data_dir)
|
||||
|
||||
device_num = int(os.getenv('RANK_SIZE'))
|
||||
if device_num == 1:
|
||||
ds_train = create_dataset(os.path.join(data_dir, "train"), cfg.batch_size)
|
||||
ds_train = create_dataset(os.path.join(data_dir + "/MNISTData", "train"), cfg.batch_size)
|
||||
if device_num > 1:
|
||||
ds_train = create_dataset_parallel(os.path.join(data_dir, "train"), cfg.batch_size)
|
||||
ds_train = create_dataset_parallel(os.path.join(data_dir + "/MNISTData", "train"), cfg.batch_size)
|
||||
if ds_train.get_dataset_size() == 0:
|
||||
raise ValueError("Please check dataset size > 0 and batch_size <= dataset size")
|
||||
|
||||
|
|
@ -136,4 +134,4 @@ if __name__ == "__main__":
|
|||
for n in new_models:
|
||||
ckpt_url = base_path + "/" + n
|
||||
obs_ckpt_url = args.train_url + "/" + n
|
||||
openi_ckpt_to_env(ckpt_url, obs_ckpt_url)
|
||||
uploadfile_to_obs(ckpt_url, obs_ckpt_url)
|
||||
Loading…
Reference in New Issue