MCPcopy Create free account
hub / github.com/apple/axlearn / attend

Method attend

axlearn/common/encoder.py:137–147  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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
150class CausalEncoder(Encoder):

Callers 1

compute_logitsMethod · 0.45

Calls 1

child_contextFunction · 0.90

Tested by

no test coverage detected