|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from torch_geometric.nn import RGCNConv |
|
|
|
|
| class ContextEnhancedRGCN(nn.Module): |
| def __init__( |
| self, |
| num_nodes, |
| num_relations, |
| event_node_idx_cpu, |
| event_x_cpu, |
| context_dim=768, |
| node_emb_dim=16, |
| hidden_dim=16, |
| dropout=0.2, |
| num_bases=2 |
| ): |
| super().__init__() |
|
|
| self.node_emb = nn.Embedding(num_nodes, node_emb_dim) |
| self.context_proj = nn.Linear(context_dim, node_emb_dim) |
|
|
| self.register_buffer( |
| "event_node_idx", |
| event_node_idx_cpu.to(torch.long) |
| ) |
|
|
| self.register_buffer( |
| "event_x", |
| torch.tensor(event_x_cpu, dtype=torch.float) |
| ) |
|
|
| self.rgcn1 = RGCNConv( |
| node_emb_dim, |
| hidden_dim, |
| num_relations=num_relations, |
| num_bases=num_bases |
| ) |
|
|
| self.rgcn2 = RGCNConv( |
| hidden_dim, |
| hidden_dim, |
| num_relations=num_relations, |
| num_bases=num_bases |
| ) |
|
|
| self.dropout_layer = nn.Dropout(dropout) |
|
|
| def initial_features(self): |
| x = self.node_emb.weight.clone() |
| event_feat = self.context_proj(self.event_x) |
| x[self.event_node_idx] = x[self.event_node_idx] + event_feat |
| return x |
|
|
| def encode(self, edge_index, edge_type): |
| x = self.initial_features() |
| x = self.rgcn1(x, edge_index, edge_type) |
| x = F.relu(x) |
| x = self.dropout_layer(x) |
| x = self.rgcn2(x, edge_index, edge_type) |
| return x |
|
|
| def forward(self, event_idx, candidate_seed_idx, edge_index, edge_type): |
| z = self.encode(edge_index, edge_type) |
|
|
| event_z = z[event_idx] |
| seed_z = z[candidate_seed_idx] |
|
|
| event_z = event_z.unsqueeze(1).expand_as(seed_z) |
|
|
| scores = (event_z * seed_z).sum(dim=-1) |
| return scores |
|
|