!21568 parallel_train_and_eval_fix

Merge pull request !21568 from yao_yf/parallel_train_and_eval_fix
This commit is contained in:
i-robot 2021-08-10 01:44:23 +00:00 committed by Gitee
commit 75e8783495
11 changed files with 346 additions and 116 deletions

View File

@ -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;
}

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

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

View File

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

View File

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

View File

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