| 154 | EPS = 1e-15 |
| 155 | |
| 156 | def __init__(self, in_dim, emb_dim, num_layer, kernel='gcn', drop_ratio=0, |
| 157 | act='relu', norm=None): |
| 158 | super().__init__() |
| 159 | |
| 160 | self.emd_dim = emb_dim |
| 161 | |
| 162 | self.encoder = Encoder(in_dim, emb_dim, num_layer, kernel, drop_ratio, act, norm, concat=False, last_act=False) |
| 163 | # self.pred_head = LinearPred(emb_dim, emb_dim, 2, 2) |
| 164 | |
| 165 | self.weight = nn.Parameter(torch.empty(emb_dim, emb_dim)) |
| 166 | # just try |
| 167 | uniform(self.emd_dim, self.weight) |
| 168 | # self.weight = nn.Parameter(torch.empty(768, 768)) |
| 169 | # uniform(768, self.weight) |
| 170 | |
| 171 | |
| 172 | def forward(self, x, edge_index, edge_weigt=None, batch=None): |