jubayer009's picture
Upload 7 files
762eeee verified
Raw
History Blame Contribute Delete
15.2 kB
"""
load_utils.py
==============
Artifact-loading helpers for Notebook 2 (the interactive CHATBOT notebook).
Every loader here fails LOUDLY and CLEARLY rather than silently returning
empty/default data -- a missing knowledge graph should never be mistaken for
"the graph legitimately has zero nodes". Use `GraphRAGArtifactError` to catch
all of this module's own raised errors in one place if you want a single
top-level try/except in the notebook.
Typical usage (see Notebook 2's early cells for the full call):
import config
import load_utils as lu
root = config.resolve_artifacts_dir()
lu.check_manifest(root) # fails fast on version mismatch
local_corpus = lu.load_local_corpus(root)
embeddings, chunk_ids = lu.load_embeddings(root)
bm25 = lu.load_bm25(root, local_corpus) # unpickles, or rebuilds from corpus
graph, kg_meta = lu.load_knowledge_graph(root)
kg_name_lookup = lu.load_kg_name_lookup(root)
default_disease_pair = lu.load_default_disease_pair(root)
"""
import os
import json
import pickle
import networkx as nx
import numpy as np
from typing import Dict, List, Tuple, Optional, Any
import config
class GraphRAGArtifactError(Exception):
"""Raised for any missing/corrupted/incompatible saved artifact, with a
message that tells the person exactly what went wrong and how to fix it."""
# ==============================================================================
# Low-level, defensive readers
# ==============================================================================
def _require_file(path: str, what: str) -> None:
if not os.path.exists(path):
raise GraphRAGArtifactError(
f"Missing artifact: {what} was expected at '{path}' but the file doesn't exist.\n"
"Fix: run Notebook 1 (the offline build notebook) to completion, publish its "
f"'{config.SAVE_DIR_NAME}' output folder as a Kaggle Dataset, and attach it to this "
"notebook via 'Add Input' -- or double-check the artifacts directory you resolved."
)
def _require_dir(path: str, what: str) -> None:
if not os.path.isdir(path):
raise GraphRAGArtifactError(
f"Missing artifact folder: {what} was expected at '{path}' but the directory doesn't "
"exist. This usually means the attached dataset is from an incomplete or failed run of "
"Notebook 1, or the wrong dataset/path was attached."
)
def load_json(path: str, what: str) -> Any:
_require_file(path, what)
try:
with open(path, "r", encoding="utf-8") as f:
return json.load(f)
except json.JSONDecodeError as exc:
raise GraphRAGArtifactError(
f"Corrupted artifact: {what} at '{path}' is not valid JSON ({exc}). "
"This file may have been truncated during upload/download -- try re-saving/re-attaching "
"Notebook 1's output dataset."
) from exc
def load_jsonl(path: str, what: str) -> List[Dict]:
_require_file(path, what)
records = []
try:
with open(path, "r", encoding="utf-8") as f:
for line_num, line in enumerate(f, start=1):
line = line.strip()
if not line:
continue
try:
records.append(json.loads(line))
except json.JSONDecodeError as exc:
raise GraphRAGArtifactError(
f"Corrupted artifact: {what} at '{path}' has invalid JSON on line "
f"{line_num} ({exc}). Try re-saving/re-attaching Notebook 1's output dataset."
) from exc
except OSError as exc:
raise GraphRAGArtifactError(f"Could not read {what} at '{path}': {exc}") from exc
return records
# ==============================================================================
# Manifest / schema-version check -- call this FIRST.
# ==============================================================================
def check_manifest(root: str) -> Dict:
"""
Load and validate manifest.json. Raises GraphRAGArtifactError with a clear
message if:
* the manifest is missing entirely (Notebook 1 was never run, or the
wrong folder was attached),
* the manifest's schema_version doesn't match this codebase's
`config.ARTIFACT_SCHEMA_VERSION` (a version mismatch -- the artifact
layout may have changed since this save was produced).
"""
path = config.path_for(root, "manifest", "manifest")
manifest = load_json(path, "the artifact manifest")
saved_version = manifest.get("schema_version")
if saved_version != config.ARTIFACT_SCHEMA_VERSION:
raise GraphRAGArtifactError(
f"Version mismatch: these artifacts were saved with schema_version={saved_version}, "
f"but this notebook expects schema_version={config.ARTIFACT_SCHEMA_VERSION}. "
"Re-run Notebook 1 with the current save_utils.py/config.py to regenerate compatible "
"artifacts, or use a matching version of load_utils.py/config.py to read the old ones."
)
print(f"Manifest OK: schema_version={saved_version}, saved_at={manifest.get('saved_at_utc')}, "
f"{manifest.get('artifact_count', '?')} artifact file(s), "
f"{manifest.get('warning_count', 0)} warning(s) at save time.")
return manifest
# ==============================================================================
# Corpus / embeddings
# ==============================================================================
def load_local_corpus(root: str) -> List[Dict]:
path = config.path_for(root, "corpus", "local_corpus")
corpus = load_jsonl(path, "the local corpus")
print(f"Loaded local corpus: {len(corpus)} chunk(s) from {path}")
return corpus
def load_chunk_lookup(root: str, local_corpus: Optional[List[Dict]] = None) -> Dict[str, Dict]:
path = config.path_for(root, "corpus", "chunk_lookup")
if os.path.exists(path):
return load_json(path, "the chunk lookup")
if local_corpus is not None:
return {item["chunk_id"]: item for item in local_corpus}
raise GraphRAGArtifactError(f"Missing artifact: chunk lookup not found at '{path}' and no "
"local_corpus was supplied to rebuild it from.")
def load_embeddings(root: str) -> Tuple[np.ndarray, List[str]]:
emb_path = config.path_for(root, "embeddings", "corpus_embeddings")
ids_path = config.path_for(root, "embeddings", "corpus_chunk_ids")
_require_file(emb_path, "the corpus embeddings")
try:
embeddings = np.load(emb_path)
except Exception as exc:
raise GraphRAGArtifactError(
f"Corrupted artifact: could not load embeddings numpy file '{emb_path}' ({exc})."
) from exc
chunk_ids = load_json(ids_path, "the embeddings' chunk-ID order")
if embeddings.shape[0] != len(chunk_ids):
raise GraphRAGArtifactError(
f"Corrupted/mismatched artifacts: embeddings array has {embeddings.shape[0]} row(s) but "
f"the chunk-ID order file lists {len(chunk_ids)} id(s). These two files must come from "
"the SAME Notebook 1 run -- re-save/re-attach both together."
)
print(f"Loaded embeddings: shape={embeddings.shape} for {len(chunk_ids)} chunk(s)")
return embeddings, chunk_ids
def load_embedding_model_info(root: str) -> Dict:
path = config.path_for(root, "embeddings", "embedding_model_info")
return load_json(path, "embedding model info") if os.path.exists(path) else {}
def rebuild_embedding_model(root: str, embedding_model_cls, corpus_texts: Optional[List[str]] = None):
"""
Reconstruct an EmbeddingModel instance matching the ORIGINAL run's backend
where possible. Downloading/loading the transformer model weights still
happens here (unavoidable -- that's a one-time-per-session model load, not
the expensive corpus-encoding step this whole save/load split exists to
skip). If the original run used the TF-IDF+SVD fallback, its fitted
vectorizer/SVD are restored from the pickled fallback artifact instead of
being re-fit, so the query-time vector space matches the saved embeddings
exactly.
"""
info = load_embedding_model_info(root)
model = embedding_model_cls(model_name=info.get("model_name") or None)
if model.backend == "tfidf_svd":
fallback_path = os.path.join(root, config.SUBDIRS["embeddings"], "tfidf_svd_fallback.pkl")
if os.path.exists(fallback_path):
with open(fallback_path, "rb") as f:
fitted = pickle.load(f)
model._tfidf, model._svd = fitted["tfidf"], fitted["svd"]
print("Restored the saved TF-IDF+SVD fallback vectorizer (query embeddings will land in "
"the same vector space as the saved corpus embeddings).")
elif corpus_texts:
model.fit_fallback(corpus_texts)
print("WARNING: no saved TF-IDF/SVD fallback found -- refit it from the loaded corpus. "
"This should reproduce the original space (fixed random_state) given the same corpus.")
else:
raise GraphRAGArtifactError(
"The embedding model's fallback backend (TF-IDF+SVD) needs either the saved "
"'tfidf_svd_fallback.pkl' artifact or the corpus text to refit -- neither was provided."
)
return model
# ==============================================================================
# BM25
# ==============================================================================
def load_bm25(root: str, bm25_wrapper_cls, local_corpus: Optional[List[Dict]] = None):
"""
Load the pickled BM25 retriever if present and it unpickles cleanly.
Otherwise (no pickle was saved, or it fails to load -- e.g. a rank_bm25
version mismatch), REBUILD it from the saved corpus. Rebuilding is fast
(BM25/TF-IDF index construction over a few thousand chunks is seconds,
not the expensive step -- that was the original PDF parsing/chunking).
"""
pkl_path = config.path_for(root, "bm25", "bm25_pickle")
if os.path.exists(pkl_path):
try:
with open(pkl_path, "rb") as f:
bm25 = pickle.load(f)
print(f"Loaded pickled BM25 retriever from {pkl_path} (backend={getattr(bm25, 'backend', '?')})")
return bm25
except Exception as exc:
print(f"[WARNING] Could not unpickle BM25 retriever ({exc}); rebuilding from the saved corpus instead.")
if local_corpus is None:
raise GraphRAGArtifactError(
"No usable BM25 pickle found, and no local_corpus was supplied to rebuild BM25 from. "
"Pass the corpus loaded via load_local_corpus()."
)
bm25 = bm25_wrapper_cls(local_corpus)
print(f"Rebuilt BM25 retriever from the saved corpus (backend={bm25.backend}, "
f"{len(local_corpus)} document(s)) -- fast, no PDF re-processing involved.")
return bm25
# ==============================================================================
# Knowledge graph
# ==============================================================================
def load_knowledge_graph(root: str, clinical_kg_cls) -> Tuple[Any, Dict]:
"""
Reconstruct a ClinicalKnowledgeGraph instance from the saved node-link
JSON, WITHOUT re-running the (heavier, merge-logic-driven) KG-file-loading
pipeline from Chapter 4.3. Node/edge attributes, graph version, and
changelog are restored exactly as they were at save time.
Returns:
Tuple[ClinicalKnowledgeGraph, Dict]: (kg, kg_meta)
"""
graph_path = config.path_for(root, "graph", "kg_node_link_json")
meta_path = config.path_for(root, "graph", "kg_meta")
_require_file(graph_path, "the knowledge graph")
try:
with open(graph_path, "r", encoding="utf-8") as f:
data = json.load(f)
graph = nx.node_link_graph(data, directed=True, multigraph=True)
except json.JSONDecodeError as exc:
raise GraphRAGArtifactError(
f"Corrupted artifact: knowledge graph JSON at '{graph_path}' is invalid ({exc})."
) from exc
except Exception as exc:
raise GraphRAGArtifactError(
f"Graph loading failure: could not reconstruct the networkx graph from '{graph_path}' "
f"({exc}). This can happen if the file was saved with an incompatible networkx version."
) from exc
kg_meta = load_json(meta_path, "knowledge graph metadata") if os.path.exists(meta_path) else {}
kg = clinical_kg_cls()
kg.graph = graph
kg.version = kg_meta.get("version", 1)
kg.changelog = kg_meta.get("changelog", [])
if kg.graph.number_of_nodes() == 0:
print("[WARNING] Loaded knowledge graph has 0 nodes -- graph-based reasoning will find no paths. "
"Check that Notebook 1's KG-loading chapter actually populated CKG before saving.")
print(f"Loaded knowledge graph: {kg.graph.number_of_nodes()} nodes, "
f"{kg.graph.number_of_edges()} edges (graph v{kg.version})")
return kg, kg_meta
def load_kg_name_lookup(root: str) -> Dict[str, str]:
path = config.path_for(root, "graph", "kg_name_lookup")
return load_json(path, "the KG name/synonym lookup")
def load_default_disease_pair(root: str) -> Tuple[str, str]:
path = config.path_for(root, "graph", "default_disease_pair")
pair = load_json(path, "the default disease pair")
return tuple(pair) if pair else ("", "")
# ==============================================================================
# Ontology / entities / relations / config (all lightweight, informational)
# ==============================================================================
def load_ontology(root: str) -> Tuple[Dict[str, set], Dict[str, Dict]]:
sets_path = config.path_for(root, "ontology", "entity_sets")
map_path = config.path_for(root, "ontology", "ontology_map")
raw_sets = load_json(sets_path, "the curated entity ontology")
all_entity_sets = {label: set(terms) for label, terms in raw_sets.items()}
ontology_map = load_json(map_path, "the MeSH/UMLS/SNOMED ontology map") if os.path.exists(map_path) else {}
return all_entity_sets, ontology_map
def load_entities_relations(root: str) -> Tuple[List[Dict], List[Dict]]:
entities_path = config.path_for(root, "metadata", "entities")
relations_path = config.path_for(root, "metadata", "relations")
entities = load_jsonl(entities_path, "extracted entities") if os.path.exists(entities_path) else []
relations = load_jsonl(relations_path, "extracted relations") if os.path.exists(relations_path) else []
return entities, relations
def load_project_config(root: str) -> Dict:
path = config.path_for(root, "config", "project_config")
return load_json(path, "the project configuration") if os.path.exists(path) else {}
def load_dependency_report(root: str) -> Dict:
path = config.path_for(root, "config", "dependency_report")
return load_json(path, "the dependency report") if os.path.exists(path) else {}