Computes logits with token embedding. Args: x: an int Tensor of shape [batch_size, seq_len, hidden_dim]. Returns: A float Tensor of shape [batch_size, seq_len, vocab_size].
(self, x: Tensor)
| 135 | return self.attention_mask(segment_ids=segment_ids, positions=positions) |
| 136 | |
| 137 | def attend(self, x: Tensor) -> Tensor: |
| 138 | """Computes logits with token embedding. |
| 139 | |
| 140 | Args: |
| 141 | x: an int Tensor of shape [batch_size, seq_len, hidden_dim]. |
| 142 | |
| 143 | Returns: |
| 144 | A float Tensor of shape [batch_size, seq_len, vocab_size]. |
| 145 | """ |
| 146 | with child_context("emb_attend", module=self.emb): |
| 147 | return self.emb.attend(x) |
| 148 | |
| 149 | |
| 150 | class CausalEncoder(Encoder): |
no test coverage detected