Browse Source

!3406 fix optimizer parallel problems

Merge pull request !3406 from gziyan/fix_optimizer_parallel_r0.6
tags/v0.6.0-beta
mindspore-ci-bot Gitee 6 years ago
parent
commit
e62137f7c0
4 changed files with 7 additions and 7 deletions
  1. +1
    -1
      mindspore/nn/optim/adam.py
  2. +1
    -1
      mindspore/nn/optim/lamb.py
  3. +3
    -3
      mindspore/nn/optim/optimizer.py
  4. +2
    -2
      mindspore/nn/wrap/grad_reducer.py

+ 1
- 1
mindspore/nn/optim/adam.py View File

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

+ 1
- 1
mindspore/nn/optim/lamb.py View File

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


+ 3
- 3
mindspore/nn/optim/optimizer.py View File

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



+ 2
- 2
mindspore/nn/wrap/grad_reducer.py View File

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


Loading…
Cancel
Save