tahamajs's picture
download
raw
13.7 kB
#!/usr/bin/env python3
"""
Fixed notebook execution with proper method injection
"""
import json
import sys
from io import StringIO
import traceback
# Import all required libraries
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")
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 time
import random
from abc import ABC, abstractmethod
import pickle
from itertools import combinations, product
import warnings
warnings.filterwarnings("ignore")
# Set seeds
torch.manual_seed(42)
np.random.seed(42)
random.seed(42)
plt.style.use("seaborn-v0_8")
sns.set_palette("husl")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print("๐Ÿ““ Running CA10 Knowledge Graph Reasoning Notebook")
print("=" * 70)
# Load notebook
with open("notebooks/CA10.ipynb", "r") as f:
nb = json.load(f)
# Execute Cell 2 (imports)
print("\n๐Ÿ”„ Cell 2: Imports and Setup")
print("-" * 70)
code2 = "".join(nb["cells"][2]["source"])
exec(code2, globals())
print("โœ… Cell 2: Complete")
# Execute Cell 3 (KG Data Structures) with fix
print("\n๐Ÿ”„ Cell 3: Knowledge Graph Data Structures")
print("-" * 70)
# First, define the core structures
@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"
range: str = "entity"
inverse: str = None
class KnowledgeGraph:
"""Main Knowledge Graph class"""
def __init__(self, name: str = "KG"):
self.name = name
self.entities = {}
self.relations = {}
self.triples = set()
self.adjacency = defaultdict(lambda: defaultdict(set))
self.reverse_adjacency = defaultdict(lambda: defaultdict(set))
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)
self.adjacency[subject][predicate].add(object)
self.reverse_adjacency[object][predicate].add(subject)
self.stats["num_triples"] = len(self.triples)
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"""
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"""
subgraph = KnowledgeGraph(f"{self.name}_subgraph")
included_entities = set(entities)
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)
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
)
for relation_id, relation in self.relations.items():
subgraph.add_relation(
relation_id,
relation.name,
relation.domain,
relation.range,
relation.inverse,
)
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()
entities_to_show = list(self.entities.keys())[:max_entities]
for entity_id in entities_to_show:
entity = self.entities[entity_id]
G.add_node(entity_id, label=entity.name, type=entity.type)
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)
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)
plt.figure(figsize=(15, 10))
nx.draw_networkx_nodes(
G, pos, node_color="lightblue", node_size=1000, alpha=0.7
)
nx.draw_networkx_edges(G, pos, alpha=0.5, edge_color="gray")
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)
edge_label_dict = {}
for (u, v), relations in edge_labels.items():
edge_label_dict[(u, v)] = ",".join(relations[:2])
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()
plt.savefig(f"kg_viz_{self.name}.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']}")
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_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 - FIXED METHOD"""
return {
"Entities": self.stats["num_entities"],
"Relations": self.stats["num_relations"],
"Triples": self.stats["num_triples"],
}
# Now define helper function
def create_sample_kg() -> KnowledgeGraph:
"""Create a sample knowledge graph"""
kg = KnowledgeGraph("Scientists_KG")
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"),
]
for entity_id, name, entity_type in scientists + countries + theories:
kg.add_entity(entity_id, name, entity_type)
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)
facts = [
("einstein", "born_in", "germany"),
("newton", "born_in", "england"),
("curie", "born_in", "poland"),
("darwin", "born_in", "england"),
("tesla", "born_in", "serbia"),
("einstein", "developed", "relativity"),
("newton", "developed", "gravity"),
("darwin", "developed", "evolution"),
("curie", "developed", "radioactivity"),
("newton", "influenced", "einstein"),
("darwin", "influenced", "einstein"),
("curie", "influenced", "tesla"),
("einstein", "studied", "gravity"),
("tesla", "studied", "relativity"),
]
for subject, predicate, object in facts:
kg.add_triple(subject, predicate, object)
return kg
# Test the KG
sample_kg = create_sample_kg()
sample_kg.print_stats()
print("โœ… Cell 3: Complete (with get_statistics fix)")
# Now continue with other cells in globals() context
print("\n๐Ÿ”„ Cell 5: Knowledge Graph Embeddings")
print("-" * 70)
try:
code5 = "".join(nb["cells"][5]["source"])
exec(code5, globals())
print("โœ… Cell 5: Complete")
except Exception as e:
print(f"โš ๏ธ Cell 5: {str(e)}")
print("\n๐Ÿ”„ Cell 7: Graph Neural Networks")
print("-" * 70)
try:
code7 = "".join(nb["cells"][7]["source"])
exec(code7, globals())
print("โœ… Cell 7: Complete")
except Exception as e:
print(f"โš ๏ธ Cell 7: {str(e)}")
print("\n๐Ÿ”„ Cell 9: Multi-hop Reasoning")
print("-" * 70)
try:
code9 = "".join(nb["cells"][9]["source"])
exec(code9, globals())
print("โœ… Cell 9: Complete")
except Exception as e:
print(f"โš ๏ธ Cell 9: {str(e)}")
print("\n๐Ÿ”„ Cell 11: Logical Reasoning")
print("-" * 70)
try:
code11 = "".join(nb["cells"][11]["source"])
exec(code11, globals())
print("โœ… Cell 11: Complete")
except Exception as e:
print(f"โš ๏ธ Cell 11: {str(e)}")
print("\n๐Ÿ”„ Cell 13: Comprehensive Analysis")
print("-" * 70)
try:
code13 = "".join(nb["cells"][13]["source"])
exec(code13, globals())
print("โœ… Cell 13: Complete")
except Exception as e:
print(f"โš ๏ธ Cell 13: {str(e)}")
import traceback
traceback.print_exc()
print("\n" + "=" * 70)
print("๐ŸŽ‰ Notebook execution complete!")
print("=" * 70)

Xet Storage Details

Size:
13.7 kB
ยท
Xet hash:
5d60721c3b065848016c5ee158b5a2b6adced60804a32608a63f93ae8e49e411

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