diff --git a/openi.py b/openi.py index ae46e2e..474beb5 100644 --- a/openi.py +++ b/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) \ No newline at end of file + uploadfolder_to_obs(self.train_dir,self.obs_train_url) \ No newline at end of file diff --git a/train_continue.py b/train_continue.py index c5e5346..f57fc56 100644 --- a/train_continue.py +++ b/train_continue.py @@ -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) \ No newline at end of file + uploadfile_to_obs(ckpt_url, obs_ckpt_url) \ No newline at end of file