Imagine you want to write PyTorch code but don't know which exact function names, parameters, or class names to use. This project does that FOR YOU β automatically.
It uses a massive knowledge graph (24,485 PyTorch API concepts, 47,958 connections) and a Heterogeneous Graph Neural Network (HGTConv) to find and assemble the most relevant code for any question you ask.
| Feature | Structural-Coder (Ours) | Standalone AI (e.g. llama3.1:8b) |
|---|---|---|
| PyTorch knowledge graph (24K nodes, 47K edges) | β | β |
| Heterogeneous GNN (9 distinct node types) | β HGTConv | β |
| Grounded in real API symbols | β Always | β Often hallucinates |
Fully-qualified import paths (e.g. torch.nn.Module) |
β From GNN output | β Guesses |
| Code validation (C0βC5 checks) | β Automatic | β Never |
| Avg Final Score (benchmark) | π 0.65+ | 0.28 |
We won all 10 out of 10 benchmark queries against llama3.1:8b.
Structural-Coder/
β
βββ README.md β You are here
βββ requirements.txt
β
βββ data/ β Input Data (Knowledge Graph)
β βββ nodes.csv β 24,485 PyTorch API nodes (Id, Label, Name, URL)
β βββ edges.csv β 47,958 connections (Source, Target, Type)
β
βββ outputs/ β GNN Training Outputs (from gnn_encoder_improved.ipynb)
β βββ gnn_embeddings.jsonl β 24,485 node embeddings (256-D) with display_name + node_type
β βββ gnn_embeddings.pt β Same embeddings as PyTorch dict (faster to load)
β βββ best_model.pt β Trained HGTConv model weights
β βββ hetero_metadata.json β Node types + edge types metadata
β βββ hetero_graph.pt β Full HeteroData graph object
β βββ train_graph.pt β Training split of graph
β βββ split_state.pt β Train/val/test split state
β βββ training_summary.json β Metrics, thresholds, training history
β
βββ src/
β βββ graph_rag/ β Heterogeneous GNN retrieval engine
β β βββ gnn_encoder.py β HGTConv model + embedding loaders
β β βββ retriever.py β Topological anchor + neighborhood retrieval
β βββ integration_pipeline/ β Validation + graph loading
β β βββ graph_loader.py β CSV graph reader (Node, Edge, CsvGraph)
β β βββ validator.py β C0-C5 active code checks
β β βββ pipeline.py β Integration orchestrator
β βββ research_pipeline/ β Research evaluation orchestrator
β βββ pipeline.py β Main pipeline: loading, retrieval, LLM, scoring
β
βββ benchmark/
β βββ interactive_comparison.py β Live side-by-side tester
β βββ run_comparison.py β Full batch benchmark
β βββ queries/queries.json β 10 test questions
β βββ outputs/ β Live benchmark results
β
βββ notebooks/
β βββ gnn_encoder_improved.ipynb β Source of the HGTConv architecture + training
β
βββ tests/
This diagram shows exactly how each data file flows through the system:
data/nodes.csv βββββββββββββββ
24,485 nodes β
Fields: Id, Label, Name, ββββ graph_loader.py βββ CsvGraph (in-memory graph)
URL β β
data/edges.csv βββββββββββββββ β
47,958 edges β
Fields: Source, Target, β
Type β
βΌ
outputs/gnn_embeddings.jsonl βββ load_embeddings_from_jsonl()
24,485 embeddings (256-D) β
Fields: element_id, ββββ node_ids + embeddings tensor
node_type, display_name, β
embedding ββββ display_names dict βββ Enriches CsvGraph node names
β (e.g. "SymInt" β "torch.SymInt")
β
βΌ
GraphRAGRetriever
β
βββ Phase 1: Lexical Anchor Discovery (text match)
βββ Phase 2: Topological Neighborhood Expansion (GNN cosine)
βββ Phase 3: Hybrid Re-Ranking (text + GNN + degree + type)
β
βΌ
Retrieved Context (top 5 nodes)
β
βΌ
_build_ollama_prompt() ββ Splits APIs vs Concepts
β
βΌ
Ollama LLM (llama3.1:8b) ββ Advisory prompt
β
βΌ
Active Validator (C0βC5)
β
βΌ
β
Grounded, Validated PyTorch Code
| File | Used at Runtime? | Purpose |
|---|---|---|
gnn_embeddings.jsonl |
β Yes | Pre-computed 256-D vectors for all 24,485 nodes. Contains display_name (fully-qualified Python paths like torch.nn.Module) used to enrich node names for the LLM. |
gnn_embeddings.pt |
β Optional | Same embeddings in PyTorch dict format β faster to load than JSONL. |
best_model.pt |
β No | The trained HGTConv model weights. Used during training to produce the embeddings, not needed at inference. |
hetero_metadata.json |
β No | Node/edge type schema. Reference only. |
hetero_graph.pt |
β No | Full HeteroData graph. Used during GNN training only. |
train_graph.pt |
β No | Training split. Used during GNN training only. |
split_state.pt |
β No | Train/val/test split. Used during GNN training only. |
training_summary.json |
β No | Metrics + history from GNN training (AUC, accuracy per relation). |
pip install -r requirements.txtOLLAMA_MODELS="/path/to/models" ollama serve
ollama pull llama3.1:8bcd Structural-Coder
../.venv/bin/python benchmark/interactive_comparison.py --model llama3.1:8b../.venv/bin/python benchmark/run_comparison.py --models llama3.1:8bThe user's query text (e.g. "flash attention with fallback") is tokenized and keyword-searched against all 24,485 nodes. This produces 1β4 Anchor Nodes β the exact entry points into the graph.
Each Anchor Node's pre-computed 256-D vector (from gnn_embeddings.jsonl) is retrieved. We compute cosine similarity against all 24,485 vectors to find the structurally nearest neighbors β APIs that share edges in the documentation graph, even if they share zero text keywords.
All candidates (lexical matches + topological neighbors + 1-hop graph expansion) are scored with:
- GNN cosine similarity (structural proximity)
- Lexical overlap (text matching, weighted 2x)
- Graph degree (well-connected nodes preferred)
- Node type bonus:
API_Class/Function/Methodboosted +1.0,API_Endpointdemoted -0.5
Top 5 nodes are split into:
- π» Valid PyTorch APIs (advisory β use only if relevant)
- π Conceptual Context (READ-ONLY, do not import)
The prompt instructs the LLM: "If the APIs seem irrelevant, IGNORE THEM and rely on your own knowledge."
Generated code is validated with 6 checks: syntax, import safety, API correctness, type checking, runtime execution, and compilation.
The src/graph_rag/gnn_encoder.py module implements these components ported directly from notebooks/gnn_encoder_improved.ipynb:
API_Class, API_Function, API_Method, API_Parameter, API_Endpoint, CodeSnippet, Concept, DeprecatedAPI, PyTorchConcept
- Input projection per node type (
nn.Linearβnn.LayerNorm) - 3-layer HGTConv (Heterogeneous Graph Transformer) with 4 attention heads
- Residual pass-through for isolated node types (prevents dead gradients)
- Jumping Knowledge (JK) connections: all layer outputs concatenated β final linear
- Factored bilinear link scorer:
score = (z_src @ U) Β· (z_dst @ V)(rank=32) - Much fewer parameters than full bilinear β prevents overfitting on rare relations
- Supervision relations:
IMPLEMENTS,CONTAINS,HAS_PARAM,CALLS,RELATED_TO,REPLACES - Message-passing relations: All above +
EXPLAINS,REFERENCES - Hard Negative Sampling: Degree-distribution weighted, true-positive filtered
- Gradient Clipping:
max_norm=2.0for stable HGT training
| Relation Type | ROC AUC | Accuracy |
|---|---|---|
| API_EndpointβAPI_Function | 1.000 | 1.000 |
| API_ClassβAPI_Method | 0.863 | 0.833 |
| API_EndpointβAPI_Class | 0.811 | 0.857 |
| API_EndpointβCodeSnippet | 0.672 | 0.650 |
| Overall | 0.668 | 0.701 |
The GNN Encoder and the Integration Pipeline are strictly decoupled. You can retrain, swap, or improve the GNN without touching a single line of the pipeline code. The pipeline simply reads the cached outputs/gnn_embeddings.jsonl (or .pt) file and enriches node names from the display_name field.
- Amitesh Sinha β Benchmarking, Pipeline Integration, Evaluation Engine
- Rahul Anand β GNN Architecture
- Mohit β Validation Pipeline, Data Engineering