tahamajs's picture
download
raw
14.1 kB
# Execute Cell 3 - Core Knowledge Graph Data Structures
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
import numpy as np
import pandas as pd
import matplotlib
matplotlib.use("Agg") # Use non-interactive backend
import matplotlib.pyplot as plt
import seaborn as sns
import networkx as nx
from typing import List, Dict, Tuple, Optional, Union, Set
from dataclasses import dataclass
from collections import defaultdict, deque
import json
import time
import random
from abc import ABC, abstractmethod
import pickle
from itertools import combinations, product
import warnings
warnings.filterwarnings("ignore")
# Set random seeds for reproducibility
torch.manual_seed(42)
np.random.seed(42)
random.seed(42)
# Set up plotting
plt.style.use("seaborn-v0_8")
sns.set_palette("husl")
# Device setup
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
@dataclass
class Triple:
"""Represents a knowledge graph triple (subject, predicate, object)"""
subject: str
predicate: str
object: str
def __str__(self):
return f"({self.subject}, {self.predicate}, {self.object})"
def __hash__(self):
return hash((self.subject, self.predicate, self.object))
@dataclass
class Entity:
"""Represents an entity in the knowledge graph"""
id: str
name: str
type: str = "unknown"
attributes: Dict[str, str] = None
def __post_init__(self):
if self.attributes is None:
self.attributes = {}
@dataclass
class Relation:
"""Represents a relation type in the knowledge graph"""
id: str
name: str
domain: str = "entity" # domain entity type
range: str = "entity" # range entity type
inverse: str = None # inverse relation
class KnowledgeGraph:
"""Main Knowledge Graph class"""
def __init__(self, name: str = "KG"):
self.name = name
self.entities = {} # entity_id -> Entity
self.relations = {} # relation_id -> Relation
self.triples = set() # set of Triple objects
self.adjacency = defaultdict(
lambda: defaultdict(set)
) # subject -> predicate -> {objects}
self.reverse_adjacency = defaultdict(
lambda: defaultdict(set)
) # object -> predicate -> {subjects}
# Statistics
self.stats = {"num_entities": 0, "num_relations": 0, "num_triples": 0}
def add_entity(
self,
entity_id: str,
name: str,
entity_type: str = "unknown",
attributes: Dict = None,
):
"""Add an entity to the knowledge graph"""
entity = Entity(entity_id, name, entity_type, attributes or {})
self.entities[entity_id] = entity
self.stats["num_entities"] = len(self.entities)
return entity
def add_relation(
self,
relation_id: str,
name: str,
domain: str = "entity",
range: str = "entity",
inverse: str = None,
):
"""Add a relation type to the knowledge graph"""
relation = Relation(relation_id, name, domain, range, inverse)
self.relations[relation_id] = relation
self.stats["num_relations"] = len(self.relations)
return relation
def add_triple(self, subject: str, predicate: str, object: str):
"""Add a triple to the knowledge graph"""
triple = Triple(subject, predicate, object)
if triple not in self.triples:
self.triples.add(triple)
# Update adjacency lists
self.adjacency[subject][predicate].add(object)
self.reverse_adjacency[object][predicate].add(subject)
# Update statistics
self.stats["num_triples"] = len(self.triples)
# Ensure entities exist (auto-create if needed)
if subject not in self.entities:
self.add_entity(subject, subject)
if object not in self.entities:
self.add_entity(object, object)
if predicate not in self.relations:
self.add_relation(predicate, predicate)
return triple
def get_neighbors(self, entity_id: str, relation: str = None) -> Set[str]:
"""Get neighbors of an entity, optionally filtered by relation"""
if relation:
return self.adjacency[entity_id].get(relation, set())
else:
neighbors = set()
for rel_neighbors in self.adjacency[entity_id].values():
neighbors.update(rel_neighbors)
return neighbors
def get_relations(self, subject: str, object: str) -> Set[str]:
"""Get all relations between two entities"""
relations = set()
for predicate, objects in self.adjacency[subject].items():
if object in objects:
relations.add(predicate)
return relations
def query_triples(
self, subject: str = None, predicate: str = None, object: str = None
) -> List[Triple]:
"""Query triples with optional constraints"""
result = []
for triple in self.triples:
if (
(subject is None or triple.subject == subject)
and (predicate is None or triple.predicate == predicate)
and (object is None or triple.object == object)
):
result.append(triple)
return result
def get_subgraph(self, entities: List[str], max_hops: int = 1) -> "KnowledgeGraph":
"""Extract a subgraph containing specified entities and their neighbors"""
subgraph = KnowledgeGraph(f"{self.name}_subgraph")
# Add initial entities
included_entities = set(entities)
# Expand to include neighbors within max_hops
for hop in range(max_hops):
new_entities = set()
for entity in included_entities:
neighbors = self.get_neighbors(entity)
new_entities.update(neighbors)
included_entities.update(new_entities)
# Add entities to subgraph
for entity_id in included_entities:
if entity_id in self.entities:
entity = self.entities[entity_id]
subgraph.add_entity(
entity_id, entity.name, entity.type, entity.attributes
)
# Add relations
for relation_id, relation in self.relations.items():
subgraph.add_relation(
relation_id,
relation.name,
relation.domain,
relation.range,
relation.inverse,
)
# Add triples involving included entities
for triple in self.triples:
if (
triple.subject in included_entities
and triple.object in included_entities
):
subgraph.add_triple(triple.subject, triple.predicate, triple.object)
return subgraph
def visualize(self, max_entities: int = 20, layout: str = "spring"):
"""Visualize the knowledge graph"""
G = nx.Graph()
# Limit entities for visualization
entities_to_show = list(self.entities.keys())[:max_entities]
# Add nodes
for entity_id in entities_to_show:
entity = self.entities[entity_id]
G.add_node(entity_id, label=entity.name, type=entity.type)
# Add edges
edge_labels = {}
for triple in self.triples:
if triple.subject in entities_to_show and triple.object in entities_to_show:
G.add_edge(triple.subject, triple.object)
edge_key = (triple.subject, triple.object)
if edge_key not in edge_labels:
edge_labels[edge_key] = []
edge_labels[edge_key].append(triple.predicate)
# Create layout
if layout == "spring":
pos = nx.spring_layout(G, k=2, iterations=50)
elif layout == "circular":
pos = nx.circular_layout(G)
else:
pos = nx.random_layout(G)
# Draw the graph
plt.figure(figsize=(15, 10))
# Draw nodes
nx.draw_networkx_nodes(
G, pos, node_color="lightblue", node_size=1000, alpha=0.7
)
# Draw edges
nx.draw_networkx_edges(G, pos, alpha=0.5, edge_color="gray")
# Add node labels
node_labels = {
node: self.entities[node].name[:10]
for node in G.nodes()
if node in self.entities
}
nx.draw_networkx_labels(G, pos, labels=node_labels, font_size=8)
# Add edge labels (relations)
edge_label_dict = {}
for (u, v), relations in edge_labels.items():
edge_label_dict[(u, v)] = ",".join(relations[:2]) # Show max 2 relations
nx.draw_networkx_edge_labels(G, pos, edge_labels=edge_label_dict, font_size=6)
plt.title(
f"Knowledge Graph: {self.name}\n{self.stats['num_entities']} entities, {self.stats['num_triples']} triples"
)
plt.axis("off")
plt.tight_layout()
# Save instead of show
plt.savefig("kg_viz.png", dpi=150, bbox_inches="tight")
plt.close()
def print_stats(self):
"""Print knowledge graph statistics"""
print(f"📊 Knowledge Graph Statistics: {self.name}")
print(f" Entities: {self.stats['num_entities']}")
print(f" Relations: {self.stats['num_relations']}")
print(f" Triples: {self.stats['num_triples']}")
# Entity type distribution
type_counts = defaultdict(int)
for entity in self.entities.values():
type_counts[entity.type] += 1
print(f"\n Entity Types:")
for entity_type, count in sorted(type_counts.items()):
print(f" {entity_type}: {count}")
# Relation distribution
relation_counts = defaultdict(int)
for triple in self.triples:
relation_counts[triple.predicate] += 1
print(f"\n Top Relations:")
for relation, count in sorted(
relation_counts.items(), key=lambda x: x[1], reverse=True
)[:5]:
print(f" {relation}: {count}")
def get_statistics(self):
"""Get knowledge graph statistics as dict"""
return {
"Entities": self.stats["num_entities"],
"Relations": self.stats["num_relations"],
"Triples": self.stats["num_triples"],
}
# Create a sample knowledge graph for demonstration
def create_sample_kg() -> KnowledgeGraph:
"""Create a sample knowledge graph about scientists and their work"""
kg = KnowledgeGraph("Scientists_KG")
# Add entities
scientists = [
("einstein", "Albert Einstein", "person"),
("newton", "Isaac Newton", "person"),
("curie", "Marie Curie", "person"),
("darwin", "Charles Darwin", "person"),
("tesla", "Nikola Tesla", "person"),
]
countries = [
("germany", "Germany", "country"),
("england", "England", "country"),
("poland", "Poland", "country"),
("usa", "United States", "country"),
("serbia", "Serbia", "country"),
]
theories = [
("relativity", "Theory of Relativity", "theory"),
("gravity", "Law of Universal Gravitation", "theory"),
("evolution", "Theory of Evolution", "theory"),
("radioactivity", "Radioactivity Theory", "theory"),
]
# Add all entities
for entity_id, name, entity_type in scientists + countries + theories:
kg.add_entity(entity_id, name, entity_type)
# Add relations
relations = [
("born_in", "born in", "person", "country"),
("developed", "developed", "person", "theory"),
("influenced", "influenced", "person", "person"),
("located_in", "located in", "country", "country"),
("studied", "studied", "person", "theory"),
]
for rel_id, name, domain, range_type in relations:
kg.add_relation(rel_id, name, domain, range_type)
# Add triples (facts)
facts = [
# Birth places
("einstein", "born_in", "germany"),
("newton", "born_in", "england"),
("curie", "born_in", "poland"),
("darwin", "born_in", "england"),
("tesla", "born_in", "serbia"),
# Developments
("einstein", "developed", "relativity"),
("newton", "developed", "gravity"),
("darwin", "developed", "evolution"),
("curie", "developed", "radioactivity"),
# Influences
("newton", "influenced", "einstein"),
("darwin", "influenced", "einstein"),
("curie", "influenced", "tesla"),
# Studies
("einstein", "studied", "gravity"),
("tesla", "studied", "relativity"),
]
for subject, predicate, object in facts:
kg.add_triple(subject, predicate, object)
return kg
# Execute the main test code
print("=" * 80)
print("🔬 Testing Knowledge Graph Implementation")
print("=" * 80)
# Create sample knowledge graph
sample_kg = create_sample_kg()
sample_kg.print_stats()
print(f"\n🔍 Sample Queries:")
print(f"Scientists born in England:")
english_scientists = sample_kg.query_triples(predicate="born_in", object="england")
for triple in english_scientists:
scientist_name = sample_kg.entities[triple.subject].name
print(f" • {scientist_name}")
print(f"\nTheories developed by Einstein:")
einstein_theories = sample_kg.query_triples(subject="einstein", predicate="developed")
for triple in einstein_theories:
theory_name = sample_kg.entities[triple.object].name
print(f" • {theory_name}")
print(f"\nEinstein's neighbors:")
einstein_neighbors = sample_kg.get_neighbors("einstein")
for neighbor in einstein_neighbors:
neighbor_name = sample_kg.entities[neighbor].name
print(f" • {neighbor_name}")
# Visualize the sample knowledge graph
sample_kg.visualize(max_entities=15)
print("\n✅ Visualization saved to kg_viz.png")
print("\n" + "=" * 80)
print("✅ Cell 3 executed successfully!")
print("=" * 80)

Xet Storage Details

Size:
14.1 kB
·
Xet hash:
faf45dc8cb1091ff990aa237a176e255d06c9704a300365521c86e7a9a08f479

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