tahamajs's picture
download
raw
25.5 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
import numpy as np
from typing import Dict, List, Tuple, Optional
import random
from collections import defaultdict
import matplotlib.pyplot as plt
from sklearn.decomposition import PCA
from sklearn.metrics.pairwise import cosine_similarity
import sys
import os
from torch_geometric.nn import GCNConv, GATConv, GraphSAGE
from torch_geometric.data import Data, DataLoader
import math
sys.path.append(os.path.join(os.path.dirname(__file__), ".."))
from core.knowledge_graph import KnowledgeGraph
class KGEmbeddingModel(nn.Module):
def __init__(self, num_entities: int, num_relations: int, embedding_dim: int):
super(KGEmbeddingModel, self).__init__()
self.num_entities = num_entities
self.num_relations = num_relations
self.embedding_dim = embedding_dim
self.entity_embeddings = nn.Embedding(num_entities, embedding_dim)
self.relation_embeddings = nn.Embedding(num_relations, embedding_dim)
self._init_embeddings()
def _init_embeddings(self):
nn.init.uniform_(self.entity_embeddings.weight, -1, 1)
nn.init.uniform_(self.relation_embeddings.weight, -1, 1)
def forward(self, triples):
raise NotImplementedError
def get_embeddings(self):
return self.entity_embeddings.weight.data, self.relation_embeddings.weight.data
class TransE(KGEmbeddingModel):
def __init__(
self, num_entities: int, num_relations: int, embedding_dim: int, norm_p: int = 2
):
super(TransE, self).__init__(num_entities, num_relations, embedding_dim)
self.norm_p = norm_p
def forward(self, triples):
heads = self.entity_embeddings(triples[:, 0])
relations = self.relation_embeddings(triples[:, 1])
tails = self.entity_embeddings(triples[:, 2])
scores = heads + relations - tails
return -torch.norm(scores, p=self.norm_p, dim=1)
def predict_tail(self, head: int, relation: int, k: int = 10):
self.eval()
with torch.no_grad():
head_emb = self.entity_embeddings(torch.tensor([head]))
rel_emb = self.relation_embeddings(torch.tensor([relation]))
all_tails = torch.arange(self.num_entities)
tail_embs = self.entity_embeddings(all_tails)
scores = head_emb + rel_emb - tail_embs
distances = torch.norm(scores, p=self.norm_p, dim=1)
top_k_indices = torch.topk(distances, k=k, largest=False).indices
return top_k_indices.tolist()
class DistMult(KGEmbeddingModel):
def forward(self, triples):
heads = self.entity_embeddings(triples[:, 0])
relations = self.relation_embeddings(triples[:, 1])
tails = self.entity_embeddings(triples[:, 2])
scores = torch.sum(heads * relations * tails, dim=1)
return scores
class ComplEx(KGEmbeddingModel):
def __init__(self, num_entities: int, num_relations: int, embedding_dim: int):
super(ComplEx, self).__init__(num_entities, num_relations, embedding_dim)
self.entity_embeddings = nn.Embedding(num_entities, embedding_dim * 2)
self.relation_embeddings = nn.Embedding(num_relations, embedding_dim * 2)
self._init_embeddings()
def forward(self, triples):
heads = self.entity_embeddings(triples[:, 0])
relations = self.relation_embeddings(triples[:, 1])
tails = self.entity_embeddings(triples[:, 2])
dim = heads.size(1) // 2
heads_real, heads_imag = heads[:, :dim], heads[:, dim:]
relations_real, relations_imag = relations[:, :dim], relations[:, dim:]
tails_real, tails_imag = tails[:, :dim], tails[:, dim:]
score_real = (
heads_real * relations_real * tails_real
+ heads_imag * relations_real * tails_imag
+ heads_real * relations_imag * tails_imag
- heads_imag * relations_imag * tails_real
)
return torch.sum(score_real, dim=1)
class KGEmbeddingTrainer:
def __init__(self, model: KGEmbeddingModel, kg: KnowledgeGraph):
self.model = model
self.kg = kg
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.model.to(self.device)
self.entity_to_idx = {
entity_id: idx for idx, entity_id in enumerate(kg.entities.keys())
}
self.relation_to_idx = {
rel_id: idx for idx, rel_id in enumerate(kg.relations.keys())
}
self.idx_to_entity = {
idx: entity_id for entity_id, idx in self.entity_to_idx.items()
}
self.idx_to_relation = {
idx: rel_id for rel_id, idx in self.relation_to_idx.items()
}
self.train_triples = self._prepare_triples()
def _prepare_triples(self):
triples = []
for triple in self.kg.triples:
if (
triple.subject in self.entity_to_idx
and triple.predicate in self.relation_to_idx
and triple.object in self.entity_to_idx
):
head_idx = self.entity_to_idx[triple.subject]
rel_idx = self.relation_to_idx[triple.predicate]
tail_idx = self.entity_to_idx[triple.object]
triples.append([head_idx, rel_idx, tail_idx])
return torch.tensor(triples, dtype=torch.long)
def generate_negative_samples(
self, positive_triples: torch.Tensor, num_negatives: int = 1
) -> torch.Tensor:
negative_triples = []
for _ in range(num_negatives):
corrupted_triples = positive_triples.clone()
for i in range(corrupted_triples.size(0)):
if random.random() < 0.5:
corrupted_h = random.randint(0, self.model.num_entities - 1)
while corrupted_h == corrupted_triples[i, 0]:
corrupted_h = random.randint(0, self.model.num_entities - 1)
corrupted_triples[i, 0] = corrupted_h
else:
corrupted_t = random.randint(0, self.model.num_entities - 1)
while corrupted_t == corrupted_triples[i, 2]:
corrupted_t = random.randint(0, self.model.num_entities - 1)
corrupted_triples[i, 2] = corrupted_t
negative_triples.append(corrupted_triples)
return torch.stack(negative_triples)
def train_step(
self, batch_size: int = 128, margin: float = 1.0, num_negatives: int = 1
) -> float:
self.model.train()
indices = torch.randperm(self.train_triples.size(0))[:batch_size]
positive_batch = self.train_triples[indices]
negative_batch = self.generate_negative_samples(positive_batch, num_negatives)
positive_batch = positive_batch.to(self.device)
negative_batch = negative_batch.to(self.device)
positive_scores = self.model(positive_batch)
negative_scores = self.model(negative_batch.view(-1, 3))
positive_scores = positive_scores.unsqueeze(1).expand(-1, num_negatives)
loss = F.margin_ranking_loss(
positive_scores.view(-1),
negative_scores,
torch.ones_like(negative_scores),
margin=margin,
)
return loss.item()
def train(self, epochs: int = 100, lr: float = 0.001, batch_size: int = 128):
optimizer = optim.Adam(self.model.parameters(), lr=lr)
losses = []
print(f"🚀 Training {type(self.model).__name__} model...")
print(f" Entities: {self.model.num_entities}")
print(f" Relations: {self.model.num_relations}")
print(f" Training triples: {len(self.train_triples)}")
print(f" Device: {self.device}")
for epoch in range(epochs):
loss = self.train_step(batch_size)
losses.append(loss)
if epoch % 20 == 0:
print(f" Epoch {epoch:3d}: Loss = {loss:.4f}")
print(f"✅ Training completed! Final loss: {losses[-1]:.4f}")
return losses
def evaluate_link_prediction(self, test_triples: torch.Tensor = None, k: int = 10):
if test_triples is None:
test_indices = torch.randperm(self.train_triples.size(0))[:100]
test_triples = self.train_triples[test_indices]
self.model.eval()
hits_at_k = 0
mean_rank = 0
with torch.no_grad():
for triple in test_triples:
head, rel, tail = triple[0], triple[1], triple[2]
all_triples = torch.zeros(
(self.model.num_entities, 3), dtype=torch.long
)
all_triples[:, 0] = head
all_triples[:, 1] = rel
all_triples[:, 2] = torch.arange(self.model.num_entities)
scores = self.model(all_triples.to(self.device))
sorted_indices = torch.argsort(scores, descending=True)
rank = (sorted_indices == tail).nonzero(as_tuple=True)[0].item() + 1
mean_rank += rank
if rank <= k:
hits_at_k += 1
hits_at_k_score = hits_at_k / len(test_triples)
mean_rank_score = mean_rank / len(test_triples)
results = {
"hits@k": hits_at_k_score,
"mean_rank": mean_rank_score,
"num_test_triples": len(test_triples),
}
return results
def find_similar_entities(self, entity_id: str, k: int = 5):
if entity_id not in self.entity_to_idx:
return []
self.model.eval()
with torch.no_grad():
entity_idx = self.entity_to_idx[entity_id]
entity_embed = self.model.entity_embeddings.weight[entity_idx]
all_embeds = self.model.entity_embeddings.weight
similarities = F.cosine_similarity(entity_embed.unsqueeze(0), all_embeds)
top_k_indices = torch.topk(similarities, k=k + 1).indices[1:]
similar_entities = []
for idx in top_k_indices:
similar_entity_id = self.idx_to_entity[idx.item()]
similarity = similarities[idx].item()
similar_entities.append((similar_entity_id, similarity))
return similar_entities
def visualize_embeddings(self, method: str = "pca", max_entities: int = 50):
self.model.eval()
with torch.no_grad():
entity_embeds, relation_embeds = self.model.get_embeddings()
entity_indices = list(range(min(max_entities, len(self.entity_to_idx))))
selected_embeds = entity_embeds[entity_indices]
if method == "pca":
pca = PCA(n_components=2)
embeddings_2d = pca.fit_transform(selected_embeds.cpu().numpy())
plt.figure(figsize=(12, 8))
colors = []
entity_names = []
for idx in entity_indices:
entity_id = self.idx_to_entity[idx]
entity = self.kg.entities[entity_id]
entity_names.append(entity.name)
if entity.entity_type == "person":
colors.append("red")
elif entity.entity_type == "concept":
colors.append("blue")
elif entity.entity_type == "field":
colors.append("green")
else:
colors.append("orange")
plt.scatter(
embeddings_2d[:, 0], embeddings_2d[:, 1], c=colors, alpha=0.7, s=100
)
for i, name in enumerate(entity_names):
plt.annotate(
name,
(embeddings_2d[i, 0], embeddings_2d[i, 1]),
xytext=(5, 5),
textcoords="offset points",
fontsize=8,
)
plt.xlabel(f"PC1 ({pca.explained_variance_ratio_[0]:.2%} variance)")
plt.ylabel(f"PC2 ({pca.explained_variance_ratio_[1]:.2%} variance)")
plt.title(f"{type(self.model).__name__} Entity Embeddings (PCA)")
plt.grid(True, alpha=0.3)
from matplotlib.patches import Patch
legend_elements = [
Patch(facecolor="red", label="Person"),
Patch(facecolor="blue", label="Concept"),
Patch(facecolor="green", label="Field"),
Patch(facecolor="orange", label="Other"),
]
plt.legend(handles=legend_elements)
plt.tight_layout()
plt.show()
return embeddings_2d if method == "pca" else None
# ============================================================================
# ADVANCED EMBEDDING MODELS
# ============================================================================
class SimplE(KGEmbeddingModel):
"""
SimplE: Simple Embedding for Link Prediction in Knowledge Graphs
Paper: https://arxiv.org/abs/1802.04868
SimplE learns separate embeddings for head and tail entities, making it
more expressive while remaining interpretable and efficient.
"""
def __init__(self, num_entities: int, num_relations: int, embedding_dim: int):
super(SimplE, self).__init__(num_entities, num_relations, embedding_dim)
# Separate embeddings for head and tail
self.entity_head_embeddings = nn.Embedding(num_entities, embedding_dim)
self.entity_tail_embeddings = nn.Embedding(num_entities, embedding_dim)
# Inverse relation embeddings
self.relation_inv_embeddings = nn.Embedding(num_relations, embedding_dim)
self._init_simple_embeddings()
def _init_simple_embeddings(self):
nn.init.uniform_(self.entity_head_embeddings.weight, -1, 1)
nn.init.uniform_(self.entity_tail_embeddings.weight, -1, 1)
nn.init.uniform_(self.relation_embeddings.weight, -1, 1)
nn.init.uniform_(self.relation_inv_embeddings.weight, -1, 1)
def forward(self, triples):
heads = self.entity_head_embeddings(triples[:, 0])
relations = self.relation_embeddings(triples[:, 1])
tails = self.entity_tail_embeddings(triples[:, 2])
# Inverse direction
heads_inv = self.entity_head_embeddings(triples[:, 2])
relations_inv = self.relation_inv_embeddings(triples[:, 1])
tails_inv = self.entity_tail_embeddings(triples[:, 0])
# Average of both directions
score_forward = torch.sum(heads * relations * tails, dim=1)
score_backward = torch.sum(heads_inv * relations_inv * tails_inv, dim=1)
return (score_forward + score_backward) / 2.0
class RotatE(KGEmbeddingModel):
"""
RotatE: Knowledge Graph Embedding by Relational Rotation in Complex Space
Paper: https://arxiv.org/abs/1902.10197
RotatE models relations as rotations in complex space, allowing it to
model symmetric, antisymmetric, inverse, and composition patterns.
"""
def __init__(
self,
num_entities: int,
num_relations: int,
embedding_dim: int,
gamma: float = 12.0,
):
super(RotatE, self).__init__(num_entities, num_relations, embedding_dim)
self.gamma = gamma
# Complex embeddings (real and imaginary parts)
self.entity_embeddings = nn.Embedding(num_entities, embedding_dim * 2)
self.relation_embeddings = nn.Embedding(num_relations, embedding_dim)
self.embedding_range = (self.gamma + 2.0) / embedding_dim
nn.init.uniform_(
self.entity_embeddings.weight, -self.embedding_range, self.embedding_range
)
nn.init.uniform_(
self.relation_embeddings.weight, -self.embedding_range, self.embedding_range
)
def forward(self, triples):
heads = self.entity_embeddings(triples[:, 0])
relations = self.relation_embeddings(triples[:, 1])
tails = self.entity_embeddings(triples[:, 2])
dim = self.embedding_dim
# Split into real and imaginary parts
head_re, head_im = heads[:, :dim], heads[:, dim:]
tail_re, tail_im = tails[:, :dim], tails[:, dim:]
# Relation as rotation (convert to phase)
phase_relation = relations / (self.embedding_range / math.pi)
rel_re = torch.cos(phase_relation)
rel_im = torch.sin(phase_relation)
# Complex multiplication: head * relation
rotated_re = head_re * rel_re - head_im * rel_im
rotated_im = head_re * rel_im + head_im * rel_re
# Distance to tail
score_re = rotated_re - tail_re
score_im = rotated_im - tail_im
score = torch.stack([score_re, score_im], dim=0)
score = score.norm(dim=0).sum(dim=1)
return self.gamma - score
class ConvE(KGEmbeddingModel):
"""
ConvE: Convolutional 2D Knowledge Graph Embeddings
Paper: https://arxiv.org/abs/1707.01476
ConvE uses 2D convolutions over entity and relation embeddings to
model complex interaction patterns.
"""
def __init__(
self,
num_entities: int,
num_relations: int,
embedding_dim: int,
input_drop: float = 0.2,
hidden_drop: float = 0.3,
feat_drop: float = 0.2,
):
super(ConvE, self).__init__(num_entities, num_relations, embedding_dim)
self.inp_drop = nn.Dropout(input_drop)
self.hidden_drop = nn.Dropout(hidden_drop)
self.feature_drop = nn.Dropout(feat_drop)
# Reshape dimensions for 2D convolution
self.embedding_height = 10
self.embedding_width = embedding_dim // self.embedding_height
# Convolutional layers
self.conv1 = nn.Conv2d(1, 32, (3, 3), padding=1)
self.bn0 = nn.BatchNorm2d(1)
self.bn1 = nn.BatchNorm2d(32)
self.bn2 = nn.BatchNorm1d(embedding_dim)
# Fully connected layer
flat_sz = self.embedding_height * self.embedding_width * 32
self.fc = nn.Linear(flat_sz, embedding_dim)
def forward(self, triples):
batch_size = triples.size(0)
heads = self.entity_embeddings(triples[:, 0])
relations = self.relation_embeddings(triples[:, 1])
tails = self.entity_embeddings(triples[:, 2])
# Reshape for 2D convolution
heads = heads.view(-1, 1, self.embedding_height, self.embedding_width)
relations = relations.view(-1, 1, self.embedding_height, self.embedding_width)
# Stack head and relation
stacked = torch.cat([heads, relations], dim=2)
stacked = self.bn0(stacked)
stacked = self.inp_drop(stacked)
# Convolution
x = self.conv1(stacked)
x = self.bn1(x)
x = F.relu(x)
x = self.feature_drop(x)
# Flatten and project
x = x.view(batch_size, -1)
x = self.fc(x)
x = self.hidden_drop(x)
x = self.bn2(x)
x = F.relu(x)
# Score with tail
scores = torch.sum(x * tails, dim=1)
return scores
class GraphAttentionEmbedding(nn.Module):
"""
Graph Attention Network for Knowledge Graph Embeddings
Uses attention mechanism to aggregate neighbor information
"""
def __init__(
self,
num_entities: int,
num_relations: int,
embedding_dim: int,
num_heads: int = 4,
dropout: float = 0.2,
):
super(GraphAttentionEmbedding, self).__init__()
self.num_entities = num_entities
self.num_relations = num_relations
self.embedding_dim = embedding_dim
self.num_heads = num_heads
# Entity and relation embeddings
self.entity_embeddings = nn.Embedding(num_entities, embedding_dim)
self.relation_embeddings = nn.Embedding(num_relations, embedding_dim)
# Multi-head attention layers
self.attention_layers = nn.ModuleList(
[
nn.MultiheadAttention(embedding_dim, num_heads, dropout=dropout)
for _ in range(2)
]
)
# Feed-forward network
self.ffn = nn.Sequential(
nn.Linear(embedding_dim, embedding_dim * 4),
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(embedding_dim * 4, embedding_dim),
)
# Layer normalization
self.layer_norms = nn.ModuleList(
[nn.LayerNorm(embedding_dim) for _ in range(4)]
)
self.dropout = nn.Dropout(dropout)
self._init_embeddings()
def _init_embeddings(self):
nn.init.xavier_uniform_(self.entity_embeddings.weight)
nn.init.xavier_uniform_(self.relation_embeddings.weight)
def forward(self, triples, entity_neighbors=None):
"""
Args:
triples: (batch_size, 3) tensor of (head, relation, tail) indices
entity_neighbors: Optional dict mapping entity indices to their neighbors
"""
heads = self.entity_embeddings(triples[:, 0])
relations = self.relation_embeddings(triples[:, 1])
tails = self.entity_embeddings(triples[:, 2])
# Apply attention if neighbors are provided
if entity_neighbors is not None:
# Aggregate neighbor information for heads
heads = self._aggregate_neighbors(heads, triples[:, 0], entity_neighbors)
tails = self._aggregate_neighbors(tails, triples[:, 2], entity_neighbors)
# Combine head, relation, and tail
combined = heads + relations
# Apply attention layers
for i, attn_layer in enumerate(self.attention_layers):
# Self-attention
attn_out, _ = attn_layer(
combined.unsqueeze(0), combined.unsqueeze(0), combined.unsqueeze(0)
)
attn_out = attn_out.squeeze(0)
# Residual connection and layer norm
combined = self.layer_norms[i * 2](combined + self.dropout(attn_out))
# Feed-forward
ffn_out = self.ffn(combined)
combined = self.layer_norms[i * 2 + 1](combined + self.dropout(ffn_out))
# Score
scores = torch.sum(combined * tails, dim=1)
return scores
def _aggregate_neighbors(self, entity_embeds, entity_indices, entity_neighbors):
"""Aggregate information from neighboring entities"""
# This is a simplified version - in practice, you'd want to batch this more efficiently
return entity_embeds
class TransformerKGEmbedding(nn.Module):
"""
Transformer-based Knowledge Graph Embedding
Uses transformer architecture to model complex patterns in KG
"""
def __init__(
self,
num_entities: int,
num_relations: int,
embedding_dim: int,
num_layers: int = 3,
num_heads: int = 8,
dropout: float = 0.1,
):
super(TransformerKGEmbedding, self).__init__()
self.num_entities = num_entities
self.num_relations = num_relations
self.embedding_dim = embedding_dim
# Embeddings
self.entity_embeddings = nn.Embedding(num_entities, embedding_dim)
self.relation_embeddings = nn.Embedding(num_relations, embedding_dim)
self.position_embeddings = nn.Embedding(
3, embedding_dim
) # For h, r, t positions
# Transformer encoder
encoder_layer = nn.TransformerEncoderLayer(
d_model=embedding_dim,
nhead=num_heads,
dim_feedforward=embedding_dim * 4,
dropout=dropout,
batch_first=True,
)
self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)
# Output projection
self.output_proj = nn.Sequential(
nn.Linear(embedding_dim, embedding_dim),
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(embedding_dim, 1),
)
self._init_embeddings()
def _init_embeddings(self):
nn.init.xavier_uniform_(self.entity_embeddings.weight)
nn.init.xavier_uniform_(self.relation_embeddings.weight)
nn.init.xavier_uniform_(self.position_embeddings.weight)
def forward(self, triples):
batch_size = triples.size(0)
# Get embeddings
heads = self.entity_embeddings(triples[:, 0])
relations = self.relation_embeddings(triples[:, 1])
tails = self.entity_embeddings(triples[:, 2])
# Add positional embeddings
pos_ids = (
torch.arange(3, device=triples.device).unsqueeze(0).expand(batch_size, -1)
)
pos_embeds = self.position_embeddings(pos_ids)
# Stack as sequence: [head, relation, tail]
sequence = torch.stack([heads, relations, tails], dim=1) # (batch, 3, dim)
sequence = sequence + pos_embeds
# Apply transformer
transformed = self.transformer(sequence) # (batch, 3, dim)
# Pool the sequence (mean pooling)
pooled = transformed.mean(dim=1) # (batch, dim)
# Project to score
scores = self.output_proj(pooled).squeeze(-1)
return scores
def get_embeddings(self):
return self.entity_embeddings.weight.data, self.relation_embeddings.weight.data

Xet Storage Details

Size:
25.5 kB
·
Xet hash:
572c5516c56b6d60bb9a500f90b33a73c54e1485d51f9830ed389a35fc48c7e3

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.