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,
)
| 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. |