From b39f02ed2d915764e78af6cd8a5008c830b7d74d Mon Sep 17 00:00:00 2001 From: yao_yf Date: Thu, 29 Jul 2021 15:06:52 +0800 Subject: [PATCH] pangu train and eval --- .../insert_tensor_move_for_hccl_op.cc | 7 +- .../parallel/ops_info/virtual_output_info.cc | 10 +- .../ccsrc/frontend/parallel/step_parallel.cc | 48 ++++-- .../ccsrc/frontend/parallel/step_parallel.h | 9 ++ mindspore/train/model.py | 6 + .../scripts/run_distribute_train_and_eval.sh | 53 +++++++ .../official/nlp/pangu_alpha/src/callbacks.py | 108 +++++++++++++ .../official/nlp/pangu_alpha/src/dataset.py | 5 +- .../official/nlp/pangu_alpha/src/metrics.py | 58 +++++++ .../official/nlp/pangu_alpha/src/utils.py | 12 ++ model_zoo/official/nlp/pangu_alpha/train.py | 146 +++++++----------- 11 files changed, 346 insertions(+), 116 deletions(-) create mode 100644 model_zoo/official/nlp/pangu_alpha/scripts/run_distribute_train_and_eval.sh create mode 100644 model_zoo/official/nlp/pangu_alpha/src/callbacks.py create mode 100644 model_zoo/official/nlp/pangu_alpha/src/metrics.py diff --git a/mindspore/ccsrc/backend/optimizer/ascend/enhancer/insert_tensor_move_for_hccl_op.cc b/mindspore/ccsrc/backend/optimizer/ascend/enhancer/insert_tensor_move_for_hccl_op.cc index af65748a45e..07957ee3334 100644 --- a/mindspore/ccsrc/backend/optimizer/ascend/enhancer/insert_tensor_move_for_hccl_op.cc +++ b/mindspore/ccsrc/backend/optimizer/ascend/enhancer/insert_tensor_move_for_hccl_op.cc @@ -27,7 +27,7 @@ namespace opt { namespace { // insert tensormove for some cnode even if not a Ref cnode const std::set kNeedInsertTensorMoveOpSet = {kLambNextMVOpName, kLambNextMVWithDecayOpName, - kLambUpdateWithLROpName}; + kLambUpdateWithLROpName, kGetNextOpName}; bool IsParameterOrValueNode(const AnfNodePtr &node) { MS_EXCEPTION_IF_NULL(node); @@ -86,9 +86,10 @@ bool InsertTensorMoveForHcclOp::NeedInsertTensorMove(const FuncGraphPtr &graph, if (kernel_query_->IsTbeRef(input)) { return true; } - + auto kernel_with_index = AnfAlgo::VisitKernelWithReturnType(input, 0, true); + auto real_node = kernel_with_index.first; // when input is some special cnodes - if (kNeedInsertTensorMoveOpSet.find(AnfAlgo::GetCNodeName(input)) != kNeedInsertTensorMoveOpSet.end()) { + if (kNeedInsertTensorMoveOpSet.find(AnfAlgo::GetCNodeName(real_node)) != kNeedInsertTensorMoveOpSet.end()) { return true; } diff --git a/mindspore/ccsrc/frontend/parallel/ops_info/virtual_output_info.cc b/mindspore/ccsrc/frontend/parallel/ops_info/virtual_output_info.cc index ae6411f8f35..712d44e509e 100644 --- a/mindspore/ccsrc/frontend/parallel/ops_info/virtual_output_info.cc +++ b/mindspore/ccsrc/frontend/parallel/ops_info/virtual_output_info.cc @@ -64,8 +64,14 @@ Status VirtualOutputInfo::GenerateStrategies(int64_t stage_id) { } for (auto &shape : inputs_shape_) { Shape temp; - temp.emplace_back(SizeToLong(total_dev_num)); - (void)temp.insert(temp.end(), shape.size() - 1, 1); + if (!shape.empty()) { + if (shape[0] % total_dev_num == 0) { + temp.emplace_back(SizeToLong(total_dev_num)); + } else { + temp.emplace_back(1); + } + (void)temp.insert(temp.end(), shape.size() - 1, 1); + } strategy.push_back(temp); } sp = std::make_shared(stage_id, strategy); diff --git a/mindspore/ccsrc/frontend/parallel/step_parallel.cc b/mindspore/ccsrc/frontend/parallel/step_parallel.cc index 357b115a871..043f8dd9833 100644 --- a/mindspore/ccsrc/frontend/parallel/step_parallel.cc +++ b/mindspore/ccsrc/frontend/parallel/step_parallel.cc @@ -2038,7 +2038,12 @@ void SetVirtualDatasetStrategy(const CNodePtr &node) { if (shape_list[0][i].empty()) { MS_LOG(EXCEPTION) << "shape_list[ " << i << " ].size() is zero"; } - Dimensions input_strategy = {dev_num}; + Dimensions input_strategy; + if (!shape_list[0][i].empty() && shape_list[0][i][0] % dev_num == 0) { + input_strategy.push_back(dev_num); + } else if (!shape_list[0][i].empty()) { + input_strategy.push_back(1); + } for (size_t j = 1; j < shape_list[0][i].size(); j++) { input_strategy.push_back(1); } @@ -3222,12 +3227,9 @@ void MarkForwardCNode(const FuncGraphPtr &root) { } } -Status ParallelInit() { - MS_EXCEPTION_IF_NULL(ParallelContext::GetInstance()); +CommInfo GetCommInfo() { int64_t device_num = ParallelContext::GetInstance()->device_num(); int64_t global_rank = ParallelContext::GetInstance()->global_rank(); - int32_t split_stage_num = ParallelContext::GetInstance()->pipeline_stage_split_num(); - std::string parallel_mode = ParallelContext::GetInstance()->parallel_mode(); auto ms_context = MsContext::GetInstance(); MS_EXCEPTION_IF_NULL(ms_context); std::string backend = ms_context->get_param(MS_CTX_DEVICE_TARGET); @@ -3240,15 +3242,8 @@ Status ParallelInit() { world_group = NCCL_WORLD_GROUP; communication_backend = NCCL_BACKEND; } else { - MS_LOG(ERROR) << "Invalid communication backend: " << backend; - return FAILED; + MS_LOG(EXCEPTION) << "Invalid communication backend: " << backend; } - - if (split_stage_num <= 0) { - MS_LOG(ERROR) << "Invalid stage num " << split_stage_num << ", expected a positive stage number"; - return FAILED; - } - uint32_t world_rank_size = 0; if (!ParallelContext::GetInstance()->device_num_is_set()) { if (!CommManager::GetInstance().GetRankSize(world_group, &world_rank_size)) { @@ -3266,7 +3261,21 @@ Status ParallelInit() { global_rank = UintToInt(rank_id); MS_LOG(INFO) << "Get global rank from communication model, the global rank is " << global_rank; } + CommInfo comm_info{device_num, global_rank, world_group, communication_backend}; + return comm_info; +} +Status ParallelInit() { + MS_EXCEPTION_IF_NULL(ParallelContext::GetInstance()); + int32_t split_stage_num = ParallelContext::GetInstance()->pipeline_stage_split_num(); + std::string parallel_mode = ParallelContext::GetInstance()->parallel_mode(); + if (split_stage_num <= 0) { + MS_LOG(ERROR) << "Invalid stage num " << split_stage_num << ", expected a positive stage number"; + return FAILED; + } + auto comm_info = GetCommInfo(); + int64_t device_num = comm_info.device_num; + int64_t global_rank = comm_info.global_rank; if ((device_num <= 0) || (device_num > MAX_DEVICE_NUM)) { MS_LOG(ERROR) << "Invalid device num " << device_num; return FAILED; @@ -3293,13 +3302,14 @@ Status ParallelInit() { return FAILED; } - if (!InitDevice(device_num, global_rank, communication_backend, stages)) { + if (!InitDevice(device_num, global_rank, comm_info.communication_backend, stages)) { MS_LOG(ERROR) << "Init device failed"; return FAILED; } MS_LOG(INFO) << "The parallel context: dev num: " << device_num << ", global rank: " << global_rank - << ", backend: " << backend << ", gradients_mean: " << ParallelContext::GetInstance()->gradients_mean() + << ", communication_backend: " << comm_info.communication_backend + << ", gradients_mean: " << ParallelContext::GetInstance()->gradients_mean() << ", gradient_fp32_sync: " << ParallelContext::GetInstance()->gradient_fp32_sync(); return SUCCESS; @@ -3714,7 +3724,13 @@ void ReorderForPipelineSplit(const FuncGraphPtr &root, const FuncGraphManagerPtr bool IsInsertVirtualOutput(const FuncGraphPtr &root) { MS_EXCEPTION_IF_NULL(ParallelContext::GetInstance()); - return (!root->has_flag(TRAINING) && ParallelContext::GetInstance()->dataset_strategy().empty()); + auto comm_info = GetCommInfo(); + int32_t split_stage_num = ParallelContext::GetInstance()->pipeline_stage_split_num(); + int32_t per_stage_device_num = comm_info.device_num / split_stage_num; + int32_t current_stage = comm_info.global_rank / per_stage_device_num; + MS_LOG(INFO) << "The current stage is: " << current_stage; + return (!root->has_flag(TRAINING) && ParallelContext::GetInstance()->dataset_strategy().empty() && + current_stage == split_stage_num - 1); } bool StepParallel(const FuncGraphPtr &root, const opt::OptimizerPtr &optimizer) { diff --git a/mindspore/ccsrc/frontend/parallel/step_parallel.h b/mindspore/ccsrc/frontend/parallel/step_parallel.h index 71c69705080..996cc11ba33 100644 --- a/mindspore/ccsrc/frontend/parallel/step_parallel.h +++ b/mindspore/ccsrc/frontend/parallel/step_parallel.h @@ -47,6 +47,13 @@ struct LossNodeInfo { CNodePtr loss_node = nullptr; }; +struct CommInfo { + int64_t device_num = 1; + int64_t global_rank = 0; + std::string world_group; + std::string communication_backend; +}; + struct ParameterSliceInfo { Shape slice_shape; RankList group_ranks; @@ -178,6 +185,8 @@ void InsertVirtualOutput(const FuncGraphPtr &root, const std::vector std::string MirrorOpName(); +CommInfo GetCommInfo(); + void ReorderForPipelineSplit(const FuncGraphPtr &root, const FuncGraphManagerPtr &manager, int64_t pipeline_stages); } // namespace parallel } // namespace mindspore diff --git a/mindspore/train/model.py b/mindspore/train/model.py index 416d8707d86..23412cd1f5d 100644 --- a/mindspore/train/model.py +++ b/mindspore/train/model.py @@ -274,6 +274,8 @@ class Model: def _update_metrics(self, outputs): """Update metrics local values.""" + if isinstance(outputs, Tensor): + outputs = (outputs,) if not isinstance(outputs, tuple): raise ValueError("The `outputs` is not tuple.") @@ -365,6 +367,8 @@ class Model: dataset_sink_mode=True, sink_size=sink_size) self._train_network = train_network + if context.get_auto_parallel_context("pipeline_stages") > 1 and valid_dataset: + self._train_network.add_flags_recursive(is_first_iteration=True) for inputs in train_dataset_helper: self._train_network.compile(*inputs) break @@ -378,6 +382,8 @@ class Model: dataset=valid_dataset, dataset_sink_mode=True) self._eval_network = eval_network + if context.get_auto_parallel_context("pipeline_stages") > 1: + self._eval_network.add_flags_recursive(is_first_iteration=False) for inputs in valid_dataset_helper: self._eval_network.compile(*inputs) break diff --git a/model_zoo/official/nlp/pangu_alpha/scripts/run_distribute_train_and_eval.sh b/model_zoo/official/nlp/pangu_alpha/scripts/run_distribute_train_and_eval.sh new file mode 100644 index 00000000000..50dd4f01aed --- /dev/null +++ b/model_zoo/official/nlp/pangu_alpha/scripts/run_distribute_train_and_eval.sh @@ -0,0 +1,53 @@ +#!/bin/bash +# Copyright 2020 Huawei Technologies Co., Ltd +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================ + +echo "==============================================================================================================" +echo "Please run the script as: " +echo "bash run_distributed_train.sh DATA_DIR RANK_TABLE_FILE DEVICE_NUM TYPE MODE STAGE_NUM MICRO_SIZE" +echo "PER_BATCH RANK_START RANK_START LOCAL_DEVICE_NUM" +echo "for example:" +echo "#######no pipeline#######" +echo "bash run_distributed_train.sh /path/dataset /path/eval_dataset /path/hccl.json 8 fp32 2.6B 1 1 16 0 8" +echo "#######pipeline#######" +echo "bash run_distributed_train.sh /path/dataset /path/eval_dataset /path/hccl.json 16 fp32 2.6B 2 4 16 0 8" +echo "bash run_distributed_train.sh /path/dataset /path/eval_dataset /path/hccl.json 16 fp32 2.6B 2 4 16 8 8" +echo "It is better to use absolute path." +echo "==============================================================================================================" + +ROOT_PATH=`pwd` +DATA_DIR=$1 +EVAL_DATA_DIR=$2 +export RANK_TABLE_FILE=$3 +RANK_SIZE=$4 +PARAM_INIT_TYPE=$5 +MODE=$6 +STAGE_NUM=$7 +MICRO_SIZE=$8 +PER_BATCH=$9 +RANK_START=${10} +LOCAL_DEVICE_NUM=${11} + +for((i=0;i<${LOCAL_DEVICE_NUM};i++)); +do + rm ${ROOT_PATH}/device$i/ -rf + mkdir ${ROOT_PATH}/device$i + cd ${ROOT_PATH}/device$i || exit + export RANK_ID=$[i+RANK_START] + export DEVICE_ID=$i + python ${ROOT_PATH}/train.py --distribute=true --device_num=$RANK_SIZE --data_url=$DATA_DIR --run_type=train \ + --param_init_type=$PARAM_INIT_TYPE --mode=$MODE --stage_num=$STAGE_NUM --micro_size=$MICRO_SIZE \ + --per_batch_size=$PER_BATCH --train_and_eval_mode=1 --eval_data_url=$EVAL_DATA_DIR > log$i.log 2>&1 & +done diff --git a/model_zoo/official/nlp/pangu_alpha/src/callbacks.py b/model_zoo/official/nlp/pangu_alpha/src/callbacks.py new file mode 100644 index 00000000000..24e7ca008ff --- /dev/null +++ b/model_zoo/official/nlp/pangu_alpha/src/callbacks.py @@ -0,0 +1,108 @@ +# Copyright 2021 Huawei Technologies Co., Ltd +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================ +""" +Callbacks +""" + + +import time +import math +from mindspore.train.callback import Callback +from mindspore import context +from mindspore.context import ParallelMode +from mindspore.communication.management import get_rank + +class LossCallBack(Callback): + """ + Monitor the loss in training. + If the loss in NAN or INF terminating training. + """ + + def __init__(self, dataset_size=-1, local_rank=0, has_trained_epoch=0, has_trained_step=0, micro_size=1): + super(LossCallBack, self).__init__() + self._dataset_size = dataset_size + self.local_rank = local_rank + self.has_trained_epoch = has_trained_epoch + self.has_trained_step = has_trained_step + self.micro_size = micro_size + print("load has trained epoch :{} and step: {}".format(has_trained_epoch, has_trained_step), flush=True) + + def step_end(self, run_context): + """ + Print loss after each step + """ + cb_params = run_context.original_args() + if self._dataset_size > 0 and self.local_rank % 8 == 0: + percent, epoch_num = math.modf(cb_params.cur_step_num / + self._dataset_size) + if percent == 0: + epoch_num -= 1 + date = time.asctime(time.localtime(time.time())) + loss_value = cb_params.net_outputs[0].asnumpy() / self.micro_size + print("time: {} local_rank: {}, epoch: {}, step: {}, output is {}, overflow is {}, scale is {}". + format(date, int(self.local_rank), int(epoch_num) + int(self.has_trained_epoch), + cb_params.cur_step_num + int(self.has_trained_step), loss_value, + cb_params.net_outputs[1].asnumpy(), cb_params.net_outputs[2].asnumpy())) + +class EvalCallBack(Callback): + """ + Monitor the ppl loss in evaluating. + Note: + If per_print_times is 0, do NOT print loss. + + Args: + print_per_step (int): Print loss every times. Default: 1. + """ + def __init__(self, model, eval_dataset, ppl_metric, print_per_step=100, has_trained_step=0): + super(EvalCallBack, self).__init__() + if not isinstance(print_per_step, int) or print_per_step < 0: + raise ValueError("print_per_step must be int and >= 0.") + self.print_per_step = print_per_step + self.model = model + self.eval_dataset = eval_dataset + self.pplMetric = ppl_metric + self.has_trained_step = has_trained_step + self.pplMetric.clear() + self.parallel_mode = context.get_auto_parallel_context("parallel_mode") + self.strategy_ckpt_save_file = context.get_auto_parallel_context("strategy_ckpt_save_file") + self.strategy_ckpt_load_file = context.get_auto_parallel_context("strategy_ckpt_load_file") + + def step_end(self, run_context): + """ + step end + """ + cb_params = run_context.original_args() + current_step = cb_params.cur_step_num + self.has_trained_step + if current_step % self.print_per_step != 0: + return + self.pplMetric.clear() + if self.parallel_mode in (ParallelMode.SEMI_AUTO_PARALLEL, ParallelMode.AUTO_PARALLEL): + context.set_auto_parallel_context(strategy_ckpt_save_file="", + strategy_ckpt_load_file=self.strategy_ckpt_save_file) + rank_id = 0 + if self.parallel_mode in (ParallelMode.SEMI_AUTO_PARALLEL, + ParallelMode.AUTO_PARALLEL, ParallelMode.DATA_PARALLEL): + rank_id = get_rank() + start_time = time.time() + out = self.model.eval(self.eval_dataset, dataset_sink_mode=True) + end_time = time.time() + eval_time = int(end_time - start_time) + + time_str = time.strftime("%Y-%m-%d %H:%M%S", time.localtime()) + out_str = "{} == Rank: {} == EvalCallBack model.eval(): {}; eval_time: {}s". \ + format(time_str, rank_id, out.values(), eval_time) + print(out_str) + context.set_auto_parallel_context(strategy_ckpt_save_file=self.strategy_ckpt_save_file, + strategy_ckpt_load_file=self.strategy_ckpt_load_file) diff --git a/model_zoo/official/nlp/pangu_alpha/src/dataset.py b/model_zoo/official/nlp/pangu_alpha/src/dataset.py index b8966d870c4..1ebafc072fd 100644 --- a/model_zoo/official/nlp/pangu_alpha/src/dataset.py +++ b/model_zoo/official/nlp/pangu_alpha/src/dataset.py @@ -67,7 +67,7 @@ def get_input_data_batch_slice_map(input_ids, eod_id, rank, dis, eod_reset): def create_dataset(batch_size, data_path, device_num=1, rank=0, drop=True, full_batch=False, data_start_index=0, - eod_reset=False, eod_id=9, column_name='input_ids', epoch=1): + eod_reset=False, eod_id=9, column_name='input_ids', epoch=1, num_samples=None): """ Create dataset @@ -99,7 +99,8 @@ def create_dataset(batch_size, data_path, device_num=1, rank=0, drop=True, full_ data.sort() # Load data files and preprocess - dataset = ds.MindDataset(data[data_start_index:], columns_list=[column_name], shuffle=False) + dataset = ds.MindDataset(data[data_start_index:], columns_list=[column_name], + shuffle=False, num_samples=num_samples) type_cast_op = C.TypeCast(mstype.int32) type_cast_op_float = C.TypeCast(mstype.float16) diff --git a/model_zoo/official/nlp/pangu_alpha/src/metrics.py b/model_zoo/official/nlp/pangu_alpha/src/metrics.py new file mode 100644 index 00000000000..4d9e8ca5e8c --- /dev/null +++ b/model_zoo/official/nlp/pangu_alpha/src/metrics.py @@ -0,0 +1,58 @@ +# Copyright 2021 Huawei Technologies Co., Ltd +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================ +""" +Eval metrics +""" + +import math +from mindspore.nn.metrics import Metric +from mindspore import context +from mindspore.communication.management import get_rank, get_group_size + +class PPLMetric(Metric): + """ + Ppl metric + """ + + def __init__(self, data_length): + super(PPLMetric, self).__init__() + self.clear() + self.data_length = data_length + pipeline_stages = context.get_auto_parallel_context("pipeline_stages") + per_stage_device_num = get_group_size() // pipeline_stages + stage_id = get_rank() // per_stage_device_num + self.is_last_stage = (stage_id == pipeline_stages - 1) + + def clear(self): + """Clear the internal evaluation result.""" + self.PPL = [] + self.tokens_count = 0 + + def update(self, *inputs): # inputs + """Update list of ppl""" + if not self.is_last_stage: + return + logits = inputs[0].asnumpy().flatten().tolist() # logits + self.PPL.append(logits[0] * self.data_length) + self.tokens_count += 1 + + def eval(self): + if not self.is_last_stage: + return 0 + val_loss = sum(self.PPL) / (self.tokens_count * self.data_length) + ppl = math.exp(min(20, val_loss)) + print("====" * 20 + " ppl end") + print("====" * 20 + " ppl: {}".format(ppl)) + return ppl diff --git a/model_zoo/official/nlp/pangu_alpha/src/utils.py b/model_zoo/official/nlp/pangu_alpha/src/utils.py index 63a6a73cd09..83465a8d3f1 100644 --- a/model_zoo/official/nlp/pangu_alpha/src/utils.py +++ b/model_zoo/official/nlp/pangu_alpha/src/utils.py @@ -405,6 +405,10 @@ def get_args(inference=False): required=False, default=None, help='Location of data.') + parser.add_argument('--eval_data_url', + required=False, + default=None, + help='Location of eval data.') parser.add_argument('--train_url', required=False, default=None, @@ -448,6 +452,14 @@ def get_args(inference=False): type=int, default=0, help="Enable incremental training. Default 0.") + parser.add_argument("--train_and_eval_mode", + type=int, + default=0, + help="Enable evaling while training. Default 0.") + parser.add_argument("--eval_steps", + type=int, + default=10, + help="The eval step in train and eval mode. Default 10.") add_training_params(parser) if inference: add_inference_params(parser) diff --git a/model_zoo/official/nlp/pangu_alpha/train.py b/model_zoo/official/nlp/pangu_alpha/train.py index fd2a83a3784..e184260cc7a 100644 --- a/model_zoo/official/nlp/pangu_alpha/train.py +++ b/model_zoo/official/nlp/pangu_alpha/train.py @@ -18,13 +18,12 @@ PanguAlpha train script import os import math -import time from mindspore import context from mindspore.train.model import Model import mindspore.communication.management as D from mindspore.context import ParallelMode import mindspore.nn as nn -from mindspore.train.callback import TimeMonitor, Callback +from mindspore.train.callback import TimeMonitor from mindspore.nn.wrap.loss_scale import DynamicLossScaleUpdateCell import mindspore.common.dtype as mstype from mindspore.parallel import set_algo_parameters @@ -37,40 +36,10 @@ from src.pangu_alpha_wrapcell import PanguAlphaTrainOneStepWithLossScaleCell, Pa from src.pangu_alpha_config import PANGUALPHAConfig, set_parse from src.utils import LearningRate, get_args, FP32StateAdamWeightDecay from src.utils import download_data +from src.callbacks import EvalCallBack, LossCallBack +from src.metrics import PPLMetric -class LossCallBack(Callback): - """ - Monitor the loss in training. - If the loss in NAN or INF terminating training. - """ - - def __init__(self, dataset_size=-1, local_rank=0, has_trained_epoch=0, has_trained_step=0, micro_size=1): - super(LossCallBack, self).__init__() - self._dataset_size = dataset_size - self.local_rank = local_rank - self.has_trained_epoch = has_trained_epoch - self.has_trained_step = has_trained_step - self.micro_size = micro_size - print("load has trained epoch :{} and step: {}".format(has_trained_epoch, has_trained_step), flush=True) - - def step_end(self, run_context): - """ - Print loss after each step - """ - cb_params = run_context.original_args() - if self._dataset_size > 0 and self.local_rank % 8 == 0: - percent, epoch_num = math.modf(cb_params.cur_step_num / - self._dataset_size) - if percent == 0: - epoch_num -= 1 - date = time.asctime(time.localtime(time.time())) - loss_value = cb_params.net_outputs[0].asnumpy() / self.micro_size - print("time: {} local_rank: {}, epoch: {}, step: {}, output is {}, overflow is {}, scale is {}". - format(date, int(self.local_rank), int(epoch_num) + int(self.has_trained_epoch), - cb_params.cur_step_num + int(self.has_trained_step), loss_value, - cb_params.net_outputs[1].asnumpy(), cb_params.net_outputs[2].asnumpy())) - project_root = os.path.abspath( os.path.dirname(os.path.realpath(__file__)) + os.path.sep + "..") @@ -101,73 +70,59 @@ def run_train(args_opt): The main training process. """ # Set execution mode - context.set_context(mode=context.GRAPH_MODE, device_target=args_opt.device_target) - context.set_context(variable_memory_max_size="31GB") + context.set_context(mode=context.GRAPH_MODE, device_target=args_opt.device_target, variable_memory_max_size="31GB") # Set parallel context if args_opt.distribute == "true": D.init() device_num = D.get_group_size() rank = D.get_rank() print("rank_id is {}, device_num is {}".format(rank, device_num)) - context.reset_auto_parallel_context() context.set_auto_parallel_context( - parallel_mode=ParallelMode.SEMI_AUTO_PARALLEL, - gradients_mean=False, - full_batch=bool(args_opt.full_batch), - strategy_ckpt_load_file=args_opt.strategy_load_ckpt_path, + parallel_mode=ParallelMode.SEMI_AUTO_PARALLEL, gradients_mean=False, + full_batch=bool(args_opt.full_batch), strategy_ckpt_load_file=args_opt.strategy_load_ckpt_path, enable_parallel_optimizer=bool(args_opt.optimizer_shard)) set_algo_parameters(elementwise_op_strategy_follow=True) _set_multi_subgraphs() - else: rank = 0 device_num = 1 context.set_context(save_graphs=False, save_graphs_path="./graphs_of_device_id_" + str(rank)) # copy data from the cloud to the /cache/Data cache_url = '/cache/Data/' + eval_cache_url = '/cache/EvalData/' if args_opt.offline: cache_url = args_opt.data_url + eval_cache_url = args_opt.eval_data_url else: download_data(src_data_url=args_opt.data_url, tgt_data_path=cache_url, rank=rank) + download_data(src_data_url=args_opt.eval_data_url, tgt_data_path=eval_cache_url, rank=rank) # Set model property model_parallel_num = args_opt.op_level_model_parallel_num data_parallel_num = int(device_num / model_parallel_num) + if data_parallel_num <= 1 and args_opt.optimizer_shard == 1: + raise ValueError("The dp must large than 1 when applying optimizer shard.") batch_size = args_opt.per_batch_size * data_parallel_num config = PANGUALPHAConfig( - data_parallel_num=data_parallel_num, - model_parallel_num=model_parallel_num, - batch_size=batch_size, - seq_length=args_opt.seq_length, - vocab_size=args_opt.vocab_size, - embedding_size=args_opt.embedding_size, - num_layers=args_opt.num_layers, - num_heads=args_opt.num_heads, - expand_ratio=4, - dropout_rate=0.1, - compute_dtype=mstype.float16, - stage_num=args_opt.stage_num, - micro_size=args_opt.micro_size, - eod_reset=bool(args_opt.eod_reset), - load_ckpt_path=args_opt.load_ckpt_path, + data_parallel_num=data_parallel_num, model_parallel_num=model_parallel_num, + batch_size=batch_size, seq_length=args_opt.seq_length, + vocab_size=args_opt.vocab_size, embedding_size=args_opt.embedding_size, + num_layers=args_opt.num_layers, num_heads=args_opt.num_heads, + expand_ratio=4, dropout_rate=0.1, compute_dtype=mstype.float16, + stage_num=args_opt.stage_num, micro_size=args_opt.micro_size, + eod_reset=bool(args_opt.eod_reset), load_ckpt_path=args_opt.load_ckpt_path, param_init_type=mstype.float32 if args_opt.param_init_type == 'fp32' else mstype.float16, word_emb_dp=bool(args_opt.word_emb_dp)) print("===config is: ", config, flush=True) - # Define network pangu_alpha = PanguAlpha(config) loss = CrossEntropyLoss(config) - pangu_alpha_with_loss = PanguAlphaWithLoss(config, pangu_alpha, loss) - pangu_alpha_with_loss = _VirtualDatasetCell(pangu_alpha_with_loss) - + pangu_alpha_with_loss_net = PanguAlphaWithLoss(config, pangu_alpha, loss) + pangu_alpha_with_loss = _VirtualDatasetCell(pangu_alpha_with_loss_net) print("=====args_opt is: ", args_opt, flush=True) - # Warm-up and cosine decay learning rate - lr = LearningRate(learning_rate=args_opt.start_lr, - end_learning_rate=args_opt.end_lr, - warmup_steps=args_opt.warmup_step, - decay_steps=200000) - + lr = LearningRate(learning_rate=args_opt.start_lr, end_learning_rate=args_opt.end_lr, + warmup_steps=args_opt.warmup_step, decay_steps=200000) params = pangu_alpha.trainable_params() group_params = set_weight_decay(params) if args_opt.optimizer == "lamb": @@ -180,36 +135,37 @@ def run_train(args_opt): loss_scale_value = math.pow(2, 32) epoch_num = args_opt.epoch_size # Dataset loading mindrecord files - ds = create_dataset(config.batch_size, data_path=cache_url, - data_start_index=0, eod_reset=config.eod_reset, full_batch=bool(args_opt.full_batch), - eod_id=args_opt.eod_id, device_num=device_num, rank=rank, - column_name=args_opt.data_column_name, epoch=epoch_num) - step_per_epoch = ds.get_dataset_size() - callback_size = args_opt.sink_size - actual_epoch_num = int(epoch_num * step_per_epoch / callback_size) - callback = [ - TimeMonitor(callback_size), - LossCallBack(callback_size, rank, 0, 0) - ] + ds = create_dataset(config.batch_size, data_path=cache_url, data_start_index=0, eod_reset=config.eod_reset, + full_batch=bool(args_opt.full_batch), eod_id=args_opt.eod_id, device_num=device_num, + rank=rank, column_name=args_opt.data_column_name, epoch=epoch_num) + actual_epoch_num = int(epoch_num * ds.get_dataset_size() / args_opt.sink_size) + callback = [TimeMonitor(args_opt.sink_size), LossCallBack(args_opt.sink_size, rank, 0, 0)] update_cell = DynamicLossScaleUpdateCell(loss_scale_value=loss_scale_value, scale_factor=2, scale_window=1000) pangu_alpha_with_grads = PanguAlphaTrainOneStepWithLossScaleCell( pangu_alpha_with_loss, optimizer=optimizer, scale_update_cell=update_cell, enable_global_norm=True, config=config) - model = Model(pangu_alpha_with_grads) + if args_opt.train_and_eval_mode: + ds_eval = create_dataset(config.batch_size, data_path=eval_cache_url, + data_start_index=0, eod_reset=config.eod_reset, full_batch=bool(args_opt.full_batch), + eod_id=args_opt.eod_id, device_num=device_num, rank=rank, + column_name=args_opt.data_column_name, epoch=epoch_num, + num_samples=args_opt.eval_steps * config.batch_size) + ppl_metric = PPLMetric(config.seq_length) + model = Model(pangu_alpha_with_grads, eval_network=pangu_alpha_with_loss, metrics={"ppl": ppl_metric}) + callback.append(EvalCallBack(model, ds_eval, ppl_metric)) + else: + model = Model(pangu_alpha_with_grads) if args_opt.incremental_training: from mindspore.train.serialization import load_distributed_checkpoint - strategy = model.infer_train_layout(train_dataset=ds, sink_size=callback_size) + strategy = model.infer_train_layout(train_dataset=ds, sink_size=args_opt.sink_size) print("======start load_distributed checkpoint", flush=True) # For 2.6B and 13B models, the number of ckpt files is 512. - ckpt_name = 'filerted' - ckpt_file_list = [os.path.join(args_opt.load_ckpt_path, f"{ckpt_name}_{ckpt_rank}.ckpt") for ckpt_rank in + ckpt_file_list = [os.path.join(args_opt.load_ckpt_path, f"filerted_{ckpt_rank}.ckpt") for ckpt_rank in range(0, 512)] print(f"Loading from path {ckpt_file_list[0]}", flush=True) - # Load checkpoint files load_distributed_checkpoint(model.train_network, ckpt_file_list, strategy) print("Dataset size: {}, actual_epoch_num: {}".format(ds.get_dataset_size(), actual_epoch_num), flush=True) - model.train(actual_epoch_num, ds, callbacks=callback, sink_size=callback_size, dataset_sink_mode=True) - + model.train(actual_epoch_num, ds, callbacks=callback, sink_size=args_opt.sink_size, dataset_sink_mode=True) def run_train_pipeline(args_opt): r""" @@ -224,12 +180,9 @@ def run_train_pipeline(args_opt): print("rank_id is {}, device_num is {}".format(rank_id, device_num)) context.reset_auto_parallel_context() context.set_auto_parallel_context( - parallel_mode=ParallelMode.SEMI_AUTO_PARALLEL, - gradients_mean=False, - full_batch=bool(args_opt.full_batch), - loss_repeated_mean=True, - device_num=device_num, - enable_parallel_optimizer=bool(args_opt.optimizer_shard), + parallel_mode=ParallelMode.SEMI_AUTO_PARALLEL, gradients_mean=False, + full_batch=bool(args_opt.full_batch), loss_repeated_mean=True, + device_num=device_num, enable_parallel_optimizer=bool(args_opt.optimizer_shard), pipeline_stages=args_opt.stage_num) set_algo_parameters(elementwise_op_strategy_follow=True) _set_multi_subgraphs() @@ -238,13 +191,18 @@ def run_train_pipeline(args_opt): device_num = 1 # copy data from the cloud to the /cache/Data cache_url = '/cache/Data/' + eval_cache_url = '/cache/EvalData/' if args_opt.offline: cache_url = args_opt.data_url + eval_cache_url = args_opt.eval_data_url else: download_data(src_data_url=args_opt.data_url, tgt_data_path=cache_url, rank=rank_id) + download_data(src_data_url=args_opt.eval_data_url, tgt_data_path=eval_cache_url, rank=rank_id) model_parallel_num = args_opt.op_level_model_parallel_num stage_device_num = int(device_num / args_opt.stage_num) data_parallel_num = int(stage_device_num / model_parallel_num) + if data_parallel_num <= 1 and args_opt.optimizer_shard == 1: + raise ValueError("The dp must large than 1 when applying optimizer shard.") per_batch_size = args_opt.per_batch_size batch_size = per_batch_size * data_parallel_num * args_opt.micro_size config = PANGUALPHAConfig( @@ -267,8 +225,8 @@ def run_train_pipeline(args_opt): print("===config is: ", config, flush=True) pangu_alpha = PanguAlpha(config) loss = CrossEntropyLoss(config) - pangu_alpha_with_loss = PipelineCell(PanguAlphaWithLoss(config, pangu_alpha, loss), config.micro_size) - pangu_alpha_with_loss = _VirtualDatasetCell(pangu_alpha_with_loss) + pangu_alpha_with_loss_net = PipelineCell(PanguAlphaWithLoss(config, pangu_alpha, loss), config.micro_size) + pangu_alpha_with_loss = _VirtualDatasetCell(pangu_alpha_with_loss_net) print("=====args_opt is: ", args_opt, flush=True) lr = LearningRate(learning_rate=args_opt.start_lr, end_learning_rate=args_opt.end_lr, warmup_steps=args_opt.warmup_step, decay_steps=args_opt.decay_steps) @@ -294,6 +252,8 @@ def run_train_pipeline(args_opt): update_cell = DynamicLossScaleUpdateCell(loss_scale_value=loss_scale_value, scale_factor=2, scale_window=1000) pangu_alpha_with_grads = PanguAlphaTrainPipelineWithLossScaleCell( pangu_alpha_with_loss, optimizer=optimizer, config=config, scale_update_cell=update_cell) + if args_opt.train_and_eval_mode: + raise ValueError("The pipeline train_and_eval_mode is not supported yet") model = Model(pangu_alpha_with_grads) model.train(actual_epoch_num, ds, callbacks=callback, sink_size=callback_size, dataset_sink_mode=True)