| 206 | EPS = 1e-15 |
| 207 | |
| 208 | def __init__(self, in_dim, emb_dim, num_layer, kernel='gcn', drop_ratio=0, |
| 209 | act='relu', norm=None, concat=False, mask_ratio=0.5, replace_ratio=0): |
| 210 | super().__init__() |
| 211 | |
| 212 | self.emb_dim = emb_dim if not concat else num_layer * emb_dim |
| 213 | self.mask_ratio = mask_ratio |
| 214 | self.replace_ratio = replace_ratio |
| 215 | |
| 216 | self.encoder = Encoder(in_dim, emb_dim, num_layer, kernel, drop_ratio, act, norm, |
| 217 | concat=False, last_act=True, aggr='mean') |
| 218 | self.decoder = Encoder(emb_dim, in_dim, 1, kernel, drop_ratio, act, norm=None, |
| 219 | concat=False, last_act=False, aggr='mean') |
| 220 | |
| 221 | self.encoder_mask_token = nn.Parameter(torch.zeros(1, in_dim)) |
| 222 | self.encoder_to_decoder = nn.Linear(self.emb_dim, emb_dim, bias=False) |
| 223 | |
| 224 | |
| 225 | def forward(self, x, edge_index, edge_weigt=None, batch=None): |