MCPcopy Create free account
hub / github.com/modelscope/FunASR / forward

Method forward

funasr/models/language_model/rnn/attentions.py:214–259  ·  view source on GitHub ↗

AttAdd forward :param torch.Tensor enc_hs_pad: padded encoder hidden state (B x T_max x D_enc) :param list enc_hs_len: padded encoder hidden state length (B) :param torch.Tensor dec_z: decoder hidden state (B x D_dec) :param torch.Tensor att_prev: dummy (does not use

(self, enc_hs_pad, enc_hs_len, dec_z, att_prev, scaling=2.0)

Source from the content-addressed store, hash-verified

212 self.mask = None
213
214 def forward(self, enc_hs_pad, enc_hs_len, dec_z, att_prev, scaling=2.0):
215 """AttAdd forward
216
217 :param torch.Tensor enc_hs_pad: padded encoder hidden state (B x T_max x D_enc)
218 :param list enc_hs_len: padded encoder hidden state length (B)
219 :param torch.Tensor dec_z: decoder hidden state (B x D_dec)
220 :param torch.Tensor att_prev: dummy (does not use)
221 :param float scaling: scaling parameter before applying softmax
222 :return: attention weighted encoder state (B, D_enc)
223 :rtype: torch.Tensor
224 :return: previous attention weights (B x T_max)
225 :rtype: torch.Tensor
226 """
227
228 batch = len(enc_hs_pad)
229 # pre-compute all h outside the decoder loop
230 if self.pre_compute_enc_h is None or self.han_mode:
231 self.enc_h = enc_hs_pad # utt x frame x hdim
232 self.h_length = self.enc_h.size(1)
233 # utt x frame x att_dim
234 self.pre_compute_enc_h = self.mlp_enc(self.enc_h)
235
236 if dec_z is None:
237 dec_z = enc_hs_pad.new_zeros(batch, self.dunits)
238 else:
239 dec_z = dec_z.view(batch, self.dunits)
240
241 # dec_z_tiled: utt x frame x att_dim
242 dec_z_tiled = self.mlp_dec(dec_z).view(batch, 1, self.att_dim)
243
244 # dot with gvec
245 # utt x frame x att_dim -> utt x frame
246 e = self.gvec(torch.tanh(self.pre_compute_enc_h + dec_z_tiled)).squeeze(2)
247
248 # NOTE consider zero padding when compute w.
249 if self.mask is None:
250 self.mask = to_device(enc_hs_pad, make_pad_mask(enc_hs_len))
251 e.masked_fill_(self.mask, -float("inf"))
252 w = F.softmax(scaling * e, dim=1)
253
254 # weighted sum over flames
255 # utt x hdim
256 # NOTE use bmm instead of sum(*)
257 c = torch.sum(self.enc_h * w.view(batch, self.h_length, 1), dim=1)
258
259 return c, w
260
261
262class AttLoc(torch.nn.Module):

Callers

nothing calls this directly

Calls 3

to_deviceFunction · 0.90
make_pad_maskFunction · 0.90
softmaxMethod · 0.45

Tested by

no test coverage detected