Buckets:
| 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.