This commit is contained in:
liuzx 2023-05-10 16:12:36 +08:00
parent 5c8d47300b
commit dcd316f7cb
2 changed files with 11 additions and 13 deletions

View File

@ -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)

View File

@ -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)