x: decoder input h: encoder output
(self, x, x_mask, h, h_mask)
| 149 | self.norm_layers_2.append(LayerNorm(hidden_channels)) |
| 150 | |
| 151 | def forward(self, x, x_mask, h, h_mask): |
| 152 | """ |
| 153 | x: decoder input |
| 154 | h: encoder output |
| 155 | """ |
| 156 | self_attn_mask = commons.subsequent_mask(x_mask.size(2)).to( |
| 157 | device=x.device, dtype=x.dtype |
| 158 | ) |
| 159 | encdec_attn_mask = h_mask.unsqueeze(2) * x_mask.unsqueeze(-1) |
| 160 | x = x * x_mask |
| 161 | for i in range(self.n_layers): |
| 162 | y = self.self_attn_layers[i](x, x, self_attn_mask) |
| 163 | y = self.drop(y) |
| 164 | x = self.norm_layers_0[i](x + y) |
| 165 | |
| 166 | y = self.encdec_attn_layers[i](x, h, encdec_attn_mask) |
| 167 | y = self.drop(y) |
| 168 | x = self.norm_layers_1[i](x + y) |
| 169 | |
| 170 | y = self.ffn_layers[i](x, x_mask) |
| 171 | y = self.drop(y) |
| 172 | x = self.norm_layers_2[i](x + y) |
| 173 | x = x * x_mask |
| 174 | return x |
| 175 | |
| 176 | |
| 177 | class MultiHeadAttention(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected