diff --git a/mindspore/python/mindspore/nn/transformer/loss.py b/mindspore/python/mindspore/nn/transformer/loss.py index 5b1d361eb8..f21938dd8d 100644 --- a/mindspore/python/mindspore/nn/transformer/loss.py +++ b/mindspore/python/mindspore/nn/transformer/loss.py @@ -22,6 +22,8 @@ from mindspore.ops import operations as P from mindspore.ops import functional as F from mindspore.nn import Cell from mindspore.nn.loss.loss import _check_is_tensor +from mindspore.parallel._utils import _get_parallel_mode, _is_sharding_propagation +from mindspore.context import ParallelMode from .layers import _check_input_dtype, _check_input_shape from .op_parallel_config import default_dpmp_config, OpParallelConfig @@ -65,32 +67,59 @@ class CrossEntropyLoss(Cell): def __init__(self, parallel_config=default_dpmp_config): super(CrossEntropyLoss, self).__init__() - if not isinstance(parallel_config, OpParallelConfig) and not isinstance(parallel_config): - raise TypeError("For 'CrossEntropyLoss', the class variable 'parallel_config' must be OpParallelConfig" - ", but got the type: {}.".format(type(parallel_config))) - dp = parallel_config.data_parallel - mp = parallel_config.model_parallel - self.sum = P.ReduceSum().shard(((dp, mp),)) - self.onehot = P.OneHot().shard(((dp, mp), (), ())) - # on/off value for onehot, for smooth labeling, modify the off_value - self.on_value = Tensor(1.0, mstype.float32) - self.off_value = Tensor(0.0, mstype.float32) - self.max = P.ArgMaxWithValue(axis=-1, keep_dims=True).shard( - ((dp, mp),)) - self.eps_const = Tensor(1e-24, mstype.float32) - self.sub = P.Sub().shard(((dp, mp), (dp, 1))) - self.exp = P.Exp().shard(((dp, mp),)) - self.div = P.RealDiv().shard(((dp, mp), (dp, 1))) - self.log = P.Log().shard(((dp, mp),)) - self.add = P.Add().shard(((dp, mp), ())) - self.mul = P.Mul().shard( - ((dp, mp), (dp, mp))) - self.neg = P.Neg().shard(((dp, mp),)) - self.sum2 = P.ReduceSum().shard(((1,),)) + if _get_parallel_mode() in (ParallelMode.AUTO_PARALLEL,) and _is_sharding_propagation(): + if not isinstance(parallel_config, OpParallelConfig) and not isinstance(parallel_config): + raise TypeError("For 'CrossEntropyLoss', the class variable 'parallel_config' must be OpParallelConfig" + ", but got the type: {}.".format(type(parallel_config))) + dp = parallel_config.data_parallel + mp = parallel_config.model_parallel + self.sum = P.ReduceSum() + self.onehot = P.OneHot() + # on/off value for onehot, for smooth labeling, modify the off_value + self.on_value = Tensor(1.0, mstype.float32) + self.off_value = Tensor(0.0, mstype.float32) + self.max = P.ArgMaxWithValue(axis=-1, keep_dims=True).shard( + ((dp, mp),)) + self.eps_const = Tensor(1e-24, mstype.float32) + self.sub = P.Sub() + self.exp = P.Exp() + self.div = P.RealDiv() + self.log = P.Log() + self.add = P.Add() + self.mul = P.Mul() + self.neg = P.Neg() + self.sum2 = P.ReduceSum().shard(((1,),)) - self.mul2 = P.Mul().shard(((1,), (1,))) - self.add2 = P.Add() - self.div2 = P.RealDiv() + self.mul2 = P.Mul().shard(((1,), (1,))) + self.add2 = P.Add() + self.div2 = P.RealDiv() + else: + if not isinstance(parallel_config, OpParallelConfig) and not isinstance(parallel_config): + raise TypeError("For 'CrossEntropyLoss', the class variable 'parallel_config' must be OpParallelConfig" + ", but got the type: {}.".format(type(parallel_config))) + dp = parallel_config.data_parallel + mp = parallel_config.model_parallel + self.sum = P.ReduceSum().shard(((dp, mp),)) + self.onehot = P.OneHot().shard(((dp, mp), (), ())) + # on/off value for onehot, for smooth labeling, modify the off_value + self.on_value = Tensor(1.0, mstype.float32) + self.off_value = Tensor(0.0, mstype.float32) + self.max = P.ArgMaxWithValue(axis=-1, keep_dims=True).shard( + ((dp, mp),)) + self.eps_const = Tensor(1e-24, mstype.float32) + self.sub = P.Sub().shard(((dp, mp), (dp, 1))) + self.exp = P.Exp().shard(((dp, mp),)) + self.div = P.RealDiv().shard(((dp, mp), (dp, 1))) + self.log = P.Log().shard(((dp, mp),)) + self.add = P.Add().shard(((dp, mp), ())) + self.mul = P.Mul().shard( + ((dp, mp), (dp, mp))) + self.neg = P.Neg().shard(((dp, mp),)) + self.sum2 = P.ReduceSum().shard(((1,),)) + + self.mul2 = P.Mul().shard(((1,), (1,))) + self.add2 = P.Add() + self.div2 = P.RealDiv() def construct(self, logits, label, input_mask): self._check_input(logits, label, input_mask) diff --git a/mindspore/python/mindspore/nn/transformer/moe.py b/mindspore/python/mindspore/nn/transformer/moe.py index 4238a3d40f..d78beab3c7 100644 --- a/mindspore/python/mindspore/nn/transformer/moe.py +++ b/mindspore/python/mindspore/nn/transformer/moe.py @@ -26,6 +26,8 @@ from mindspore.ops import functional as F from mindspore.ops.primitive import constexpr from mindspore.nn.cell import Cell from mindspore.nn.layer import Dense +from mindspore.context import ParallelMode +from mindspore.parallel._utils import _get_parallel_mode, _is_sharding_propagation from .op_parallel_config import default_moeparallel_config __all__ = [ @@ -51,6 +53,7 @@ class MoEConfig: >>> from mindspore.nn.transformer import MoEConfig >>> moe_config = MoEConfig(expert_num=4, capacity_factor=5.0, aux_loss_factor=0.05, num_experts_chosen=1) """ + def __init__(self, expert_num=1, capacity_factor=1.1, aux_loss_factor=0.05, num_experts_chosen=1): Validator.check_positive_int(expert_num, "expert_num") @@ -133,6 +136,7 @@ class MoE(Cell): Outputs: Tensor, the output of this layer after mapping. The shape is `[batch, seq_length, hidden_size]`. """ + def __init__(self, hidden_size, ffn_hidden_size, dropout_rate, @@ -141,36 +145,66 @@ class MoE(Cell): moe_config=default_moe_config, parallel_config=default_moeparallel_config): super(MoE, self).__init__() - self.hidden_size = hidden_size - self.expert_dim = moe_config.expert_num - self.capacity_factor = moe_config.capacity_factor - self.aux_loss_factor = moe_config.aux_loss_factor - self.num_experts_chosen = moe_config.num_experts_chosen - self.dp_group = parallel_config.data_parallel - self.dp = parallel_config.data_parallel - self.ep = parallel_config.expert_parallel - from .transformer import FeedForward + if _get_parallel_mode() in (ParallelMode.AUTO_PARALLEL,) and _is_sharding_propagation(): + self.hidden_size = hidden_size + self.expert_dim = moe_config.expert_num + self.capacity_factor = moe_config.capacity_factor + self.aux_loss_factor = moe_config.aux_loss_factor + self.num_experts_chosen = moe_config.num_experts_chosen + self.dp_group = parallel_config.data_parallel + self.dp = parallel_config.data_parallel + self.ep = parallel_config.expert_parallel + from .transformer import FeedForward - self.ffn = FeedForward(hidden_size=hidden_size, - ffn_hidden_size=ffn_hidden_size, - dropout_rate=dropout_rate, - hidden_act=hidden_act, - expert_num=self.expert_dim, - param_init_type=param_init_type, - parallel_config=parallel_config) - self.reshape = P.Reshape() - self.shape = P.Shape() - self.transpose_2dim = P.Transpose().shard(((self.dp, 1),)) - self.transpose_2dim_ep = P.Transpose().shard(((self.ep, 1),)) - self.transpose_3dim = P.Transpose().shard(((self.dp, 1, 1),)) - self.transpose_4dim_ep = P.Transpose().shard(((self.ep, 1, 1, 1),)) - self.batch_mm = P.BatchMatMul().shard(((self.dp, 1, 1), (self.dp, 1, 1))) - self.batch_mm2 = P.BatchMatMul().shard(((self.dp, 1, 1), (self.dp, 1, 1))) - self.mul = P.Mul().shard(((), ())) - self.router = Router(d_model=hidden_size, moe_config=moe_config, routing_policy=None, - training=True, parallel_config=parallel_config) - self.cast = P.Cast() + self.ffn = FeedForward(hidden_size=hidden_size, + ffn_hidden_size=ffn_hidden_size, + dropout_rate=dropout_rate, + hidden_act=hidden_act, + expert_num=self.expert_dim, + param_init_type=param_init_type, + parallel_config=parallel_config) + self.reshape = P.Reshape() + self.shape = P.Shape() + self.transpose_2dim = P.Transpose().shard(((self.dp, 1),)) + self.transpose_2dim_ep = P.Transpose().shard(((self.ep, 1),)) + self.transpose_3dim = P.Transpose().shard(((self.dp, 1, 1),)) + self.transpose_4dim_ep = P.Transpose().shard(((self.ep, 1, 1, 1),)) + self.batch_mm = P.BatchMatMul().shard(((self.dp, 1, 1), (self.dp, 1, 1))) + self.batch_mm2 = P.BatchMatMul().shard(((self.dp, 1, 1), (self.dp, 1, 1))) + self.mul = P.Mul() + self.router = Router(d_model=hidden_size, moe_config=moe_config, routing_policy=None, + training=True, parallel_config=parallel_config) + self.cast = P.Cast() + else: + self.hidden_size = hidden_size + self.expert_dim = moe_config.expert_num + self.capacity_factor = moe_config.capacity_factor + self.aux_loss_factor = moe_config.aux_loss_factor + self.num_experts_chosen = moe_config.num_experts_chosen + self.dp_group = parallel_config.data_parallel + self.dp = parallel_config.data_parallel + self.ep = parallel_config.expert_parallel + from .transformer import FeedForward + self.ffn = FeedForward(hidden_size=hidden_size, + ffn_hidden_size=ffn_hidden_size, + dropout_rate=dropout_rate, + hidden_act=hidden_act, + expert_num=self.expert_dim, + param_init_type=param_init_type, + parallel_config=parallel_config) + self.reshape = P.Reshape() + self.shape = P.Shape() + self.transpose_2dim = P.Transpose().shard(((self.dp, 1),)) + self.transpose_2dim_ep = P.Transpose().shard(((self.ep, 1),)) + self.transpose_3dim = P.Transpose().shard(((self.dp, 1, 1),)) + self.transpose_4dim_ep = P.Transpose().shard(((self.ep, 1, 1, 1),)) + self.batch_mm = P.BatchMatMul().shard(((self.dp, 1, 1), (self.dp, 1, 1))) + self.batch_mm2 = P.BatchMatMul().shard(((self.dp, 1, 1), (self.dp, 1, 1))) + self.mul = P.Mul().shard(((), ())) + self.router = Router(d_model=hidden_size, moe_config=moe_config, routing_policy=None, + training=True, parallel_config=parallel_config) + self.cast = P.Cast() def construct(self, input_tensor): input_shape = F.shape(input_tensor) @@ -196,8 +230,8 @@ class MoE(Cell): expert_capacity)) # The following four ops are to implement transpose(expert_input, (2, 0, 3, 1)), for that a single transpose # has bad performance - expert_input = self.reshape(expert_input, (self.dp_group*self.hidden_size, - self.expert_dim*expert_capacity)) + expert_input = self.reshape(expert_input, (self.dp_group * self.hidden_size, + self.expert_dim * expert_capacity)) expert_input = self.transpose_2dim(expert_input, (1, 0)) expert_input = self.reshape(expert_input, (self.expert_dim, expert_capacity, self.dp_group, self.hidden_size)) @@ -213,18 +247,18 @@ class MoE(Cell): # The following five ops are to implement transpose(expert_output, (1, 3, 0, 2)), for that a single transpose # has bad performance expert_output = self.reshape(expert_output, (self.expert_dim, - self.dp_group*expert_capacity*self.hidden_size)) + self.dp_group * expert_capacity * self.hidden_size)) expert_output = self.transpose_2dim_ep(expert_output, (1, 0)) expert_output = self.reshape(expert_output, (self.dp_group, expert_capacity, - self.hidden_size*self.expert_dim)) + self.hidden_size * self.expert_dim)) expert_output = self.transpose_3dim(expert_output, (0, 2, 1)) # expert_output's shape: (self.dp_group, self.hidden_size, self.expert_dim, expert_capacity) expert_output = self.reshape(expert_output, (self.dp_group, self.hidden_size, self.expert_dim, expert_capacity)) expert_output = self.reshape(expert_output, (self.dp_group, self.hidden_size, - self.expert_dim*expert_capacity)) + self.expert_dim * expert_capacity)) combine_tensor = self.reshape(combine_tensor, (self.dp_group, tokens_per_group, - self.expert_dim*expert_capacity)) + self.expert_dim * expert_capacity)) # combine_tensor's shape: (self.dp_group, self.expert_dim*expert_capacity, tokens_per_group) combine_tensor = self.transpose_3dim(combine_tensor, (0, 2, 1)) combine_tensor = self.cast(combine_tensor, F.dtype(expert_output)) @@ -268,27 +302,48 @@ class Router(Cell): training=True, parallel_config=None): super(Router, self).__init__() - dp = parallel_config.data_parallel - self.d_model = d_model - self.expert_dim = moe_config.expert_num - self.capacity_factor = moe_config.capacity_factor - self.num_experts_chosen = moe_config.num_experts_chosen - self.training = training - self.routing_policy = routing_policy - self.noisy_policy = None # candidate: ["jitter", "rsample", "None"] - self.noisy_epsilon = 1e-2 - self.noise = Tensor(np.random.uniform(1 - self.noisy_epsilon, 1 + self.noisy_epsilon, (d_model,))) + if _get_parallel_mode() in (ParallelMode.AUTO_PARALLEL,) and _is_sharding_propagation(): + self.d_model = d_model + self.expert_dim = moe_config.expert_num + self.capacity_factor = moe_config.capacity_factor + self.num_experts_chosen = moe_config.num_experts_chosen + self.training = training + self.routing_policy = routing_policy + self.noisy_policy = None # candidate: ["jitter", "rsample", "None"] + self.noisy_epsilon = 1e-2 + self.noise = Tensor(np.random.uniform(1 - self.noisy_epsilon, 1 + self.noisy_epsilon, (d_model,))) - self.dense = Dense(in_channels=self.d_model, out_channels=self.expert_dim, has_bias=False) - self.dense.matmul.shard(((dp, 1), (1, 1))) - self.mul = P.Mul().shard(((dp, 1, 1), (dp,))) - self.cast = P.Cast() + self.dense = Dense(in_channels=self.d_model, out_channels=self.expert_dim, has_bias=False) + self.mul = P.Mul() + self.cast = P.Cast() - if self.routing_policy is None: - self.router = TopkRouter(d_model=d_model, moe_config=moe_config, training=training, - parallel_config=parallel_config) + if self.routing_policy is None: + self.router = TopkRouter(d_model=d_model, moe_config=moe_config, training=training, + parallel_config=parallel_config) + else: + self.router = routing_policy else: - self.router = routing_policy + dp = parallel_config.data_parallel + self.d_model = d_model + self.expert_dim = moe_config.expert_num + self.capacity_factor = moe_config.capacity_factor + self.num_experts_chosen = moe_config.num_experts_chosen + self.training = training + self.routing_policy = routing_policy + self.noisy_policy = None # candidate: ["jitter", "rsample", "None"] + self.noisy_epsilon = 1e-2 + self.noise = Tensor(np.random.uniform(1 - self.noisy_epsilon, 1 + self.noisy_epsilon, (d_model,))) + + self.dense = Dense(in_channels=self.d_model, out_channels=self.expert_dim, has_bias=False) + self.dense.matmul.shard(((dp, 1), (1, 1))) + self.mul = P.Mul().shard(((dp, 1, 1), (dp,))) + self.cast = P.Cast() + + if self.routing_policy is None: + self.router = TopkRouter(d_model=d_model, moe_config=moe_config, training=training, + parallel_config=parallel_config) + else: + self.router = routing_policy def construct(self, input_tensor): input_tensor = self.cast(input_tensor, mstype.float32) @@ -326,53 +381,102 @@ class TopkRouter(Cell): training=True, parallel_config=None): super(TopkRouter, self).__init__() - dp = parallel_config.data_parallel - self.d_model = d_model - self.expert_dim = moe_config.expert_num - self.capacity_factor = moe_config.capacity_factor - self.training = training - self.dp_group = dp - self.noisy_policy = None - self.cast = P.Cast() - self.reshape = P.Reshape() - self.shape = P.Shape() - self.softmax = P.Softmax(axis=-1).shard(((dp, 1, 1,),)) - self.argmax = P.ArgMaxWithValue(axis=-1, keep_dims=False).shard(((dp, 1, 1),)) - self.num_experts_chosen = moe_config.num_experts_chosen - self.onehot = P.OneHot().shard(((dp, 1, 1), (), ())) - self.onehot2 = P.OneHot().shard(((dp, 1, 1), (), ())) - self.onehot3 = P.OneHot().shard(((dp, 1, 1, 1), (), ())) - self.on_value = Tensor(1.0, mstype.float32) - self.off_value = Tensor(0.0, mstype.float32) + if _get_parallel_mode() in (ParallelMode.AUTO_PARALLEL,) and _is_sharding_propagation(): + dp = parallel_config.data_parallel + self.d_model = d_model + self.expert_dim = moe_config.expert_num + self.capacity_factor = moe_config.capacity_factor + self.training = training + self.dp_group = dp + self.noisy_policy = None + self.cast = P.Cast() + self.reshape = P.Reshape() + self.shape = P.Shape() + self.softmax = P.Softmax(axis=-1) + self.argmax = P.ArgMaxWithValue(axis=-1, keep_dims=False) + self.num_experts_chosen = moe_config.num_experts_chosen + self.onehot = P.OneHot() + self.onehot2 = P.OneHot() + self.onehot3 = P.OneHot() + self.on_value = Tensor(1.0, mstype.float32) + self.off_value = Tensor(0.0, mstype.float32) - self.reduce_mean = P.ReduceMean(keep_dims=False).shard(((dp, 1, 1),)) - self.reduce_mean2 = P.ReduceMean(keep_dims=False).shard(((dp, 1, 1),)) - self.reduce_mean3 = P.ReduceMean(keep_dims=False).shard(((dp, 1),)) - self.mul = P.Mul().shard(((dp, 1), (dp, 1))) - self.mul2 = P.Mul().shard(((1,), ())) - self.mul3 = P.Mul().shard(((1,), ())) - self.mul4 = P.Mul().shard(((dp, 1, 1), (dp, 1, 1))) - self.mul5 = P.Mul().shard(((dp, 1, 1), (dp, 1, 1))) - self.mul6 = P.Mul().shard(((dp, 1), (dp, 1))) - self.mul7 = P.Mul().shard(((dp, 1), (dp, 1))) - self.mul8 = P.Mul().shard(((dp, 1, 1), (dp, 1, 1))) - self.mul9 = P.Mul().shard(((dp, 1, 1, 1), (dp, 1, 1, 1))) - self.not_equal = P.NotEqual().shard(((dp, 1, 1, 1), ())) - self.div1 = P.RealDiv().shard(((dp, 1, 1), (dp, 1, 1))) - self.div2 = P.RealDiv().shard(((dp, 1, 1, 1), (dp, 1, 1, 1))) - self.add = P.Add().shard(((dp, 1, 1), (dp, 1, 1))) - self.add2 = P.Add().shard(((dp, 1, 1, 1), (dp, 1, 1, 1))) - self.add3 = P.Add().shard(((dp, 1), (dp, 1))) - self.add4 = P.Add().shard(((dp, 1, 1, 1), ())) - self.sub = P.Sub().shard(((), (dp, 1, 1))) + self.reduce_mean = P.ReduceMean(keep_dims=False) + self.reduce_mean2 = P.ReduceMean(keep_dims=False) + self.reduce_mean3 = P.ReduceMean(keep_dims=False) + self.mul = P.Mul() + self.mul2 = P.Mul() + self.mul3 = P.Mul() + self.mul4 = P.Mul() + self.mul5 = P.Mul() + self.mul6 = P.Mul() + self.mul7 = P.Mul() + self.mul8 = P.Mul().shard(((dp, 1, 1), (dp, 1, 1))) + self.mul9 = P.Mul().shard(((dp, 1, 1, 1), (dp, 1, 1, 1))) + self.not_equal = P.NotEqual() + self.div1 = P.RealDiv() + self.div2 = P.RealDiv() + self.add = P.Add() + self.add2 = P.Add() + self.add3 = P.Add() + self.add4 = P.Add() + self.sub = P.Sub() - self.cumsum = P.CumSum(exclusive=True).shard(((dp, 1, 1),)) - self.less = P.Less().shard(((dp, 1, 1), ())) - self.reduce_sum = P.ReduceSum(keep_dims=False).shard(((dp, 1, 1),)) - self.reduce_sum_keep = P.ReduceSum(keep_dims=True).shard(((dp, 1, 1),)) - self.reduce_sum_keep2 = P.ReduceSum(keep_dims=True).shard(((dp, 1, 1, 1),)) - self.expand = P.ExpandDims().shard(((dp, 1),)) - self.expand2 = P.ExpandDims().shard(((dp, 1, 1),)) + self.cumsum = P.CumSum(exclusive=True) + self.less = P.Less() + self.reduce_sum = P.ReduceSum(keep_dims=False) + self.reduce_sum_keep = P.ReduceSum(keep_dims=True) + self.reduce_sum_keep2 = P.ReduceSum(keep_dims=True) + self.expand = P.ExpandDims() + self.expand2 = P.ExpandDims() + else: + dp = parallel_config.data_parallel + self.d_model = d_model + self.expert_dim = moe_config.expert_num + self.capacity_factor = moe_config.capacity_factor + self.training = training + self.dp_group = dp + self.noisy_policy = None + self.cast = P.Cast() + self.reshape = P.Reshape() + self.shape = P.Shape() + self.softmax = P.Softmax(axis=-1).shard(((dp, 1, 1,),)) + self.argmax = P.ArgMaxWithValue(axis=-1, keep_dims=False).shard(((dp, 1, 1),)) + self.num_experts_chosen = moe_config.num_experts_chosen + self.onehot = P.OneHot().shard(((dp, 1, 1), (), ())) + self.onehot2 = P.OneHot().shard(((dp, 1, 1), (), ())) + self.onehot3 = P.OneHot().shard(((dp, 1, 1, 1), (), ())) + self.on_value = Tensor(1.0, mstype.float32) + self.off_value = Tensor(0.0, mstype.float32) + + self.reduce_mean = P.ReduceMean(keep_dims=False).shard(((dp, 1, 1),)) + self.reduce_mean2 = P.ReduceMean(keep_dims=False).shard(((dp, 1, 1),)) + self.reduce_mean3 = P.ReduceMean(keep_dims=False).shard(((dp, 1),)) + self.mul = P.Mul().shard(((dp, 1), (dp, 1))) + self.mul2 = P.Mul().shard(((1,), ())) + self.mul3 = P.Mul().shard(((1,), ())) + self.mul4 = P.Mul().shard(((dp, 1, 1), (dp, 1, 1))) + self.mul5 = P.Mul().shard(((dp, 1, 1), (dp, 1, 1))) + self.mul6 = P.Mul().shard(((dp, 1), (dp, 1))) + self.mul7 = P.Mul().shard(((dp, 1), (dp, 1))) + self.mul8 = P.Mul().shard(((dp, 1, 1), (dp, 1, 1))) + self.mul9 = P.Mul().shard(((dp, 1, 1, 1), (dp, 1, 1, 1))) + self.not_equal = P.NotEqual().shard(((dp, 1, 1, 1), ())) + self.div1 = P.RealDiv().shard(((dp, 1, 1), (dp, 1, 1))) + self.div2 = P.RealDiv().shard(((dp, 1, 1, 1), (dp, 1, 1, 1))) + self.add = P.Add().shard(((dp, 1, 1), (dp, 1, 1))) + self.add2 = P.Add().shard(((dp, 1, 1, 1), (dp, 1, 1, 1))) + self.add3 = P.Add().shard(((dp, 1), (dp, 1))) + self.add4 = P.Add().shard(((dp, 1, 1, 1), ())) + self.sub = P.Sub().shard(((), (dp, 1, 1))) + + self.cumsum = P.CumSum(exclusive=True).shard(((dp, 1, 1),)) + self.less = P.Less().shard(((dp, 1, 1), ())) + self.reduce_sum = P.ReduceSum(keep_dims=False).shard(((dp, 1, 1),)) + self.reduce_sum_keep = P.ReduceSum(keep_dims=True).shard(((dp, 1, 1),)) + self.reduce_sum_keep2 = P.ReduceSum(keep_dims=True).shard(((dp, 1, 1, 1),)) + self.expand = P.ExpandDims().shard(((dp, 1),)) + self.expand2 = P.ExpandDims().shard(((dp, 1, 1),)) def _auxiliary_loss(self, expert_mask, router_prob): """ @@ -432,10 +536,10 @@ class TopkRouter(Cell): expert_index, expert_gate = self.argmax(router_prob) # expert_mask's shape: (dp_group, tokens_per_group, self.expert_dim) expert_mask = self.onehot(expert_index, self.expert_dim, self.on_value, self.off_value) - #renormalize the rest prob to be of sum 1 + # renormalize the rest prob to be of sum 1 router_prob_normal = self.div1(router_prob, self.add(self.reduce_sum_keep(router_prob, -1), 1e-9)) - #the balance loss is computed at each routing step + # the balance loss is computed at each routing step loss += self._auxiliary_loss(expert_mask, router_prob_normal) output = self._maskout_overflowed_tokens(expert_mask, expert_capacity, expert_gate, @@ -456,7 +560,7 @@ class TopkRouter(Cell): self.on_value, self.off_value)) accum_combine_tensor = self.add2(accum_combine_tensor, combine_tensor) - #expert weights normalization + # expert weights normalization combine_tensor_sum = self.reduce_sum_keep2(self.reduce_sum_keep2(accum_combine_tensor, -1), -2) accum_combine_tensor = self.div2(accum_combine_tensor, self.add4(combine_tensor_sum, 1e-9)) # dispatch_tensor is of boolean type. Here, using NotEqual instead of Cast, for that 'Cast to bool' has diff --git a/mindspore/python/mindspore/nn/transformer/transformer.py b/mindspore/python/mindspore/nn/transformer/transformer.py index efc46687c9..ca54172f46 100644 --- a/mindspore/python/mindspore/nn/transformer/transformer.py +++ b/mindspore/python/mindspore/nn/transformer/transformer.py @@ -29,7 +29,7 @@ from mindspore.ops import functional as F from mindspore.nn.cell import Cell from mindspore._checkparam import Validator from mindspore import log as logger -from mindspore.parallel._utils import _get_parallel_mode +from mindspore.parallel._utils import _get_parallel_mode, _is_sharding_propagation from mindspore.context import ParallelMode from .layers import _LayerNorm, _Linear, _Dropout, _check_input_shape, \ _args_type_validator_check, _valid_type_checks, _valid_value_checks, \ @@ -405,67 +405,115 @@ class FeedForward(Cell): param_init_type=mstype.float32, parallel_config=default_dpmp_config): super(FeedForward, self).__init__() - _check_config(parallel_config) - mp = parallel_config.model_parallel - if expert_num > 1: - ep = parallel_config.expert_parallel - else: - ep = 1 - # ffn use less dp than other ops when use_moe, due to there are ops use dp and ep. - dp = int(parallel_config.data_parallel / ep) - if ffn_hidden_size % mp != 0: - raise ValueError("For 'FeedForward', the class variable 'ffn_hidden_size' must be a multiple of the" - "num of " - "model parallel, but got the ffn_hidden_size is {} and the num of model parallel is {}." - .format(ffn_hidden_size, mp)) - if hidden_size % mp != 0: - raise ValueError("For 'FeedForward', the class variable 'hidden_size' must be a multiple of the num of " - "model parallel, but got the hidden_size is {} and the num of model parallel is {}." - .format(hidden_size, mp)) - if dropout_rate < 0 or dropout_rate >= 1: - raise ValueError("For 'FeedForward', the class variable 'dropout_rate' must be in the range [0, 1.0), " - "but got the value : {}.".format(dropout_rate)) - input_size = hidden_size - output_size = ffn_hidden_size + if _get_parallel_mode() in (ParallelMode.AUTO_PARALLEL,) and _is_sharding_propagation(): + _check_config(parallel_config) + mp = parallel_config.model_parallel + if expert_num > 1: + ep = parallel_config.expert_parallel + else: + ep = 1 + # ffn use less dp than other ops when use_moe, due to there are ops use dp and ep. + dp = int(parallel_config.data_parallel / ep) + if ffn_hidden_size % mp != 0: + raise ValueError("For 'FeedForward', the class variable 'ffn_hidden_size' must be a multiple of the" + "num of model parallel, but got the ffn_hidden_size is {} and the num of model " + "parallel is {}.".format(ffn_hidden_size, mp)) + if hidden_size % mp != 0: + raise ValueError("For 'FeedForward', the class variable 'hidden_size' must be a multiple of the num of " + "model parallel, but got the hidden_size is {} and the num of model parallel is {}." + .format(hidden_size, mp)) + if dropout_rate < 0 or dropout_rate >= 1: + raise ValueError("For 'FeedForward', the class variable 'dropout_rate' must be in the range [0, 1.0), " + "but got the value : {}.".format(dropout_rate)) + input_size = hidden_size + output_size = ffn_hidden_size - # Project to ffn_hidden_size - self.mapping = _Linear(in_channels=input_size, - out_channels=output_size, - activation=hidden_act, - transpose_b=False, - expert_num=expert_num, - outer_batch=dp, - param_init_type=param_init_type) + # Project to ffn_hidden_size + self.mapping = _Linear(in_channels=input_size, + out_channels=output_size, + activation=hidden_act, + transpose_b=False, + expert_num=expert_num, + outer_batch=dp, + param_init_type=param_init_type) - if expert_num > 1: - self.mapping.shard(strategy_matmul=((dp, ep, 1, 1), (ep, 1, mp)), - strategy_bias=((dp, ep, 1, mp), (mp,)), - strategy_activation=((dp, ep, 1, mp),)) + # Project back to hidden_size + self.projection = _Linear(in_channels=output_size, + out_channels=input_size, + transpose_b=False, + expert_num=expert_num, + outer_batch=dp, + param_init_type=param_init_type) + if expert_num > 1: + self.projection.shard(strategy_matmul=((dp, ep, 1, mp), (ep, mp, 1))) + else: + self.projection.shard(strategy_matmul=((dp, mp), (mp, 1))) + self.projection.bias.parallel_optimizer = False + self.dropout = _Dropout(1 - dropout_rate) + self.dropout_3d = _Dropout(1 - dropout_rate) + self.dropout_4d = _Dropout(1 - dropout_rate) + self.cast = P.Cast() else: - self.mapping.shard(strategy_matmul=((dp, 1), (1, mp)), - strategy_bias=((dp, mp), (mp,)), - strategy_activation=((dp, mp),)) - # Project back to hidden_size - self.projection = _Linear(in_channels=output_size, - out_channels=input_size, - transpose_b=False, - expert_num=expert_num, - outer_batch=dp, - param_init_type=param_init_type) - if expert_num > 1: - self.projection.shard(strategy_matmul=((dp, ep, 1, mp), (ep, mp, 1)), - strategy_bias=((dp, ep, 1, 1), (1,))) - else: - self.projection.shard(strategy_matmul=((dp, mp), (mp, 1)), - strategy_bias=((dp, 1), (1,))) - self.projection.bias.parallel_optimizer = False - self.dropout = _Dropout(1 - dropout_rate) - self.dropout.shard(((dp, 1),)) - self.dropout_3d = _Dropout(1 - dropout_rate) - self.dropout_3d.shard(((dp, 1, 1),)) - self.dropout_4d = _Dropout(1 - dropout_rate) - self.dropout_4d.shard(((dp, ep, 1, 1),)) - self.cast = P.Cast() + _check_config(parallel_config) + mp = parallel_config.model_parallel + if expert_num > 1: + ep = parallel_config.expert_parallel + else: + ep = 1 + # ffn use less dp than other ops when use_moe, due to there are ops use dp and ep. + dp = int(parallel_config.data_parallel / ep) + if ffn_hidden_size % mp != 0: + raise ValueError("For 'FeedForward', the class variable 'ffn_hidden_size' must be a multiple of the" + "num of model parallel, but got the ffn_hidden_size is {} and the num of model " + "parallel is {}.".format(ffn_hidden_size, mp)) + if hidden_size % mp != 0: + raise ValueError("For 'FeedForward', the class variable 'hidden_size' must be a multiple of the num of " + "model parallel, but got the hidden_size is {} and the num of model parallel is {}." + .format(hidden_size, mp)) + if dropout_rate < 0 or dropout_rate >= 1: + raise ValueError("For 'FeedForward', the class variable 'dropout_rate' must be in the range [0, 1.0), " + "but got the value : {}.".format(dropout_rate)) + input_size = hidden_size + output_size = ffn_hidden_size + + # Project to ffn_hidden_size + self.mapping = _Linear(in_channels=input_size, + out_channels=output_size, + activation=hidden_act, + transpose_b=False, + expert_num=expert_num, + outer_batch=dp, + param_init_type=param_init_type) + + if expert_num > 1: + self.mapping.shard(strategy_matmul=((dp, ep, 1, 1), (ep, 1, mp)), + strategy_bias=((dp, ep, 1, mp), (mp,)), + strategy_activation=((dp, ep, 1, mp),)) + else: + self.mapping.shard(strategy_matmul=((dp, 1), (1, mp)), + strategy_bias=((dp, mp), (mp,)), + strategy_activation=((dp, mp),)) + # Project back to hidden_size + self.projection = _Linear(in_channels=output_size, + out_channels=input_size, + transpose_b=False, + expert_num=expert_num, + outer_batch=dp, + param_init_type=param_init_type) + if expert_num > 1: + self.projection.shard(strategy_matmul=((dp, ep, 1, mp), (ep, mp, 1)), + strategy_bias=((dp, ep, 1, 1), (1,))) + else: + self.projection.shard(strategy_matmul=((dp, mp), (mp, 1)), + strategy_bias=((dp, 1), (1,))) + self.projection.bias.parallel_optimizer = False + self.dropout = _Dropout(1 - dropout_rate) + self.dropout.shard(((dp, 1),)) + self.dropout_3d = _Dropout(1 - dropout_rate) + self.dropout_3d.shard(((dp, 1, 1),)) + self.dropout_4d = _Dropout(1 - dropout_rate) + self.dropout_4d.shard(((dp, ep, 1, 1),)) + self.cast = P.Cast() def construct(self, x): _check_input_shape(F.shape(x), "x", self.cls_name, [2, 3]) @@ -795,118 +843,216 @@ class MultiHeadAttention(Cell): use_past=False, parallel_config=default_dpmp_config): super(MultiHeadAttention, self).__init__() - _check_config(parallel_config) - self.is_parallel_mode = _get_parallel_mode() in (ParallelMode.SEMI_AUTO_PARALLEL, ParallelMode.AUTO_PARALLEL) - self.src_seq_length = src_seq_length - self.tgt_seq_length = tgt_seq_length - self.hidden_size = hidden_size - self.batch_size = batch_size - if hidden_dropout_rate < 0 or hidden_dropout_rate >= 1: - raise ValueError("For 'MultiHeadAttention', the class variable 'hidden_dropout_rate' must be " - "in range [0, 1.0), but got the value : {}.".format(hidden_dropout_rate)) - if attention_dropout_rate < 0 or attention_dropout_rate >= 1: - raise ValueError("For 'MultiHeadAttention', the class variable 'attention_dropout_rate' must be " - "in range [0, 1.0), but got the value : {}.".format(attention_dropout_rate)) - if hidden_size % num_heads != 0: - raise ValueError("For 'MultiHeadAttention', the class variable 'hidden_size' should be a multiple " - "of 'num_heads', but got the hidden_size is {} and the num_heads is {}." - .format(hidden_size, num_heads)) - if num_heads % parallel_config.model_parallel != 0: - raise ValueError("For 'MultiHeadAttention', the class variable 'num_heads' must be a multiple of " - "'parallel_config.model_parallel', but got the num_heads is {} " - "and the parallel_config.model_parallel is {}." - .format(num_heads, parallel_config.model_parallel)) - if self.is_parallel_mode and batch_size % parallel_config.data_parallel != 0: - raise ValueError("For 'MultiHeadAttention', the class variable 'batch_size' must be a multiple of " - "'parallel_config.data_parallel', but got the batch_size is {} " - "and the parallel_config.data_parallel is {}." - .format(batch_size, parallel_config.data_parallel)) - self.is_first_iteration = True - # Output layer - self.projection = _Linear(in_channels=hidden_size, - out_channels=hidden_size, - transpose_b=False, + if _get_parallel_mode() in (ParallelMode.AUTO_PARALLEL,) and _is_sharding_propagation(): + _check_config(parallel_config) + self.is_parallel_mode = _get_parallel_mode() in ( + ParallelMode.SEMI_AUTO_PARALLEL, ParallelMode.AUTO_PARALLEL) + self.src_seq_length = src_seq_length + self.tgt_seq_length = tgt_seq_length + self.hidden_size = hidden_size + self.batch_size = batch_size + if hidden_dropout_rate < 0 or hidden_dropout_rate >= 1: + raise ValueError("For 'MultiHeadAttention', the class variable 'hidden_dropout_rate' must be " + "in range [0, 1.0), but got the value : {}.".format(hidden_dropout_rate)) + if attention_dropout_rate < 0 or attention_dropout_rate >= 1: + raise ValueError("For 'MultiHeadAttention', the class variable 'attention_dropout_rate' must be " + "in range [0, 1.0), but got the value : {}.".format(attention_dropout_rate)) + if hidden_size % num_heads != 0: + raise ValueError("For 'MultiHeadAttention', the class variable 'hidden_size' should be a multiple " + "of 'num_heads', but got the hidden_size is {} and the num_heads is {}." + .format(hidden_size, num_heads)) + if num_heads % parallel_config.model_parallel != 0: + raise ValueError("For 'MultiHeadAttention', the class variable 'num_heads' must be a multiple of " + "'parallel_config.model_parallel', but got the num_heads is {} " + "and the parallel_config.model_parallel is {}." + .format(num_heads, parallel_config.model_parallel)) + if self.is_parallel_mode and batch_size % parallel_config.data_parallel != 0: + raise ValueError("For 'MultiHeadAttention', the class variable 'batch_size' must be a multiple of " + "'parallel_config.data_parallel', but got the batch_size is {} " + "and the parallel_config.data_parallel is {}." + .format(batch_size, parallel_config.data_parallel)) + self.is_first_iteration = True + # Output layer + self.projection = _Linear(in_channels=hidden_size, + out_channels=hidden_size, + transpose_b=False, + param_init_type=param_init_type).to_float(compute_dtype) + self.projection.shard(strategy_bias=((parallel_config.data_parallel, 1), (1,)), + strategy_matmul=((parallel_config.data_parallel, parallel_config.model_parallel), + (parallel_config.model_parallel, 1))) + self.projection.bias.parallel_optimizer = False + self.transpose = P.Transpose() + self.merger_head_transpose = P.Transpose() + self.reshape = P.Reshape() + self.n_head = num_heads + # embedding size per head + self.size_per_head = hidden_size // self.n_head + self.concat_k = P.Concat(axis=3) + self.concat_v = P.Concat(axis=2) + self.multiply_data = Tensor([ + -10000.0, + ], dtype=softmax_compute_type) + self.batch_matmul = P.BatchMatMul() + self.real_div = P.RealDiv() + self.sub = P.Sub() + self.mul = P.Mul() + self.add = P.Add() + # Normalize factor for attention, sqrt(dk) as widely used + self.scale_factor = Tensor(math.sqrt(math.sqrt(self.size_per_head))) + self.use_past = use_past + self.dropout = _Dropout(1 - hidden_dropout_rate) + self.prob_dropout = _Dropout(1 - attention_dropout_rate) + self.softmax = nn.Softmax().to_float(softmax_compute_type) + self.expand_dims = P.ExpandDims() + + # Query + self.dense1 = _Linear(hidden_size, + hidden_size, + param_init_type=param_init_type).to_float(compute_dtype) + # Key + self.dense2 = _Linear(hidden_size, + hidden_size, + param_init_type=param_init_type).to_float(compute_dtype) + # Value + self.dense3 = _Linear(hidden_size, + hidden_size, param_init_type=param_init_type).to_float(compute_dtype) - self.projection.shard(strategy_bias=((parallel_config.data_parallel, 1), (1,)), - strategy_matmul=((parallel_config.data_parallel, parallel_config.model_parallel), - (parallel_config.model_parallel, 1))) - self.projection.bias.parallel_optimizer = False - self.transpose = P.Transpose().shard(((parallel_config.data_parallel, 1, parallel_config.model_parallel, 1),)) - self.merger_head_transpose = P.Transpose().shard( - ((parallel_config.data_parallel, parallel_config.model_parallel, 1, 1),)) - self.reshape = P.Reshape() - self.n_head = num_heads - # embedding size per head - self.size_per_head = hidden_size // self.n_head - self.concat_k = P.Concat(axis=3) - self.concat_v = P.Concat(axis=2) - self.multiply_data = Tensor([ - -10000.0, - ], dtype=softmax_compute_type) - self.batch_matmul = P.BatchMatMul().shard( - ((parallel_config.data_parallel, parallel_config.model_parallel, 1, 1), - (parallel_config.data_parallel, parallel_config.model_parallel, 1, 1))) - self.real_div = P.RealDiv().shard(((parallel_config.data_parallel, parallel_config.model_parallel, 1, 1), ())) - self.sub = P.Sub().shard( - ((1,), (parallel_config.data_parallel, 1, 1, 1))) - self.mul = P.Mul().shard( - ((parallel_config.data_parallel, 1, 1, 1), (1,))) - self.add = P.Add().shard( - ((parallel_config.data_parallel, 1, 1, 1), - (parallel_config.data_parallel, parallel_config.model_parallel, 1, 1))) - # Normalize factor for attention, sqrt(dk) as widely used - self.scale_factor = Tensor(math.sqrt(math.sqrt(self.size_per_head))) - self.use_past = use_past - self.dropout = _Dropout(1 - hidden_dropout_rate) - self.dropout.shard(((parallel_config.data_parallel, 1),)) - self.prob_dropout = _Dropout(1 - attention_dropout_rate) - self.prob_dropout.shard( - ((parallel_config.data_parallel, parallel_config.model_parallel, 1, 1),)) - self.softmax = nn.Softmax().to_float(softmax_compute_type) - self.softmax.softmax.shard(((parallel_config.data_parallel, parallel_config.model_parallel, 1),)) - self.expand_dims = P.ExpandDims().shard(((parallel_config.data_parallel, 1, 1),)) - # Query - self.dense1 = _Linear(hidden_size, - hidden_size, - param_init_type=param_init_type).to_float(compute_dtype) - self.dense1.shard(strategy_matmul=((parallel_config.data_parallel, 1), (parallel_config.model_parallel, 1)), - strategy_bias=((parallel_config.data_parallel, parallel_config.model_parallel), - (parallel_config.model_parallel,))) - # Key - self.dense2 = _Linear(hidden_size, - hidden_size, - param_init_type=param_init_type).to_float(compute_dtype) - self.dense2.shard(strategy_matmul=((parallel_config.data_parallel, 1), (parallel_config.model_parallel, 1)), - strategy_bias=((parallel_config.data_parallel, parallel_config.model_parallel), - (parallel_config.model_parallel,))) + self.dtype = compute_dtype + self.softmax_dtype = softmax_compute_type + if self.use_past: + # operators used for state reuse + seq_range = np.arange(src_seq_length).reshape(1, 1, -1) + self.range = Tensor(np.tile(seq_range, (batch_size, 1, 1)), mstype.int32) + self.seq_length = src_seq_length + self.attention_mask = Tensor(np.tril(np.ones(shape=(self.seq_length, self.seq_length))), mstype.int32) + self.slice = P.StridedSlice().shard(((1, 1, 1, 1),)) + self.not_equal = P.NotEqual().shard(((1, 1, 1, 1), ())) + self.reducesum = P.ReduceSum().shard(((1, 1, 1, 1),)) + self.expand_dims = P.ExpandDims().shard(((1, 1, 1),)) + self.tensor_le = P.LessEqual().shard(((1, 1, 1), (1, 1, 1))) + self.add = P.Add().shard(((1, 1, 1, 1), (1, 1, 1, 1))) + self.equal = P.Equal().shard(((1, 1, 1), (1, 1, 1))) + self.sub1 = P.Sub().shard(((1,), ())) + self.tile = P.Tile().shard(((1, 1, 1, 1),)) + self.less = P.Less().shard(((1, 1, 1), (1, 1, 1))) + self.mul1 = P.Mul().shard(((1, 1, 1, 1), (1, 1, 1, 1))) + else: + _check_config(parallel_config) + self.is_parallel_mode = _get_parallel_mode() in ( + ParallelMode.SEMI_AUTO_PARALLEL, ParallelMode.AUTO_PARALLEL) + self.src_seq_length = src_seq_length + self.tgt_seq_length = tgt_seq_length + self.hidden_size = hidden_size + self.batch_size = batch_size + if hidden_dropout_rate < 0 or hidden_dropout_rate >= 1: + raise ValueError("For 'MultiHeadAttention', the class variable 'hidden_dropout_rate' must be " + "in range [0, 1.0), but got the value : {}.".format(hidden_dropout_rate)) + if attention_dropout_rate < 0 or attention_dropout_rate >= 1: + raise ValueError("For 'MultiHeadAttention', the class variable 'attention_dropout_rate' must be " + "in range [0, 1.0), but got the value : {}.".format(attention_dropout_rate)) + if hidden_size % num_heads != 0: + raise ValueError("For 'MultiHeadAttention', the class variable 'hidden_size' should be a multiple " + "of 'num_heads', but got the hidden_size is {} and the num_heads is {}." + .format(hidden_size, num_heads)) + if num_heads % parallel_config.model_parallel != 0: + raise ValueError("For 'MultiHeadAttention', the class variable 'num_heads' must be a multiple of " + "'parallel_config.model_parallel', but got the num_heads is {} " + "and the parallel_config.model_parallel is {}." + .format(num_heads, parallel_config.model_parallel)) + if self.is_parallel_mode and batch_size % parallel_config.data_parallel != 0: + raise ValueError("For 'MultiHeadAttention', the class variable 'batch_size' must be a multiple of " + "'parallel_config.data_parallel', but got the batch_size is {} " + "and the parallel_config.data_parallel is {}." + .format(batch_size, parallel_config.data_parallel)) + self.is_first_iteration = True + # Output layer + self.projection = _Linear(in_channels=hidden_size, + out_channels=hidden_size, + transpose_b=False, + param_init_type=param_init_type).to_float(compute_dtype) + self.projection.shard(strategy_bias=((parallel_config.data_parallel, 1), (1,)), + strategy_matmul=((parallel_config.data_parallel, parallel_config.model_parallel), + (parallel_config.model_parallel, 1))) + self.projection.bias.parallel_optimizer = False + self.transpose = P.Transpose().shard( + ((parallel_config.data_parallel, 1, parallel_config.model_parallel, 1),)) + self.merger_head_transpose = P.Transpose().shard( + ((parallel_config.data_parallel, parallel_config.model_parallel, 1, 1),)) + self.reshape = P.Reshape() + self.n_head = num_heads + # embedding size per head + self.size_per_head = hidden_size // self.n_head + self.concat_k = P.Concat(axis=3) + self.concat_v = P.Concat(axis=2) + self.multiply_data = Tensor([ + -10000.0, + ], dtype=softmax_compute_type) + self.batch_matmul = P.BatchMatMul().shard( + ((parallel_config.data_parallel, parallel_config.model_parallel, 1, 1), + (parallel_config.data_parallel, parallel_config.model_parallel, 1, 1))) + self.real_div = P.RealDiv().shard( + ((parallel_config.data_parallel, parallel_config.model_parallel, 1, 1), ())) + self.sub = P.Sub().shard( + ((1,), (parallel_config.data_parallel, 1, 1, 1))) + self.mul = P.Mul().shard( + ((parallel_config.data_parallel, 1, 1, 1), (1,))) + self.add = P.Add().shard( + ((parallel_config.data_parallel, 1, 1, 1), + (parallel_config.data_parallel, parallel_config.model_parallel, 1, 1))) + # Normalize factor for attention, sqrt(dk) as widely used + self.scale_factor = Tensor(math.sqrt(math.sqrt(self.size_per_head))) + self.use_past = use_past + self.dropout = _Dropout(1 - hidden_dropout_rate) + self.dropout.shard(((parallel_config.data_parallel, 1),)) + self.prob_dropout = _Dropout(1 - attention_dropout_rate) + self.prob_dropout.shard( + ((parallel_config.data_parallel, parallel_config.model_parallel, 1, 1),)) + self.softmax = nn.Softmax().to_float(softmax_compute_type) + self.softmax.softmax.shard(((parallel_config.data_parallel, parallel_config.model_parallel, 1),)) + self.expand_dims = P.ExpandDims().shard(((parallel_config.data_parallel, 1, 1),)) - # Value - self.dense3 = _Linear(hidden_size, - hidden_size, - param_init_type=param_init_type).to_float(compute_dtype) - self.dense3.shard(strategy_matmul=((parallel_config.data_parallel, 1), (parallel_config.model_parallel, 1)), - strategy_bias=((parallel_config.data_parallel, parallel_config.model_parallel), - (parallel_config.model_parallel,))) - self.dtype = compute_dtype - self.softmax_dtype = softmax_compute_type - if self.use_past: - # operators used for state reuse - seq_range = np.arange(src_seq_length).reshape(1, 1, -1) - self.range = Tensor(np.tile(seq_range, (batch_size, 1, 1)), mstype.int32) - self.seq_length = src_seq_length - self.attention_mask = Tensor(np.tril(np.ones(shape=(self.seq_length, self.seq_length))), mstype.int32) - self.slice = P.StridedSlice().shard(((1, 1, 1, 1),)) - self.not_equal = P.NotEqual().shard(((1, 1, 1, 1), ())) - self.reducesum = P.ReduceSum().shard(((1, 1, 1, 1),)) - self.expand_dims = P.ExpandDims().shard(((1, 1, 1),)) - self.tensor_le = P.LessEqual().shard(((1, 1, 1), (1, 1, 1))) - self.add = P.Add().shard(((1, 1, 1, 1), (1, 1, 1, 1))) - self.equal = P.Equal().shard(((1, 1, 1), (1, 1, 1))) - self.sub1 = P.Sub().shard(((1,), ())) - self.tile = P.Tile().shard(((1, 1, 1, 1),)) - self.less = P.Less().shard(((1, 1, 1), (1, 1, 1))) - self.mul1 = P.Mul().shard(((1, 1, 1, 1), (1, 1, 1, 1))) + # Query + self.dense1 = _Linear(hidden_size, + hidden_size, + param_init_type=param_init_type).to_float(compute_dtype) + self.dense1.shard(strategy_matmul=((parallel_config.data_parallel, 1), (parallel_config.model_parallel, 1)), + strategy_bias=((parallel_config.data_parallel, parallel_config.model_parallel), + (parallel_config.model_parallel,))) + # Key + self.dense2 = _Linear(hidden_size, + hidden_size, + param_init_type=param_init_type).to_float(compute_dtype) + self.dense2.shard(strategy_matmul=((parallel_config.data_parallel, 1), (parallel_config.model_parallel, 1)), + strategy_bias=((parallel_config.data_parallel, parallel_config.model_parallel), + (parallel_config.model_parallel,))) + + # Value + self.dense3 = _Linear(hidden_size, + hidden_size, + param_init_type=param_init_type).to_float(compute_dtype) + self.dense3.shard(strategy_matmul=((parallel_config.data_parallel, 1), (parallel_config.model_parallel, 1)), + strategy_bias=((parallel_config.data_parallel, parallel_config.model_parallel), + (parallel_config.model_parallel,))) + self.dtype = compute_dtype + self.softmax_dtype = softmax_compute_type + if self.use_past: + # operators used for state reuse + seq_range = np.arange(src_seq_length).reshape(1, 1, -1) + self.range = Tensor(np.tile(seq_range, (batch_size, 1, 1)), mstype.int32) + self.seq_length = src_seq_length + self.attention_mask = Tensor(np.tril(np.ones(shape=(self.seq_length, self.seq_length))), mstype.int32) + self.slice = P.StridedSlice().shard(((1, 1, 1, 1),)) + self.not_equal = P.NotEqual().shard(((1, 1, 1, 1), ())) + self.reducesum = P.ReduceSum().shard(((1, 1, 1, 1),)) + self.expand_dims = P.ExpandDims().shard(((1, 1, 1),)) + self.tensor_le = P.LessEqual().shard(((1, 1, 1), (1, 1, 1))) + self.add = P.Add().shard(((1, 1, 1, 1), (1, 1, 1, 1))) + self.equal = P.Equal().shard(((1, 1, 1), (1, 1, 1))) + self.sub1 = P.Sub().shard(((1,), ())) + self.tile = P.Tile().shard(((1, 1, 1, 1),)) + self.less = P.Less().shard(((1, 1, 1), (1, 1, 1))) + self.mul1 = P.Mul().shard(((1, 1, 1, 1), (1, 1, 1, 1))) def construct(self, query_tensor, key_tensor, value_tensor, attention_mask, key_past=None, value_past=None, batch_valid_length=None): @@ -1272,79 +1418,163 @@ class TransformerEncoderLayer(Cell): moe_config=default_moe_config, parallel_config=default_dpmp_config): super(TransformerEncoderLayer, self).__init__() - _check_config(parallel_config) - if num_heads % parallel_config.model_parallel != 0: - raise ValueError("For 'TransformerEncoderLayer', the class variable 'num_heads' must be divisibled by the " - "'parallel_config.model_parallel', but got the num_heads is {} and " - "parallel_config.model_parallel is {}.".format(num_heads, parallel_config.model_parallel)) - if hidden_size % parallel_config.model_parallel != 0: - raise ValueError("For 'TransformerEncoderLayer', the class variable 'hidden_size' must be divisibled by " - "the 'parallel_config.model_parallel', but got the hidden_size is {} and parallel_config." - " model_parallel is {}.".format(hidden_size, parallel_config.model_parallel)) - if ffn_hidden_size % parallel_config.model_parallel != 0: - raise ValueError("For 'TransformerEncoderLayer', the class variable 'ffn_hidden_size' must be divisibled " - "by the 'parallel_config.model_parallel', but got the ffn_hidden_size is {} " - "and parallel_config. model_parallel is {}." - .format(ffn_hidden_size, parallel_config.model_parallel)) - _check_moe_config(moe_config, parallel_config) - self.use_moe = (moe_config.expert_num > 1) - self.use_past = use_past - self.seq_length = seq_length - self.hidden_size = hidden_size - self.batch_size = batch_size - self.layernorm1 = _LayerNorm((hidden_size,)).to_float(layernorm_compute_type) - self.layernorm1.shard(((parallel_config.data_parallel, 1),)) - self.layernorm2 = _LayerNorm((hidden_size,)).to_float(layernorm_compute_type) - self.layernorm2.shard(((parallel_config.data_parallel, 1),)) + if _get_parallel_mode() in (ParallelMode.AUTO_PARALLEL,) and _is_sharding_propagation(): + _check_config(parallel_config) + if num_heads % parallel_config.model_parallel != 0: + raise ValueError( + "For 'TransformerEncoderLayer', the class variable 'num_heads' must be divisibled by the " + "'parallel_config.model_parallel', but got the num_heads is {} and " + "parallel_config.model_parallel is {}.".format(num_heads, parallel_config.model_parallel)) + if hidden_size % parallel_config.model_parallel != 0: + raise ValueError( + "For 'TransformerEncoderLayer', the class variable 'hidden_size' must be divisibled by " + "the 'parallel_config.model_parallel', but got the hidden_size is {} and parallel_config." + " model_parallel is {}.".format(hidden_size, parallel_config.model_parallel)) + if ffn_hidden_size % parallel_config.model_parallel != 0: + raise ValueError( + "For 'TransformerEncoderLayer', the class variable 'ffn_hidden_size' must be divisibled " + "by the 'parallel_config.model_parallel', but got the ffn_hidden_size is {} " + "and parallel_config. model_parallel is {}." + .format(ffn_hidden_size, parallel_config.model_parallel)) + _check_moe_config(moe_config, parallel_config) + self.use_moe = (moe_config.expert_num > 1) + self.use_past = use_past + self.seq_length = seq_length + self.hidden_size = hidden_size + self.batch_size = batch_size + self.layernorm1 = _LayerNorm((hidden_size,)).to_float(layernorm_compute_type) + self.layernorm2 = _LayerNorm((hidden_size,)).to_float(layernorm_compute_type) - self.attention = MultiHeadAttention(batch_size=batch_size, - src_seq_length=seq_length, - tgt_seq_length=seq_length, - hidden_size=hidden_size, - num_heads=num_heads, - hidden_dropout_rate=hidden_dropout_rate, - attention_dropout_rate=attention_dropout_rate, - softmax_compute_type=softmax_compute_type, - param_init_type=param_init_type, - use_past=use_past, - parallel_config=parallel_config.dpmp if self.use_moe else parallel_config) - if self.use_moe: - self.output = MoE(hidden_size=hidden_size, - dropout_rate=hidden_dropout_rate, - ffn_hidden_size=ffn_hidden_size, - param_init_type=param_init_type, - hidden_act=hidden_act, - moe_config=moe_config, - parallel_config=parallel_config) + self.attention = MultiHeadAttention(batch_size=batch_size, + src_seq_length=seq_length, + tgt_seq_length=seq_length, + hidden_size=hidden_size, + num_heads=num_heads, + hidden_dropout_rate=hidden_dropout_rate, + attention_dropout_rate=attention_dropout_rate, + softmax_compute_type=softmax_compute_type, + param_init_type=param_init_type, + use_past=use_past, + parallel_config=parallel_config.dpmp if self.use_moe + else parallel_config) + if self.use_moe: + self.output = MoE(hidden_size=hidden_size, + dropout_rate=hidden_dropout_rate, + ffn_hidden_size=ffn_hidden_size, + param_init_type=param_init_type, + hidden_act=hidden_act, + moe_config=moe_config, + parallel_config=parallel_config) + else: + # Feed Forward Network, FFN + self.output = FeedForward(hidden_size=hidden_size, + dropout_rate=hidden_dropout_rate, + ffn_hidden_size=ffn_hidden_size, + param_init_type=param_init_type, + hidden_act=hidden_act, + parallel_config=parallel_config) + self.post_layernorm_residual = post_layernorm_residual + self.add = P.Add().shard(((parallel_config.data_parallel, 1), (parallel_config.data_parallel, 1))) + self.add_3d = P.Add().shard(((parallel_config.data_parallel, 1, 1), (parallel_config.data_parallel, 1, 1))) + self.dtype = mstype.float16 + self.key_past = None + self.value_past = None + + if self.use_past: + # operator used for state reuse + self.reducesum = P.ReduceSum().shard(((1, 1, 1, 1),)) + self.not_equal = P.NotEqual().shard(((1, 1, 1, 1), ())) + self.slice = P.StridedSlice().shard(((1, 1, 1, 1),)) + size_per_head = int(hidden_size / num_heads) + self.key_shape = (batch_size, num_heads, size_per_head, seq_length) + self.value_shape = (batch_size, num_heads, seq_length, size_per_head) + # parameters saving key and value states + self.key_past = Parameter(Tensor(np.zeros(shape=self.key_shape), self.dtype), name="key_past") + self.value_past = Parameter(Tensor(np.zeros(shape=self.value_shape), self.dtype), name="value_past") + self.tile = P.Tile().shard(((1, 1),)) + self.mul = P.Mul().shard(((1, 1, 1, 1), (1,))) + self.assign = P.Assign().shard(((1, 1, 1, 1), (1, 1, 1, 1))) + elif _get_parallel_mode() not in (ParallelMode.AUTO_PARALLEL,): + _check_config(parallel_config) + if num_heads % parallel_config.model_parallel != 0: + raise ValueError( + "For 'TransformerEncoderLayer', the class variable 'num_heads' must be divisibled by the " + "'parallel_config.model_parallel', but got the num_heads is {} and " + "parallel_config.model_parallel is {}.".format(num_heads, parallel_config.model_parallel)) + if hidden_size % parallel_config.model_parallel != 0: + raise ValueError( + "For 'TransformerEncoderLayer', the class variable 'hidden_size' must be divisibled by " + "the 'parallel_config.model_parallel', but got the hidden_size is {} and parallel_config." + " model_parallel is {}.".format(hidden_size, parallel_config.model_parallel)) + if ffn_hidden_size % parallel_config.model_parallel != 0: + raise ValueError( + "For 'TransformerEncoderLayer', the class variable 'ffn_hidden_size' must be divisibled " + "by the 'parallel_config.model_parallel', but got the ffn_hidden_size is {} " + "and parallel_config. model_parallel is {}." + .format(ffn_hidden_size, parallel_config.model_parallel)) + _check_moe_config(moe_config, parallel_config) + self.use_moe = (moe_config.expert_num > 1) + self.use_past = use_past + self.seq_length = seq_length + self.hidden_size = hidden_size + self.batch_size = batch_size + self.layernorm1 = _LayerNorm((hidden_size,)).to_float(layernorm_compute_type) + self.layernorm1.shard(((parallel_config.data_parallel, 1),)) + self.layernorm2 = _LayerNorm((hidden_size,)).to_float(layernorm_compute_type) + self.layernorm2.shard(((parallel_config.data_parallel, 1),)) + + self.attention = MultiHeadAttention(batch_size=batch_size, + src_seq_length=seq_length, + tgt_seq_length=seq_length, + hidden_size=hidden_size, + num_heads=num_heads, + hidden_dropout_rate=hidden_dropout_rate, + attention_dropout_rate=attention_dropout_rate, + softmax_compute_type=softmax_compute_type, + param_init_type=param_init_type, + use_past=use_past, + parallel_config=parallel_config.dpmp if self.use_moe + else parallel_config) + if self.use_moe: + self.output = MoE(hidden_size=hidden_size, + dropout_rate=hidden_dropout_rate, + ffn_hidden_size=ffn_hidden_size, + param_init_type=param_init_type, + hidden_act=hidden_act, + moe_config=moe_config, + parallel_config=parallel_config) + else: + # Feed Forward Network, FFN + self.output = FeedForward(hidden_size=hidden_size, + dropout_rate=hidden_dropout_rate, + ffn_hidden_size=ffn_hidden_size, + param_init_type=param_init_type, + hidden_act=hidden_act, + parallel_config=parallel_config) + self.post_layernorm_residual = post_layernorm_residual + self.add = P.Add().shard(((parallel_config.data_parallel, 1), (parallel_config.data_parallel, 1))) + self.add_3d = P.Add().shard(((parallel_config.data_parallel, 1, 1), (parallel_config.data_parallel, 1, 1))) + self.dtype = mstype.float16 + self.key_past = None + self.value_past = None + + if self.use_past: + # operator used for state reuse + self.reducesum = P.ReduceSum().shard(((1, 1, 1, 1),)) + self.not_equal = P.NotEqual().shard(((1, 1, 1, 1), ())) + self.slice = P.StridedSlice().shard(((1, 1, 1, 1),)) + size_per_head = int(hidden_size / num_heads) + self.key_shape = (batch_size, num_heads, size_per_head, seq_length) + self.value_shape = (batch_size, num_heads, seq_length, size_per_head) + # parameters saving key and value states + self.key_past = Parameter(Tensor(np.zeros(shape=self.key_shape), self.dtype), name="key_past") + self.value_past = Parameter(Tensor(np.zeros(shape=self.value_shape), self.dtype), name="value_past") + self.tile = P.Tile().shard(((1, 1),)) + self.mul = P.Mul().shard(((1, 1, 1, 1), (1,))) + self.assign = P.Assign().shard(((1, 1, 1, 1), (1, 1, 1, 1))) else: - # Feed Forward Network, FFN - self.output = FeedForward(hidden_size=hidden_size, - dropout_rate=hidden_dropout_rate, - ffn_hidden_size=ffn_hidden_size, - param_init_type=param_init_type, - hidden_act=hidden_act, - parallel_config=parallel_config) - self.post_layernorm_residual = post_layernorm_residual - self.add = P.Add().shard(((parallel_config.data_parallel, 1), (parallel_config.data_parallel, 1))) - self.add_3d = P.Add().shard(((parallel_config.data_parallel, 1, 1), (parallel_config.data_parallel, 1, 1))) - self.dtype = mstype.float16 - self.key_past = None - self.value_past = None - - if self.use_past: - # operator used for state reuse - self.reducesum = P.ReduceSum().shard(((1, 1, 1, 1),)) - self.not_equal = P.NotEqual().shard(((1, 1, 1, 1), ())) - self.slice = P.StridedSlice().shard(((1, 1, 1, 1),)) - size_per_head = int(hidden_size / num_heads) - self.key_shape = (batch_size, num_heads, size_per_head, seq_length) - self.value_shape = (batch_size, num_heads, seq_length, size_per_head) - # parameters saving key and value states - self.key_past = Parameter(Tensor(np.zeros(shape=self.key_shape), self.dtype), name="key_past") - self.value_past = Parameter(Tensor(np.zeros(shape=self.value_shape), self.dtype), name="value_past") - self.tile = P.Tile().shard(((1, 1),)) - self.mul = P.Mul().shard(((1, 1, 1, 1), (1,))) - self.assign = P.Assign().shard(((1, 1, 1, 1), (1, 1, 1, 1))) + raise RuntimeError(f"The {self.cls_name} only support sharding propagation or " + f"semi-auto parallel mode now.") def construct(self, x, input_mask, init_reset=True, batch_valid_length=None): self._check_input(x, input_mask, init_reset, batch_valid_length) @@ -1573,103 +1803,208 @@ class TransformerDecoderLayer(Cell): moe_config=default_moe_config, parallel_config=default_dpmp_config): super(TransformerDecoderLayer, self).__init__() - _check_config(parallel_config) - if num_heads % parallel_config.model_parallel != 0: - raise ValueError("For 'TransformerDecoderLayer', the class variable 'num_heads' must be divisibled by " - "'parallel_config.model_parallel', but got the num_heads is {} and " - "parallel_config.model_parallel is {}.".format(num_heads, parallel_config.model_parallel)) - if hidden_size % parallel_config.model_parallel != 0: - raise ValueError("For 'TransformerDecoderLayer', the class variable 'hidden_size' must be divisibled by " - "'parallel_config.model_parallel', but got the hidden_size is {} and " - "parallel_config.model_parallel is {}." - .format(hidden_size, parallel_config.model_parallel)) - if ffn_hidden_size % parallel_config.model_parallel != 0: - raise ValueError("For 'TransformerDecoderLayer', the class variable 'ffn_hidden_size' must be " - "divisibled by 'parallel_config.model_parallel', but got the ffn_hidden_size is {} " - "and parallel_config.model_parallel is {}." - .format(ffn_hidden_size, parallel_config.model_parallel)) - _check_moe_config(moe_config, parallel_config) - self.use_moe = (moe_config.expert_num > 1) - if use_past: - raise ValueError(f"The {self.cls_name} does not support use_past=True.") - self.batch_size = batch_size - self.use_past = use_past - self.softmax_compute_type = softmax_compute_type + if _get_parallel_mode() in (ParallelMode.AUTO_PARALLEL,) and _is_sharding_propagation(): + _check_config(parallel_config) + if num_heads % parallel_config.model_parallel != 0: + raise ValueError("For 'TransformerDecoderLayer', the class variable 'num_heads' must be divisibled by " + "'parallel_config.model_parallel', but got the num_heads is {} and " + "parallel_config.model_parallel is {}.".format(num_heads, + parallel_config.model_parallel)) + if hidden_size % parallel_config.model_parallel != 0: + raise ValueError( + "For 'TransformerDecoderLayer', the class variable 'hidden_size' must be divisibled by " + "'parallel_config.model_parallel', but got the hidden_size is {} and " + "parallel_config.model_parallel is {}." + .format(hidden_size, parallel_config.model_parallel)) + if ffn_hidden_size % parallel_config.model_parallel != 0: + raise ValueError("For 'TransformerDecoderLayer', the class variable 'ffn_hidden_size' must be " + "divisibled by 'parallel_config.model_parallel', but got the ffn_hidden_size is {} " + "and parallel_config.model_parallel is {}." + .format(ffn_hidden_size, parallel_config.model_parallel)) + _check_moe_config(moe_config, parallel_config) + self.use_moe = (moe_config.expert_num > 1) + if use_past: + raise ValueError(f"The {self.cls_name} does not support use_past=True.") + self.batch_size = batch_size + self.use_past = use_past + self.softmax_compute_type = softmax_compute_type - self.src_seq_length = src_seq_length - self.tgt_seq_length = tgt_seq_length - self.use_past = use_past - self.hidden_size = hidden_size + self.src_seq_length = src_seq_length + self.tgt_seq_length = tgt_seq_length + self.use_past = use_past + self.hidden_size = hidden_size - self.layernorm1 = _LayerNorm((hidden_size,)).to_float(layernorm_compute_type) - self.layernorm1.shard(((parallel_config.data_parallel, 1),)) - self.layernorm2 = _LayerNorm((hidden_size,)).to_float(layernorm_compute_type) - self.layernorm2.shard(((parallel_config.data_parallel, 1),)) - self.attention = MultiHeadAttention(hidden_size=hidden_size, - num_heads=num_heads, - batch_size=batch_size, - src_seq_length=tgt_seq_length, - tgt_seq_length=tgt_seq_length, - hidden_dropout_rate=hidden_dropout_rate, - attention_dropout_rate=attention_dropout_rate, - use_past=use_past, - softmax_compute_type=softmax_compute_type, - param_init_type=param_init_type, - parallel_config=parallel_config.dpmp if self.use_moe else parallel_config) + self.layernorm1 = _LayerNorm((hidden_size,)).to_float(layernorm_compute_type) + self.layernorm2 = _LayerNorm((hidden_size,)).to_float(layernorm_compute_type) + self.attention = MultiHeadAttention(hidden_size=hidden_size, + num_heads=num_heads, + batch_size=batch_size, + src_seq_length=tgt_seq_length, + tgt_seq_length=tgt_seq_length, + hidden_dropout_rate=hidden_dropout_rate, + attention_dropout_rate=attention_dropout_rate, + use_past=use_past, + softmax_compute_type=softmax_compute_type, + param_init_type=param_init_type, + parallel_config=parallel_config.dpmp if self.use_moe + else parallel_config) - # Cross attention with the output of encoder as memory tensor - self.cross_attention = MultiHeadAttention(hidden_size=hidden_size, - num_heads=num_heads, - batch_size=batch_size, - src_seq_length=tgt_seq_length, - tgt_seq_length=src_seq_length, - hidden_dropout_rate=hidden_dropout_rate, - attention_dropout_rate=attention_dropout_rate, - softmax_compute_type=softmax_compute_type, - use_past=use_past, - param_init_type=param_init_type, - parallel_config=parallel_config.dpmp - if self.use_moe else parallel_config) - self.cross_attention_layernorm = _LayerNorm((hidden_size,)).to_float( - layernorm_compute_type) - self.cross_attention_layernorm.shard(((parallel_config.data_parallel, 1),)) + # Cross attention with the output of encoder as memory tensor + self.cross_attention = MultiHeadAttention(hidden_size=hidden_size, + num_heads=num_heads, + batch_size=batch_size, + src_seq_length=tgt_seq_length, + tgt_seq_length=src_seq_length, + hidden_dropout_rate=hidden_dropout_rate, + attention_dropout_rate=attention_dropout_rate, + softmax_compute_type=softmax_compute_type, + use_past=use_past, + param_init_type=param_init_type, + parallel_config=parallel_config.dpmp + if self.use_moe else parallel_config) + self.cross_attention_layernorm = _LayerNorm((hidden_size,)).to_float( + layernorm_compute_type) - if self.use_moe: - self.output = MoE(hidden_size=hidden_size, - dropout_rate=hidden_dropout_rate, - ffn_hidden_size=ffn_hidden_size, - param_init_type=param_init_type, - hidden_act=hidden_act, - moe_config=moe_config, - parallel_config=parallel_config) + if self.use_moe: + self.output = MoE(hidden_size=hidden_size, + dropout_rate=hidden_dropout_rate, + ffn_hidden_size=ffn_hidden_size, + param_init_type=param_init_type, + hidden_act=hidden_act, + moe_config=moe_config, + parallel_config=parallel_config) + else: + # Feed Forward Network, FFN + self.output = FeedForward(hidden_size=hidden_size, + dropout_rate=hidden_dropout_rate, + ffn_hidden_size=ffn_hidden_size, + hidden_act=hidden_act, + param_init_type=param_init_type, + parallel_config=parallel_config) + self.post_layernorm_residual = post_layernorm_residual + self.add = P.Add() + self.add_3d = P.Add() + self.dtype = mstype.float16 + self.key_past = None + self.value_past = None + if self.use_past: + # operator used for state reuse + self.reducesum = P.ReduceSum().shard(((1, 1, 1, 1),)) + self.not_equal = P.NotEqual().shard(((1, 1, 1, 1), ())) + self.slice = P.StridedSlice().shard(((1, 1, 1, 1),)) + size_per_head = int(hidden_size / num_heads) + self.key_shape = (batch_size, num_heads, size_per_head, tgt_seq_length) + self.value_shape = (batch_size, num_heads, tgt_seq_length, size_per_head) + # parameters saving key and value states + self.key_past = Parameter(Tensor(np.zeros(shape=self.key_shape), self.dtype), name="key_past") + self.value_past = Parameter(Tensor(np.zeros(shape=self.value_shape), self.dtype), name="value_past") + self.tile = P.Tile().shard(((1, 1),)) + self.mul = P.Mul().shard(((1, 1, 1, 1), (1,))) + self.assign = P.Assign().shard(((1, 1, 1, 1), (1, 1, 1, 1))) + elif _get_parallel_mode() not in (ParallelMode.AUTO_PARALLEL,): + _check_config(parallel_config) + if num_heads % parallel_config.model_parallel != 0: + raise ValueError("For 'TransformerDecoderLayer', the class variable 'num_heads' must be divisibled by " + "'parallel_config.model_parallel', but got the num_heads is {} and " + "parallel_config.model_parallel is {}.".format(num_heads, + parallel_config.model_parallel)) + if hidden_size % parallel_config.model_parallel != 0: + raise ValueError( + "For 'TransformerDecoderLayer', the class variable 'hidden_size' must be divisibled by " + "'parallel_config.model_parallel', but got the hidden_size is {} and " + "parallel_config.model_parallel is {}." + .format(hidden_size, parallel_config.model_parallel)) + if ffn_hidden_size % parallel_config.model_parallel != 0: + raise ValueError("For 'TransformerDecoderLayer', the class variable 'ffn_hidden_size' must be " + "divisibled by 'parallel_config.model_parallel', but got the ffn_hidden_size is {} " + "and parallel_config.model_parallel is {}." + .format(ffn_hidden_size, parallel_config.model_parallel)) + _check_moe_config(moe_config, parallel_config) + self.use_moe = (moe_config.expert_num > 1) + if use_past: + raise ValueError(f"The {self.cls_name} does not support use_past=True.") + self.batch_size = batch_size + self.use_past = use_past + self.softmax_compute_type = softmax_compute_type + + self.src_seq_length = src_seq_length + self.tgt_seq_length = tgt_seq_length + self.use_past = use_past + self.hidden_size = hidden_size + + self.layernorm1 = _LayerNorm((hidden_size,)).to_float(layernorm_compute_type) + self.layernorm1.shard(((parallel_config.data_parallel, 1),)) + self.layernorm2 = _LayerNorm((hidden_size,)).to_float(layernorm_compute_type) + self.layernorm2.shard(((parallel_config.data_parallel, 1),)) + self.attention = MultiHeadAttention(hidden_size=hidden_size, + num_heads=num_heads, + batch_size=batch_size, + src_seq_length=tgt_seq_length, + tgt_seq_length=tgt_seq_length, + hidden_dropout_rate=hidden_dropout_rate, + attention_dropout_rate=attention_dropout_rate, + use_past=use_past, + softmax_compute_type=softmax_compute_type, + param_init_type=param_init_type, + parallel_config=parallel_config.dpmp if self.use_moe + else parallel_config) + + # Cross attention with the output of encoder as memory tensor + self.cross_attention = MultiHeadAttention(hidden_size=hidden_size, + num_heads=num_heads, + batch_size=batch_size, + src_seq_length=tgt_seq_length, + tgt_seq_length=src_seq_length, + hidden_dropout_rate=hidden_dropout_rate, + attention_dropout_rate=attention_dropout_rate, + softmax_compute_type=softmax_compute_type, + use_past=use_past, + param_init_type=param_init_type, + parallel_config=parallel_config.dpmp + if self.use_moe else parallel_config) + self.cross_attention_layernorm = _LayerNorm((hidden_size,)).to_float( + layernorm_compute_type) + self.cross_attention_layernorm.shard(((parallel_config.data_parallel, 1),)) + + if self.use_moe: + self.output = MoE(hidden_size=hidden_size, + dropout_rate=hidden_dropout_rate, + ffn_hidden_size=ffn_hidden_size, + param_init_type=param_init_type, + hidden_act=hidden_act, + moe_config=moe_config, + parallel_config=parallel_config) + else: + # Feed Forward Network, FFN + self.output = FeedForward(hidden_size=hidden_size, + dropout_rate=hidden_dropout_rate, + ffn_hidden_size=ffn_hidden_size, + hidden_act=hidden_act, + param_init_type=param_init_type, + parallel_config=parallel_config) + self.post_layernorm_residual = post_layernorm_residual + self.add = P.Add().shard(((parallel_config.data_parallel, 1), (parallel_config.data_parallel, 1))) + self.add_3d = P.Add().shard(((parallel_config.data_parallel, 1, 1), (parallel_config.data_parallel, 1, 1))) + self.dtype = mstype.float16 + self.key_past = None + self.value_past = None + if self.use_past: + # operator used for state reuse + self.reducesum = P.ReduceSum().shard(((1, 1, 1, 1),)) + self.not_equal = P.NotEqual().shard(((1, 1, 1, 1), ())) + self.slice = P.StridedSlice().shard(((1, 1, 1, 1),)) + size_per_head = int(hidden_size / num_heads) + self.key_shape = (batch_size, num_heads, size_per_head, tgt_seq_length) + self.value_shape = (batch_size, num_heads, tgt_seq_length, size_per_head) + # parameters saving key and value states + self.key_past = Parameter(Tensor(np.zeros(shape=self.key_shape), self.dtype), name="key_past") + self.value_past = Parameter(Tensor(np.zeros(shape=self.value_shape), self.dtype), name="value_past") + self.tile = P.Tile().shard(((1, 1),)) + self.mul = P.Mul().shard(((1, 1, 1, 1), (1,))) + self.assign = P.Assign().shard(((1, 1, 1, 1), (1, 1, 1, 1))) else: - # Feed Forward Network, FFN - self.output = FeedForward(hidden_size=hidden_size, - dropout_rate=hidden_dropout_rate, - ffn_hidden_size=ffn_hidden_size, - hidden_act=hidden_act, - param_init_type=param_init_type, - parallel_config=parallel_config) - self.post_layernorm_residual = post_layernorm_residual - self.add = P.Add().shard(((parallel_config.data_parallel, 1), (parallel_config.data_parallel, 1))) - self.add_3d = P.Add().shard(((parallel_config.data_parallel, 1, 1), (parallel_config.data_parallel, 1, 1))) - self.dtype = mstype.float16 - self.key_past = None - self.value_past = None - if self.use_past: - # operator used for state reuse - self.reducesum = P.ReduceSum().shard(((1, 1, 1, 1),)) - self.not_equal = P.NotEqual().shard(((1, 1, 1, 1), ())) - self.slice = P.StridedSlice().shard(((1, 1, 1, 1),)) - size_per_head = int(hidden_size / num_heads) - self.key_shape = (batch_size, num_heads, size_per_head, tgt_seq_length) - self.value_shape = (batch_size, num_heads, tgt_seq_length, size_per_head) - # parameters saving key and value states - self.key_past = Parameter(Tensor(np.zeros(shape=self.key_shape), self.dtype), name="key_past") - self.value_past = Parameter(Tensor(np.zeros(shape=self.value_shape), self.dtype), name="value_past") - self.tile = P.Tile().shard(((1, 1),)) - self.mul = P.Mul().shard(((1, 1, 1, 1), (1,))) - self.assign = P.Assign().shard(((1, 1, 1, 1), (1, 1, 1, 1))) + raise RuntimeError(f"The {self.cls_name} only support sharding propagation or " + f"semi-auto parallel mode now.") def construct(self, hidden_stats, decoder_mask, @@ -2012,39 +2347,77 @@ class TransformerEncoder(Cell): moe_config=default_moe_config, parallel_config=default_transformer_config): super(TransformerEncoder, self).__init__() - _check_config(parallel_config) - _check_moe_config(moe_config, parallel_config) - self.use_moe = (moe_config.expert_num > 1) - self.add = P.Add().shard(((), ())) - self.aux_loss = Tensor(0.0, mstype.float32) - if _get_parallel_mode() in (ParallelMode.AUTO_PARALLEL,): - raise RuntimeError(f"The {self.cls_name} does not support auto parallel mode now.") - self.num_layers = num_layers - self.blocks = nn.CellList() - for i in range(num_layers): - block = TransformerEncoderLayer(hidden_size=hidden_size, - batch_size=batch_size, - ffn_hidden_size=ffn_hidden_size, - seq_length=seq_length, - attention_dropout_rate=attention_dropout_rate, - hidden_dropout_rate=hidden_dropout_rate, - layernorm_compute_type=layernorm_compute_type, - softmax_compute_type=softmax_compute_type, - num_heads=num_heads, - hidden_act=hidden_act, - post_layernorm_residual=post_layernorm_residual, - param_init_type=param_init_type, - use_past=use_past, - moe_config=moe_config, - parallel_config=parallel_config.moe_parallel_config if self.use_moe - else parallel_config.dp_mp_config) - # If the user doesn't pass the fusion function, use the default one - if not lambda_func: - lambda_func = _get_lambda_func() + if _get_parallel_mode() in (ParallelMode.AUTO_PARALLEL,) and _is_sharding_propagation(): + _check_config(parallel_config) + _check_moe_config(moe_config, parallel_config) + self.use_moe = (moe_config.expert_num > 1) + self.add = P.Add() + self.aux_loss = Tensor(0.0, mstype.float32) + self.num_layers = num_layers + self.blocks = nn.CellList() + for i in range(num_layers): + block = TransformerEncoderLayer(hidden_size=hidden_size, + batch_size=batch_size, + ffn_hidden_size=ffn_hidden_size, + seq_length=seq_length, + attention_dropout_rate=attention_dropout_rate, + hidden_dropout_rate=hidden_dropout_rate, + layernorm_compute_type=layernorm_compute_type, + softmax_compute_type=softmax_compute_type, + num_heads=num_heads, + hidden_act=hidden_act, + post_layernorm_residual=post_layernorm_residual, + param_init_type=param_init_type, + use_past=use_past, + moe_config=moe_config, + parallel_config=parallel_config.moe_parallel_config if self.use_moe + else parallel_config.dp_mp_config) + # If the user doesn't pass the fusion function, use the default one + if not lambda_func: + lambda_func = _get_lambda_func() - lambda_func(block, layer_id=i, layers=num_layers, - offset=offset, parallel_config=parallel_config) - self.blocks.append(block) + lambda_func(block, layer_id=i, layers=num_layers, + offset=offset, parallel_config=parallel_config) + self.blocks.append(block) + elif _get_parallel_mode() not in (ParallelMode.AUTO_PARALLEL,): + _check_config(parallel_config) + _check_moe_config(moe_config, parallel_config) + self.use_moe = (moe_config.expert_num > 1) + self.add = P.Add().shard(((), ())) + self.aux_loss = Tensor(0.0, mstype.float32) + logger.warning("For parallel mode, sharding propagation is recommended, you can use it by setting " + "'set_auto_parallel_context(parallel_mode=ParallelMode.AUTO_PARALLEL, " + "search_mode=\"sharding_propagation\")' and " + "'set_algo_parameters(elementwise_op_strategy_follow=False, fully_use_devices=False)'") + self.num_layers = num_layers + self.blocks = nn.CellList() + for i in range(num_layers): + block = TransformerEncoderLayer(hidden_size=hidden_size, + batch_size=batch_size, + ffn_hidden_size=ffn_hidden_size, + seq_length=seq_length, + attention_dropout_rate=attention_dropout_rate, + hidden_dropout_rate=hidden_dropout_rate, + layernorm_compute_type=layernorm_compute_type, + softmax_compute_type=softmax_compute_type, + num_heads=num_heads, + hidden_act=hidden_act, + post_layernorm_residual=post_layernorm_residual, + param_init_type=param_init_type, + use_past=use_past, + moe_config=moe_config, + parallel_config=parallel_config.moe_parallel_config if self.use_moe + else parallel_config.dp_mp_config) + # If the user doesn't pass the fusion function, use the default one + if not lambda_func: + lambda_func = _get_lambda_func() + + lambda_func(block, layer_id=i, layers=num_layers, + offset=offset, parallel_config=parallel_config) + self.blocks.append(block) + else: + raise RuntimeError(f"The {self.cls_name} only support sharding propagation or " + f"semi-auto parallel mode now.") def construct(self, hidden_states, attention_mask, init_reset=True, batch_valid_length=None): present_layer = () @@ -2207,42 +2580,83 @@ class TransformerDecoder(Cell): moe_config=default_moe_config, parallel_config=default_transformer_config): super(TransformerDecoder, self).__init__() - _check_config(parallel_config) + if _get_parallel_mode() in (ParallelMode.AUTO_PARALLEL,) and _is_sharding_propagation(): + _check_config(parallel_config) - self.add = P.Add().shard(((), ())) - self.aux_loss = Tensor(0.0, mstype.float32) - if _get_parallel_mode() in (ParallelMode.AUTO_PARALLEL,): - raise RuntimeError(f"The {self.cls_name} does not support auto parallel mode now.") - self.num_layers = num_layers - self.blocks = nn.CellList() - _check_moe_config(moe_config, parallel_config) - self.use_moe = (moe_config.expert_num > 1) - for i in range(num_layers): - block = TransformerDecoderLayer(hidden_size=hidden_size, - batch_size=batch_size, - ffn_hidden_size=ffn_hidden_size, - src_seq_length=src_seq_length, - tgt_seq_length=tgt_seq_length, - attention_dropout_rate=attention_dropout_rate, - hidden_dropout_rate=hidden_dropout_rate, - num_heads=num_heads, - layernorm_compute_type=layernorm_compute_type, - softmax_compute_type=softmax_compute_type, - hidden_act=hidden_act, - use_past=use_past, - param_init_type=param_init_type, - post_layernorm_residual=post_layernorm_residual, - moe_config=moe_config, - parallel_config=parallel_config.moe_parallel_config if self.use_moe - else parallel_config.dp_mp_config) - # If the user doesn't pass the fusion function, use the default one - if not lambda_func: - lambda_func = _get_lambda_func() + self.add = P.Add() + self.aux_loss = Tensor(0.0, mstype.float32) + self.num_layers = num_layers + self.blocks = nn.CellList() + _check_moe_config(moe_config, parallel_config) + self.use_moe = (moe_config.expert_num > 1) + for i in range(num_layers): + block = TransformerDecoderLayer(hidden_size=hidden_size, + batch_size=batch_size, + ffn_hidden_size=ffn_hidden_size, + src_seq_length=src_seq_length, + tgt_seq_length=tgt_seq_length, + attention_dropout_rate=attention_dropout_rate, + hidden_dropout_rate=hidden_dropout_rate, + num_heads=num_heads, + layernorm_compute_type=layernorm_compute_type, + softmax_compute_type=softmax_compute_type, + hidden_act=hidden_act, + use_past=use_past, + param_init_type=param_init_type, + post_layernorm_residual=post_layernorm_residual, + moe_config=moe_config, + parallel_config=parallel_config.moe_parallel_config if self.use_moe + else parallel_config.dp_mp_config) + # If the user doesn't pass the fusion function, use the default one + if not lambda_func: + lambda_func = _get_lambda_func() - lambda_func(block, layer_id=i, layers=num_layers, - offset=offset, parallel_config=parallel_config) + lambda_func(block, layer_id=i, layers=num_layers, + offset=offset, parallel_config=parallel_config) - self.blocks.append(block) + self.blocks.append(block) + elif _get_parallel_mode() not in (ParallelMode.AUTO_PARALLEL,): + _check_config(parallel_config) + + self.add = P.Add().shard(((), ())) + self.aux_loss = Tensor(0.0, mstype.float32) + logger.warning("For parallel mode, sharding propagation is recommended, you can use it by setting " + "'set_auto_parallel_context(parallel_mode=ParallelMode.AUTO_PARALLEL, " + "search_mode=\"sharding_propagation\")' and " + "'set_algo_parameters(elementwise_op_strategy_follow=False, fully_use_devices=False)'") + self.num_layers = num_layers + self.blocks = nn.CellList() + _check_moe_config(moe_config, parallel_config) + self.use_moe = (moe_config.expert_num > 1) + for i in range(num_layers): + block = TransformerDecoderLayer(hidden_size=hidden_size, + batch_size=batch_size, + ffn_hidden_size=ffn_hidden_size, + src_seq_length=src_seq_length, + tgt_seq_length=tgt_seq_length, + attention_dropout_rate=attention_dropout_rate, + hidden_dropout_rate=hidden_dropout_rate, + num_heads=num_heads, + layernorm_compute_type=layernorm_compute_type, + softmax_compute_type=softmax_compute_type, + hidden_act=hidden_act, + use_past=use_past, + param_init_type=param_init_type, + post_layernorm_residual=post_layernorm_residual, + moe_config=moe_config, + parallel_config=parallel_config.moe_parallel_config if self.use_moe + else parallel_config.dp_mp_config) + # If the user doesn't pass the fusion function, use the default one + if not lambda_func: + lambda_func = _get_lambda_func() + + lambda_func(block, layer_id=i, layers=num_layers, + offset=offset, parallel_config=parallel_config) + + self.blocks.append(block) + else: + raise RuntimeError(f"The {self.cls_name} only support sharding propagation or " + f"semi-auto parallel mode now.") def construct(self, hidden_states, attention_mask, encoder_output=None, memory_mask=None, init_reset=True, batch_valid_length=None): @@ -2431,70 +2845,139 @@ class Transformer(Cell): moe_config=default_moe_config, parallel_config=default_transformer_config): super(Transformer, self).__init__() - _check_config(parallel_config) - self.batch_size = batch_size - self.hidden_size = hidden_size - self.src_seq_length = src_seq_length - self.tgt_seq_length = tgt_seq_length - self.use_past = use_past - if encoder_layers <= 0 < decoder_layers: - raise ValueError(f"Transformer doest support encoder layer {encoder_layers} and decoder" - f"layer {decoder_layers}, please use TransformerDecoder") - if encoder_layers > 0 and decoder_layers > 0 and use_past: - raise ValueError(f"The {self.cls_name} with encoder and decoder does not support use_past=True.") - if _get_parallel_mode() in (ParallelMode.AUTO_PARALLEL,): - raise RuntimeError(f"The {self.cls_name} does not support auto parallel mode now.") - # The shard setting of Transformer is set within the TransformerEncoderLayer - if not lambda_func: - lambda_func = _get_lambda_func(total_layer=encoder_layers + decoder_layers) - _check_moe_config(moe_config, parallel_config) - self.use_moe = (moe_config.expert_num > 1) - self.add = P.Add().shard(((), ())) - self.aux_loss = Tensor(0.0, mstype.float32) - if encoder_layers > 0: - self.encoder = TransformerEncoder(num_layers=encoder_layers, - batch_size=batch_size, - hidden_size=hidden_size, - ffn_hidden_size=ffn_hidden_size, - num_heads=num_heads, - seq_length=src_seq_length, - attention_dropout_rate=attention_dropout_rate, - hidden_dropout_rate=hidden_dropout_rate, - hidden_act=hidden_act, - layernorm_compute_type=layernorm_compute_type, - softmax_compute_type=softmax_compute_type, - post_layernorm_residual=post_layernorm_residual, - param_init_type=param_init_type, - lambda_func=lambda_func, - use_past=use_past, - moe_config=moe_config, - parallel_config=parallel_config) - else: - self.encoder = None + if _get_parallel_mode() in (ParallelMode.AUTO_PARALLEL,) and _is_sharding_propagation(): + _check_config(parallel_config) + self.batch_size = batch_size + self.hidden_size = hidden_size + self.src_seq_length = src_seq_length + self.tgt_seq_length = tgt_seq_length + self.use_past = use_past + if encoder_layers <= 0 < decoder_layers: + raise ValueError(f"Transformer doest support encoder layer {encoder_layers} and decoder" + f"layer {decoder_layers}, please use TransformerDecoder") + if encoder_layers > 0 and decoder_layers > 0 and use_past: + raise ValueError(f"The {self.cls_name} with encoder and decoder does not support use_past=True.") + # The shard setting of Transformer is set within the TransformerEncoderLayer + if not lambda_func: + lambda_func = _get_lambda_func(total_layer=encoder_layers + decoder_layers) + _check_moe_config(moe_config, parallel_config) + self.use_moe = (moe_config.expert_num > 1) + self.add = P.Add() + self.aux_loss = Tensor(0.0, mstype.float32) + if encoder_layers > 0: + self.encoder = TransformerEncoder(num_layers=encoder_layers, + batch_size=batch_size, + hidden_size=hidden_size, + ffn_hidden_size=ffn_hidden_size, + num_heads=num_heads, + seq_length=src_seq_length, + attention_dropout_rate=attention_dropout_rate, + hidden_dropout_rate=hidden_dropout_rate, + hidden_act=hidden_act, + layernorm_compute_type=layernorm_compute_type, + softmax_compute_type=softmax_compute_type, + post_layernorm_residual=post_layernorm_residual, + param_init_type=param_init_type, + lambda_func=lambda_func, + use_past=use_past, + moe_config=moe_config, + parallel_config=parallel_config) + else: + self.encoder = None - # Offset is needed as the encoder has consumed some flags. - # so the decoder need to increase the flags based on the encoder layer - self.decoder = None - if decoder_layers > 0: - self.decoder = TransformerDecoder(num_layers=decoder_layers, - batch_size=batch_size, - hidden_size=hidden_size, - ffn_hidden_size=ffn_hidden_size, - num_heads=num_heads, - src_seq_length=src_seq_length, - tgt_seq_length=tgt_seq_length, - attention_dropout_rate=attention_dropout_rate, - hidden_dropout_rate=hidden_dropout_rate, - hidden_act=hidden_act, - post_layernorm_residual=post_layernorm_residual, - layernorm_compute_type=layernorm_compute_type, - softmax_compute_type=softmax_compute_type, - lambda_func=lambda_func, - use_past=use_past, - param_init_type=param_init_type, - offset=encoder_layers, - moe_config=moe_config, - parallel_config=parallel_config) + # Offset is needed as the encoder has consumed some flags. + # so the decoder need to increase the flags based on the encoder layer + self.decoder = None + if decoder_layers > 0: + self.decoder = TransformerDecoder(num_layers=decoder_layers, + batch_size=batch_size, + hidden_size=hidden_size, + ffn_hidden_size=ffn_hidden_size, + num_heads=num_heads, + src_seq_length=src_seq_length, + tgt_seq_length=tgt_seq_length, + attention_dropout_rate=attention_dropout_rate, + hidden_dropout_rate=hidden_dropout_rate, + hidden_act=hidden_act, + post_layernorm_residual=post_layernorm_residual, + layernorm_compute_type=layernorm_compute_type, + softmax_compute_type=softmax_compute_type, + lambda_func=lambda_func, + use_past=use_past, + param_init_type=param_init_type, + offset=encoder_layers, + moe_config=moe_config, + parallel_config=parallel_config) + elif _get_parallel_mode() not in (ParallelMode.AUTO_PARALLEL,): + _check_config(parallel_config) + self.batch_size = batch_size + self.hidden_size = hidden_size + self.src_seq_length = src_seq_length + self.tgt_seq_length = tgt_seq_length + self.use_past = use_past + if encoder_layers <= 0 < decoder_layers: + raise ValueError(f"Transformer doest support encoder layer {encoder_layers} and decoder" + f"layer {decoder_layers}, please use TransformerDecoder") + if encoder_layers > 0 and decoder_layers > 0 and use_past: + raise ValueError(f"The {self.cls_name} with encoder and decoder does not support use_past=True.") + logger.warning("For parallel mode, sharding propagation is recommended, you can use it by setting " + "'set_auto_parallel_context(parallel_mode=ParallelMode.AUTO_PARALLEL, " + "search_mode=\"sharding_propagation\")' and " + "'set_algo_parameters(elementwise_op_strategy_follow=False, fully_use_devices=False)'") + # The shard setting of Transformer is set within the TransformerEncoderLayer + if not lambda_func: + lambda_func = _get_lambda_func(total_layer=encoder_layers + decoder_layers) + _check_moe_config(moe_config, parallel_config) + self.use_moe = (moe_config.expert_num > 1) + self.add = P.Add().shard(((), ())) + self.aux_loss = Tensor(0.0, mstype.float32) + if encoder_layers > 0: + self.encoder = TransformerEncoder(num_layers=encoder_layers, + batch_size=batch_size, + hidden_size=hidden_size, + ffn_hidden_size=ffn_hidden_size, + num_heads=num_heads, + seq_length=src_seq_length, + attention_dropout_rate=attention_dropout_rate, + hidden_dropout_rate=hidden_dropout_rate, + hidden_act=hidden_act, + layernorm_compute_type=layernorm_compute_type, + softmax_compute_type=softmax_compute_type, + post_layernorm_residual=post_layernorm_residual, + param_init_type=param_init_type, + lambda_func=lambda_func, + use_past=use_past, + moe_config=moe_config, + parallel_config=parallel_config) + else: + self.encoder = None + + # Offset is needed as the encoder has consumed some flags. + # so the decoder need to increase the flags based on the encoder layer + self.decoder = None + if decoder_layers > 0: + self.decoder = TransformerDecoder(num_layers=decoder_layers, + batch_size=batch_size, + hidden_size=hidden_size, + ffn_hidden_size=ffn_hidden_size, + num_heads=num_heads, + src_seq_length=src_seq_length, + tgt_seq_length=tgt_seq_length, + attention_dropout_rate=attention_dropout_rate, + hidden_dropout_rate=hidden_dropout_rate, + hidden_act=hidden_act, + post_layernorm_residual=post_layernorm_residual, + layernorm_compute_type=layernorm_compute_type, + softmax_compute_type=softmax_compute_type, + lambda_func=lambda_func, + use_past=use_past, + param_init_type=param_init_type, + offset=encoder_layers, + moe_config=moe_config, + parallel_config=parallel_config) + else: + raise RuntimeError(f"The {self.cls_name} only support sharding propagation or " + f"semi-auto parallel mode now.") def construct(self, encoder_inputs, encoder_masks, diff --git a/mindspore/python/mindspore/parallel/_utils.py b/mindspore/python/mindspore/parallel/_utils.py index 85e9c2390f..cc7e9f9561 100644 --- a/mindspore/python/mindspore/parallel/_utils.py +++ b/mindspore/python/mindspore/parallel/_utils.py @@ -30,6 +30,12 @@ def _get_parallel_mode(): return auto_parallel_context().get_parallel_mode() +def _is_sharding_propagation(): + """Is sharding propagation.""" + return (auto_parallel_context().get_strategy_search_mode() == "sharding_propagation") or ( + auto_parallel_context().get_sharding_propagation()) + + def _is_in_auto_parallel_mode(): return _get_parallel_mode() in [ParallelMode.SEMI_AUTO_PARALLEL, ParallelMode.AUTO_PARALLEL]