MCPcopy Create free account
hub / github.com/alibaba/euler / __init__

Method __init__

examples/distmult/distmult.py:22–48  ·  view source on GitHub ↗
(self, node_type, edge_type,
                 node_max_id, edge_max_id,
                 ent_dim, rel_dim,
                 num_negs=5, margin=1,
                 metric_name='mrr',
                 corrupt='both',
                 l2_regular=False,
                 regular_param=0.0001,
                 **kwargs)

Source from the content-addressed store, hash-verified

20
21class 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)

Callers

nothing calls this directly

Calls 1

getMethod · 0.80

Tested by

no test coverage detected