| 20 | |
| 21 | class DistMult(tf_euler.utils.layers.Layer): |
| 22 | def __init__(self, node_type, edge_type, |
| 23 | node_max_id, edge_max_id, |
| 24 | ent_dim, rel_dim, |
| 25 | num_negs=5, margin=1, |
| 26 | metric_name='mrr', |
| 27 | corrupt='both', |
| 28 | l2_regular=False, |
| 29 | regular_param=0.0001, |
| 30 | **kwargs): |
| 31 | super(DistMult, self).__init__(**kwargs) |
| 32 | self.node_type = node_type |
| 33 | self.edge_type = edge_type |
| 34 | self.node_max_id = node_max_id |
| 35 | self.edge_max_id = edge_max_id |
| 36 | self.ent_dim = ent_dim |
| 37 | self.rel_dim = rel_dim |
| 38 | self.num_negs = num_negs |
| 39 | self.metric_name = metric_name |
| 40 | self.corrupt = corrupt |
| 41 | self.entity_encoder = tf_euler.utils.layers.Embedding(node_max_id+1, |
| 42 | ent_dim) |
| 43 | self.relation_encoder = tf_euler.utils.layers.Embedding(edge_max_id+1, |
| 44 | rel_dim) |
| 45 | self.margin = margin |
| 46 | self.metric_class = tf_euler.utils.metrics.get(metric_name) |
| 47 | self.l2_regular = l2_regular |
| 48 | self.regular_param = regular_param |
| 49 | |
| 50 | def generate_negative(self, batch_size): |
| 51 | return tf_euler.sample_node(batch_size * self.num_negs, self.node_type) |