From 9f264b6e550c8d5f845aa20ae93a77c92802c795 Mon Sep 17 00:00:00 2001 From: Ziyan Date: Fri, 24 Jul 2020 11:20:50 +0800 Subject: [PATCH] fix optimizer parallel problems --- mindspore/nn/optim/adam.py | 2 +- mindspore/nn/optim/lamb.py | 2 +- mindspore/nn/optim/optimizer.py | 6 +++--- mindspore/nn/wrap/grad_reducer.py | 4 ++-- 4 files changed, 7 insertions(+), 7 deletions(-) diff --git a/mindspore/nn/optim/adam.py b/mindspore/nn/optim/adam.py index b823bf69f3..9c834f26de 100755 --- a/mindspore/nn/optim/adam.py +++ b/mindspore/nn/optim/adam.py @@ -450,5 +450,5 @@ class AdamWeightDecay(Optimizer): self.parameters, self.moments1, self.moments2, gradients, self.decay_flags, self.optim_filter) if self.use_parallel: - optim_result = self.broadcast_params(optim_result) + self.broadcast_params(optim_result) return optim_result diff --git a/mindspore/nn/optim/lamb.py b/mindspore/nn/optim/lamb.py index 0d2552b8c1..7fc79bf418 100755 --- a/mindspore/nn/optim/lamb.py +++ b/mindspore/nn/optim/lamb.py @@ -312,7 +312,7 @@ class Lamb(Optimizer): self.decay_flags, self.optim_filter) if self.use_parallel: - optim_result = self.broadcast_params(optim_result) + self.broadcast_params(optim_result) if not self.dynamic_lr: F.control_depend(lr, self.assignadd(self.global_step, 1)) diff --git a/mindspore/nn/optim/optimizer.py b/mindspore/nn/optim/optimizer.py index 9379e395ae..acfc09630f 100755 --- a/mindspore/nn/optim/optimizer.py +++ b/mindspore/nn/optim/optimizer.py @@ -466,7 +466,7 @@ class Optimizer(Cell): param_group.append(F.make_tuple()) key_group.append(F.make_tuple()) for i in range(self.param_length): - param_group[self.param_rank[i]] = param_group[self.param_rank[i]] + (optim_result[i],) + param_group[self.param_rank[i]] = param_group[self.param_rank[i]] + (self.parameters[i],) key = P.MakeRefKey(self.param_names[i])() key_group[self.param_rank[i]] = key_group[self.param_rank[i]] + (key,) new_param_group = [] @@ -476,9 +476,9 @@ class Optimizer(Cell): new_param_group.append(next_params) for i in range(F.tuple_len(next_params)): F.assign(key_group[root][i], next_params[i]) - status = True + status = F.control_depend(optim_result, new_param_group[0][0]) for i in range(self.dev_num - 1): - status = F.control_depend(new_param_group[i][0], new_param_group[i+1]) + status = F.depend(F.control_depend(new_param_group[i], new_param_group[i+1][0]), status) return status diff --git a/mindspore/nn/wrap/grad_reducer.py b/mindspore/nn/wrap/grad_reducer.py index e67e74d9ef..1039be0619 100644 --- a/mindspore/nn/wrap/grad_reducer.py +++ b/mindspore/nn/wrap/grad_reducer.py @@ -25,7 +25,7 @@ import mindspore.common.dtype as mstype reduce_opt = C.MultitypeFuncGraph("reduce_opt") -def _init_allreduce_operators(length): +def _init_allreduce_operators(length, split_indices): """ initialize allreduce communication operators""" group = 1 fusion = () @@ -318,7 +318,7 @@ class DistributedGradReducer(Cell): split_indices = auto_parallel_context().get_all_reduce_fusion_split_indices() if is_parallel_optimizer and split_indices: self.split_fusion = True - self.op_list = _init_allreduce_operators(len(parameters)) + self.op_list = _init_allreduce_operators(len(parameters), split_indices) else: self.split_fusion = False self.allreduce = AllReduce().add_prim_attr('fusion', 1)