| 37 | |
| 38 | class BaseAggregator(layers.Layer): |
| 39 | def __init__(self, dim, activation=tf.nn.relu, concat=False, **kwargs): |
| 40 | super(BaseAggregator, self).__init__(**kwargs) |
| 41 | if concat: |
| 42 | if dim % 2: |
| 43 | raise ValueError('dim must be divided exactly ' |
| 44 | 'by 2 if concat is True.') |
| 45 | dim //= 2 |
| 46 | self.concat = concat |
| 47 | self.self_layer = layers.Dense(dim, |
| 48 | activation=activation, |
| 49 | use_bias=False) |
| 50 | self.neigh_layer = layers.Dense(dim, |
| 51 | activation=activation, |
| 52 | use_bias=False) |
| 53 | |
| 54 | def call(self, inputs): |
| 55 | self_embedding, neigh_embedding = inputs |