MCPcopy Create free account
hub / github.com/espnet/espnet / forward

Method forward

espnet2/asr/decoder/linear_decoder.py:47–77  ·  view source on GitHub ↗

Forward method. Args: hs_pad: (B, Tmax, D) hlens: (B,) Returns: output: (B, n_classes)

(
        self,
        hs_pad: torch.Tensor,
        hlens: torch.Tensor,
        ys_in_pad: torch.Tensor = None,
        ys_in_lens: torch.Tensor = None,
    )

Source from the content-addressed store, hash-verified

45 self.pooling = pooling
46
47 def forward(
48 self,
49 hs_pad: torch.Tensor,
50 hlens: torch.Tensor,
51 ys_in_pad: torch.Tensor = None,
52 ys_in_lens: torch.Tensor = None,
53 ) -> Tuple[torch.Tensor, torch.Tensor]:
54 """Forward method.
55
56 Args:
57 hs_pad: (B, Tmax, D)
58 hlens: (B,)
59 Returns:
60 output: (B, n_classes)
61 """
62
63 mask = make_pad_mask(lengths=hlens, xs=hs_pad, length_dim=1).to(hs_pad.device)
64 if self.dropout is not None:
65 hs_pad = self.dropout(hs_pad)
66 if self.pooling == "mean":
67 unmasked_entries = (~mask).to(dtype=hs_pad.dtype)
68 input_feature = (hs_pad * unmasked_entries).sum(dim=1)
69 input_feature = input_feature / unmasked_entries.sum(dim=1)
70 elif self.pooling == "max":
71 input_feature = hs_pad.masked_fill(mask, float("-inf"))
72 input_feature, _ = torch.max(input_feature, dim=1)
73 elif self.pooling == "CLS":
74 input_feature = hs_pad[:, 0, :]
75
76 output = self.linear_out(input_feature)
77 return output
78
79 def score(self, ys, state, x):
80 """Classify x.

Callers 2

scoreMethod · 0.95
test_linear_decoderFunction · 0.95

Calls 2

make_pad_maskFunction · 0.90
toMethod · 0.80

Tested by 1

test_linear_decoderFunction · 0.76