| 41 | raise NotImplementedError |
| 42 | |
| 43 | def __call__(self, inputs): |
| 44 | src, pos, negs = self.to_sample(inputs) |
| 45 | embedding = self.embed(src) # [batch, 1, dim] |
| 46 | embedding_pos = self.embed(pos) # [batch, num_negs, dim] |
| 47 | embedding_negs = self.embed(negs) # [batch, num_negs, dim] |
| 48 | |
| 49 | # [batch, 1, num_negs] |
| 50 | logits = tf.matmul(embedding, embedding_pos, transpose_b=True) |
| 51 | # [batch, 1, num_negs] |
| 52 | neg_logits = tf.matmul(embedding, embedding_negs, transpose_b=True) |
| 53 | true_xent = tf.nn.sigmoid_cross_entropy_with_logits( |
| 54 | labels=tf.ones_like(logits), logits=logits) |
| 55 | negative_xent = tf.nn.sigmoid_cross_entropy_with_logits( |
| 56 | labels=tf.zeros_like(neg_logits), logits=neg_logits) |
| 57 | loss = tf.reduce_mean(tf.concat([tf.reshape(true_xent, [-1, 1]), |
| 58 | tf.reshape(negative_xent, |
| 59 | [-1, 1])], 0)) |
| 60 | predict = tf.nn.sigmoid(logits) |
| 61 | neg_predict = tf.nn.sigmoid(neg_logits) |
| 62 | label = tf.ones_like(logits) |
| 63 | neg_label = tf.zeros_like(neg_logits) |
| 64 | acc = tf_euler.utils.metrics.acc_score( |
| 65 | tf.concat([label, neg_label], axis=2), |
| 66 | tf.concat([predict, neg_predict], axis=2)) |
| 67 | |
| 68 | embedding = self.embed(inputs) |
| 69 | |
| 70 | return (embedding, loss, 'acc', acc) |
| 71 | |