| 108 | return latent_node_feats_list, latent_edge_feats_list |
| 109 | |
| 110 | class EdgeEncodeNet(nn.Module): |
| 111 | def __init__(self, cfg): |
| 112 | super(EdgeEncodeNet, self).__init__() |
| 113 | input_dim = cfg.IN_DIM |
| 114 | out_dim = cfg.OUT_DIM |
| 115 | fc_dims = cfg.FC_DIMS |
| 116 | norm_layer = nn.LayerNorm |
| 117 | act_func = nn.ReLU(inplace=True) |
| 118 | dropout_p = cfg.DROPOUT_P |
| 119 | layers = build_layers(input_dim, fc_dims+[out_dim], dropout_p, norm_layer, act_func) |
| 120 | layers += [nn.Linear(out_dim, out_dim)] |
| 121 | self.mlp = nn.Sequential(*layers) |
| 122 | |
| 123 | def forward(self, edge_attr): |
| 124 | return self.mlp(edge_attr) |
| 125 | |
| 126 | class EdgeUpdateNet(nn.Module): |
| 127 | def __init__(self, cfg, edge_model_in_dim): |