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