| 40 | class SharedGroupGNNNet(object): |
| 41 | |
| 42 | def __init__(self, conv, group_flow, dims, |
| 43 | group_fanouts, group_metapath, |
| 44 | add_self_loops=True): |
| 45 | conv_class = utils.get_conv_class(conv) |
| 46 | self.convs = [] |
| 47 | for dim in dims[:-1]: |
| 48 | self.convs.append(self.get_conv(conv_class, dim)) |
| 49 | group_flow_class = [utils.get_flow_class(flow) for flow in group_flow] |
| 50 | self.fc = tf.layers.Dense(dims[-1]) |
| 51 | if 'whole' in group_flow: |
| 52 | raise ValueError('Group GNN does not support whole dataflow') |
| 53 | self.group_sampler = [flow_class(fanouts, metapath, add_self_loops) |
| 54 | for flow_class, fanouts, metapath |
| 55 | in zip(group_flow_class, group_fanouts, group_metapath)] |
| 56 | |
| 57 | def get_conv(self, conv_class, dim): |
| 58 | return conv_class(dim) |