forward_attn_mask and config update

This commit is contained in:
Eren Golge committed 2019-06-06 11:14:20 +02:00
1 parent 127a6b68e0
commit 31fe02412c
7 files changed
+32 -15

No files matched your search

+4 -2
View File
@@ -108,7 +108,7 @@ class Attention(nn.Module):
def __init__(self, attention_rnn_dim, embedding_dim, attention_dim,
location_attention, attention_location_n_filters,
attention_location_kernel_size, windowing, norm, forward_attn,
trans_agent):
trans_agent, forward_attn_mask):
super(Attention, self).__init__()
self.query_layer = Linear(
attention_rnn_dim, attention_dim, bias=False, init_gain='tanh')
@@ -128,6 +128,7 @@ class Attention(nn.Module):
self.norm = norm
self.forward_attn = forward_attn
self.trans_agent = trans_agent
self.forward_attn_mask = forward_attn_mask
self.location_attention = location_attention
def init_win_idx(self):
@@ -203,9 +204,10 @@ class Attention(nn.Module):
alpha = (((1 - self.u) * self.alpha.clone().to(inputs.device) +
self.u * prev_alpha) + 1e-8) * alignment
# force incremental alignment - TODO: make configurable
if not self.training:
if not self.training and self.forward_attn_mask:
_, n = prev_alpha.max(1)
val, n2 = alpha.max(1)
print(True)
for b in range(alignment.shape[0]):
alpha[b, n[b] + 2:] = 0
alpha[b, :(n[b] - 1)] = 0 # ignore all previous states to prevent repetition.
+3 -2
View File
@@ -271,7 +271,7 @@ class Decoder(nn.Module):
def __init__(self, in_features, memory_dim, r, memory_size, attn_windowing,
attn_norm, prenet_type, prenet_dropout, forward_attn,
trans_agent, location_attn, separate_stopnet):
trans_agent, forward_attn_mask, location_attn, separate_stopnet):
super(Decoder, self).__init__()
self.r = r
self.in_features = in_features
@@ -296,7 +296,8 @@ class Decoder(nn.Module):
windowing=attn_windowing,
norm=attn_norm,
forward_attn=forward_attn,
trans_agent=trans_agent)
trans_agent=trans_agent,
forward_attn_mask=forward_attn_mask)
# (processed_memory | attention context) -> |Linear| -> decoder_RNN_input
self.project_to_decoder_in = nn.Linear(256 + in_features, 256)
# decoder_RNN_input -> |RNN| -> RNN_state
+12 -4
View File
@@ -97,7 +97,7 @@ class Encoder(nn.Module):
class Decoder(nn.Module):
def __init__(self, in_features, inputs_dim, r, attn_win, attn_norm,
prenet_type, prenet_dropout, forward_attn, trans_agent,
location_attn, separate_stopnet):
forward_attn_mask, location_attn, separate_stopnet):
super(Decoder, self).__init__()
self.mel_channels = inputs_dim
self.r = r
@@ -118,9 +118,17 @@ class Decoder(nn.Module):
self.attention_rnn = nn.LSTMCell(self.prenet_dim + in_features,
self.attention_rnn_dim)
self.attention_layer = Attention(self.attention_rnn_dim, in_features,
128, location_attn, 32, 31, attn_win,
attn_norm, forward_attn, trans_agent)
self.attention_layer = Attention(attention_rnn_dim=self.attention_rnn_dim,
embedding_dim=in_features,
attention_dim=128,
location_attention=location_attn,
attention_location_n_filters=32,
attention_location_kernel_size=31,
windowing=attn_win,
norm=attn_norm,
forward_attn=forward_attn,
trans_agent=trans_agent,
forward_attn_mask=forward_attn_mask)
self.decoder_rnn = nn.LSTMCell(self.attention_rnn_dim + in_features,
self.decoder_rnn_dim, 1)