Daniel0315's picture
Upload folder using huggingface_hub
68af900 verified
Raw
History Blame Contribute Delete
1.95 kB
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