[Cherry pick] fix srn for sub_layers (#2694)

* fix srn for sublayer

* update for paddle2.1
This commit is contained in:
xiaoting 2021-05-07 10:50:12 +08:00 committed by GitHub
parent 3ed769d155
commit 81aba737d3
No known key found for this signature in database
GPG Key ID: 4AEE18F83AFDEB23
1 changed files with 2 additions and 5 deletions

View File

@ -285,8 +285,7 @@ class PrePostProcessLayer(nn.Layer):
elif cmd == "n": # add layer normalization elif cmd == "n": # add layer normalization
self.functors.append( self.functors.append(
self.add_sublayer( self.add_sublayer(
"layer_norm_%d" % len( "layer_norm_%d" % len(self.sublayers()),
self.sublayers(include_sublayers=False)),
paddle.nn.LayerNorm( paddle.nn.LayerNorm(
normalized_shape=d_model, normalized_shape=d_model,
weight_attr=fluid.ParamAttr( weight_attr=fluid.ParamAttr(
@ -320,9 +319,7 @@ class PrepareEncoder(nn.Layer):
self.src_emb_dim = src_emb_dim self.src_emb_dim = src_emb_dim
self.src_max_len = src_max_len self.src_max_len = src_max_len
self.emb = paddle.nn.Embedding( self.emb = paddle.nn.Embedding(
num_embeddings=self.src_max_len, num_embeddings=self.src_max_len, embedding_dim=self.src_emb_dim)
embedding_dim=self.src_emb_dim,
sparse=True)
self.dropout_rate = dropout_rate self.dropout_rate = dropout_rate
def forward(self, src_word, src_pos): def forward(self, src_word, src_pos):