!32630 Modify Optimizer to support flatten parameters

Merge pull request !32630 from hewei/flatten_weights
This commit is contained in:
i-robot 2022-04-08 06:48:09 +00:00 committed by Gitee
commit 3416b3a94b
No known key found for this signature in database
GPG Key ID: 173E9B9CA92EEF8F
4 changed files with 108 additions and 11 deletions

View File

@ -231,15 +231,15 @@ class Parameter(Tensor_):
"""Set `set_data` of current `Parameter`."""
if isinstance(data, bool):
raise ValueError('Parameter data can not be `bool`')
if isinstance(data, Tensor) and data.has_init:
if isinstance(data, Tensor):
if not data.has_init:
# make a copy of Tensor to init the parameter.
return (Tensor, data.asnumpy())
if not _is_fl_mode():
if _is_in_parallel_mode() or _is_role_worker() or _is_role_sched() or _is_role_pserver():
# do not init data while in auto parallel.
return (Tensor, None, data.dtype, data.shape, data.init)
data = data.init_data().asnumpy()
elif isinstance(data, Tensor):
# make a copy of Tensor to init the parameter
return (Tensor, data.asnumpy(),)
return (Tensor, data.init_data())
if isinstance(data, int):
return (Tensor, data, mstype.int32)
if isinstance(data, float):
@ -537,6 +537,16 @@ class Parameter(Tensor_):
raise ValueError(f"Can not change the shape of Parameter which has been initialized."
f" Current shape is {current_shape}, and incoming is {data_shape}.")
@staticmethod
def _from_tensor(tensor, *args, **kwargs):
"""Create a `Parameter` that data is shared from a `Tensor`."""
if not isinstance(tensor, Tensor_):
raise TypeError(f"The type of input must be Tensor, but got {type(tensor)}.")
param = Tensor_.__new__(Parameter)
Tensor_.__init__(param, tensor)
Parameter.__init__(param, tensor, *args, **kwargs)
return param
def set_data(self, data, slice_shape=False):
"""
Set Parameter's data.

View File

@ -1,4 +1,4 @@
# Copyright 2020-2021 Huawei Technologies Co., Ltd
# Copyright 2020-2022 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.
@ -156,12 +156,14 @@ class Optimizer(Cell):
self._unique = True
self._target = context.get_context("device_target")
self._use_flattened_params = False
self.dynamic_lr = False
self.assignadd = P.AssignAdd()
self.global_step = Parameter(initializer(0, [1], mindspore.int32), name='global_step')
self.is_group = False
self.is_group_lr = False
self.is_group_params_ordered = False
self.use_parallel = False
learning_rate = self._preprocess_single_lr(learning_rate)
if isinstance(parameters[0], dict):
self.is_group = True
@ -193,6 +195,7 @@ class Optimizer(Cell):
self.exec_weight_decay = any(self.decay_flags)
self.grad_centralization_flags = tuple(self.group_grad_centralization)
else:
parameters = self._get_flattened_params(parameters)
self.parameters = ParameterTuple(parameters)
decay_filter = lambda x: 'beta' not in x.name and 'gamma' not in x.name
self.decay_flags = tuple(decay_filter(x) for x in self.parameters)
@ -222,6 +225,27 @@ class Optimizer(Cell):
self._use_parallel_optimizer()
self.enable_tuple_broaden = True
def _get_flattened_params(self, parameters):
"""Get parameters for each contiguous memory chunks used by input parameters if they are flattened."""
if self.is_group:
# We don't use flattened parameters when parameters are grouped.
return parameters
# Check whether parameters are flattened.
flattened = Tensor._is_flattened(parameters) # pylint: disable=W0212
if not flattened:
# Parameters are not flattened.
return parameters
# Try to get chunk tensors from flattened parameters.
chunk_tensors = Tensor._get_flattened_tensors(parameters) # pylint: disable=W0212
if not chunk_tensors:
# Failed to get chunk tensors.
logger.warning("Parameters are not properly falttened, fallback to not flattened parameters.")
return parameters
# Convert chunk tensors to parameters.
self._use_flattened_params = True
return [Parameter._from_tensor(t, name='_chunk_param_' + str(t.dtype)) # pylint: disable=W0212
for t in chunk_tensors]
def _use_parallel_optimizer(self):
"""Indicates whether to use automatic parallelism."""
if context.get_auto_parallel_context("enable_parallel_optimizer"):
@ -235,10 +259,7 @@ class Optimizer(Cell):
raise RuntimeError("For 'Optimizer', parallel optimizer is not supported in {}, you should set "
"parallel mode to 'data_parallel', 'semi_auto_parallel' or 'auto_parallel'."
.format(_get_parallel_mode()))
else:
self.use_parallel = False
else:
self.use_parallel = False
if self.use_parallel:
if not self._support_parallel_optimizer:
raise RuntimeError("For 'Optimizer', parallel optimizer shard doest not support "

View File

@ -16,6 +16,7 @@
import numpy as np
import pytest
import mindspore as ms
from mindspore import Tensor
from mindspore.common.parameter import Parameter
from mindspore.nn.optim import Optimizer, SGD, Adam, AdamWeightDecay
@ -101,3 +102,50 @@ class TestUnsupportParam():
with pytest.raises(TypeError):
paramsTensor = Parameter(Tensor(np.zeros([1, 2, 3])), "x")
SGD(paramsTensor)
class TestFlattenParams:
""" Test Optimizer with flatten parameters """
def __init__(self):
self.p1 = None
self.p2 = None
self.p3 = None
self.params = []
def setup_method(self):
self.p1 = Parameter(Tensor([1], ms.float32), name="p1")
self.p2 = Parameter(Tensor([2], ms.float32), name="p2")
self.p3 = Parameter(Tensor([3], ms.float32), name="p3")
self.params = [self.p1, self.p2, self.p3]
def test_not_flattened_params(self):
"""
Feature: Flatten weights.
Description: Optimizer with not flattened parameters.
Expectation: The Optimizer works as expected.
"""
opt = Optimizer(0.1, self.params)
assert not opt._use_flattened_params # pylint: disable=W0212
assert len(opt.parameters) == 3
assert len(opt.cache_enable) == 3
def test_with_flattened_params(self):
"""
Feature: Flatten weights.
Description: Optimizer with flattened parameters.
Expectation: The Optimizer works as expected.
"""
Tensor._flatten_tensors(self.params) # pylint: disable=W0212
opt = Optimizer(0.1, self.params)
assert opt._use_flattened_params # pylint: disable=W0212
assert len(opt.parameters) == 1
assert len(opt.cache_enable) == 1
assert opt.parameters[0].dtype == ms.float32
assert opt.parameters[0].shape == (3,)
assert opt.parameters[0].size == 3
assert np.allclose(opt.parameters[0].asnumpy(), np.array([1, 2, 3]))
self.p1.asnumpy()[0] = 6
self.p2.asnumpy()[0] = 6
self.p3.asnumpy()[0] = 6
assert np.allclose(opt.parameters[0].asnumpy(), np.array([6, 6, 6]))

View File

@ -1,5 +1,5 @@
# Copyright 2020 Huawei Technologies Co., Ltd
# Copyright 2020-2022 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.
@ -22,6 +22,7 @@ from mindspore._checkparam import Validator
from mindspore.common import dtype as mstype
from mindspore.common.initializer import initializer
def test_parameter_init():
dat = np.array([[1, 2, 3], [2, 3, 4]])
tensor = Tensor(dat)
@ -134,6 +135,7 @@ def test_check_str_by_regular():
with pytest.raises(ValueError):
Validator.check_str_by_regular(str6)
def test_parameter_compute():
para_1 = Parameter(initializer('ones', [1, 2, 3], mstype.int32), 'test1')
para_2 = Parameter(initializer('ones', [1, 2, 3], mstype.int32), 'test2')
@ -242,6 +244,7 @@ def test_parameter_as_output():
context.set_auto_parallel_context(parallel_mode="semi_auto_parallel")
initial_input = initializer('One', shape=(2,), dtype=mstype.int32)
updated_input = Tensor([2, 2], mstype.int32)
class Net(nn.Cell):
def __init__(self, initial, updated):
super().__init__()
@ -250,6 +253,7 @@ def test_parameter_as_output():
self.p = Parameter(self.initial, name="weight")
self.new_p = self.p.init_data()
self.new_p.set_data(self.updated)
def construct(self):
return self.new_p
@ -257,3 +261,17 @@ def test_parameter_as_output():
output = net()
assert np.array_equal(output.asnumpy(), np.array([2, 2], np.int32))
context.reset_auto_parallel_context()
def test_parameter_init_from_tensor():
"""
Feature: Parameter initialize.
Description: Parameter initialized from a given tensor, data is shared.
Expectation: The Parameter and the tensor share same data buffer.
"""
tensor = Tensor([1], mstype.float32)
param = Parameter._from_tensor(tensor, name="mypara") # pylint: disable=W0212
assert param.name == "mypara"
assert np.allclose(param.asnumpy(), np.array([1]))
tensor.asnumpy()[0] = 2
assert np.allclose(param.asnumpy(), np.array([2]))