(self, inputs)
| 72 | return src, pos, negs |
| 73 | |
| 74 | def __call__(self, inputs): |
| 75 | src, pos, negs = self.to_sample(inputs) |
| 76 | embedding = self.embed(src) |
| 77 | embedding_pos = self.embed_context(pos) |
| 78 | embedding_negs = self.embed_context(negs) |
| 79 | |
| 80 | logits = tf.matmul(embedding, embedding_pos, transpose_b=True) |
| 81 | neg_logits = tf.matmul(embedding, embedding_negs, transpose_b=True) |
| 82 | metric = self.metric_class(logits, neg_logits) |
| 83 | true_xent = tf.nn.sigmoid_cross_entropy_with_logits( |
| 84 | labels=tf.ones_like(logits), logits=logits) |
| 85 | negative_xent = tf.nn.sigmoid_cross_entropy_with_logits( |
| 86 | labels=tf.zeros_like(neg_logits), logits=neg_logits) |
| 87 | loss = tf.reduce_mean(tf.concat([tf.reshape(true_xent, [-1, 1]), |
| 88 | tf.reshape(negative_xent, |
| 89 | [-1, 1])], 0)) |
| 90 | embedding = self.embed(inputs) |
| 91 | return (embedding, loss, self.metric_name, metric) |
nothing calls this directly
no test coverage detected