forked from huawei/mindspore2022
!21568 parallel_train_and_eval_fix
Merge pull request !21568 from yao_yf/parallel_train_and_eval_fix
This commit is contained in:
commit
75e8783495
|
|
@ -27,7 +27,7 @@ namespace opt {
|
|||
namespace {
|
||||
// insert tensormove for some cnode even if not a Ref cnode
|
||||
const std::set<std::string> 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;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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<Strategy>(stage_id, strategy);
|
||||
|
|
|
|||
|
|
@ -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<std::string>(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) {
|
||||
|
|
|
|||
|
|
@ -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<AnfNodePtr>
|
|||
|
||||
std::string MirrorOpName();
|
||||
|
||||
CommInfo GetCommInfo();
|
||||
|
||||
void ReorderForPipelineSplit(const FuncGraphPtr &root, const FuncGraphManagerPtr &manager, int64_t pipeline_stages);
|
||||
} // namespace parallel
|
||||
} // namespace mindspore
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Reference in New Issue