Lab3 / README.md
NamPhoenix's picture
Update README.md
159e6f1 verified
|
Raw
History Blame Contribute Delete
6.59 kB

A newer version of the Gradio SDK is available: 6.22.0

Upgrade
metadata
title: Lab3
emoji: 📊
colorFrom: indigo
colorTo: yellow
sdk: gradio
sdk_version: 6.13.0
app_file: app.py
pinned: false

GNN-based Medical RAG System

Healthcare Question Answering using Graph Neural Network-enhanced Retrieval Augmented Generation.


Table of Contents

  1. Algorithm Overview
  2. Pipeline Steps
  3. System Architecture
  4. Repository Structure
  5. References

Algorithm Overview

This system implements a GNN-based RAG pipeline for medical QA, combining semantic retrieval with structural graph reasoning. Each query is answered by retrieving relevant documents through a hybrid scoring mechanism that blends raw embedding similarity with GNN-enriched structural embeddings.


Pipeline Steps

1. Dataset

Property Value
Source keivalya/MedQuad-MedicalQnADataset (HuggingFace)
Size 250 answer documents
Truncation 800 characters per document

MedQuad contains medical question-answer pairs from trusted sources including NIH, CDC, and NLM. Each document is converted to a LlamaIndex Document object, with the answer as the main text and the original question stored as metadata.


2. LlamaIndex Property Graph Construction

We use LlamaIndex PropertyGraphIndex with ImplicitPathExtractor to parse documents into a structured graph index:

index = PropertyGraphIndex.from_documents(
    documents,
    embed_model=embed_model,
    kg_extractors=[ImplicitPathExtractor()],
)

ImplicitPathExtractor requires no LLM — it creates implicit node relationships based on document structure. Text nodes are extracted from index.docstore.docs for downstream graph construction.


3. PyTorch Geometric Graph Construction

LlamaIndex text nodes are converted to a PyTorch Geometric Data object:

  • Nodes — Each text chunk is one graph node.
    Node features x ∈ ℝ^(N×384) are the BGE-Small embeddings.

  • Edges — Two nodes are connected if their cosine similarity exceeds the threshold:

edge(i,j) exists    cosine_sim(embedi, embedj)>0.70\text{edge}(i, j) \text{ exists} \iff \text{cosine\_sim}(\text{embed}_i,\ \text{embed}_j) > 0.70

Self-loops are added for GCN stability: $\tilde{A} = A + I$


4. Graph Neural Network (GCN)

A 2-layer Graph Convolutional Network following Kipf & Welling (2017) enriches node embeddings with structural context.

Architecture:

x [N, 384] → GCNConv(384→256) → ReLU → Dropout(0.3)
           → GCNConv(256→384) → structural_embs [N, 384]

GCN Update Rule:

H(l+1)=σ ⁣(D~1/2A~D~1/2H(l)W(l))H^{(l+1)} = \sigma\!\left(\tilde{D}^{-1/2}\,\tilde{A}\,\tilde{D}^{-1/2}\,H^{(l)}\,W^{(l)}\right)

where $\tilde{A} = A + I$ (adjacency + self-loops) and $\tilde{D}$ is the corresponding degree matrix.

Key insight: The GNN acts as a structural feature transformer, not a classifier. No training is required — the graph topology itself provides the structural signal. After 2 message-passing layers, each node's embedding captures information from its 2-hop neighborhood.


5. Hybrid Retrieval (GNNHybridRetriever)

We subclass LlamaIndex's BaseRetriever to implement dual-score hybrid retrieval.

Score Formula Description
Semantic $\text{sem}(q, i) = \cos(\text{embed}(q),\ \text{raw_embed}[i])$ Pure text similarity
Structural $\text{struct}(q, i) = \cos(\text{embed}(q),\ \text{gnn_embed}[i])$ Neighborhood-aware similarity
Hybrid $\text{score}(q, i) = \alpha \cdot \text{sem}(q, i) + (1-\alpha) \cdot \text{struct}(q, i)$ Combined score

Default $\alpha = 0.6$ (60% semantic, 40% structural).

Why hybrid is better:

Method Weakness
Semantic only Misses documents with different terminology
Structural only May retrieve irrelevant neighbors of relevant nodes
Hybrid Robust to vocabulary mismatch while leveraging graph topology ✓

The top-K nodes by hybrid score are returned as context for generation.


6. Generation (Qwen3.5-4B-Neo GGUF)

We use llama-cpp-python to run Qwen3.5-4B-Neo-GGUF (Q4_K_M quantization) on CPU for efficient inference on free-tier hardware:

llm = Llama(
    model_path=model_path,
    n_ctx=2048,
    n_threads=4,
    n_gpu_layers=0,   # CPU only
)

The prompt template explicitly constrains the model to answer only from retrieved context, reducing hallucination risk.


System Architecture

MedQuad Dataset (250 docs)
│
▼
LlamaIndex PropertyGraphIndex
(ImplicitPathExtractor)
│
▼
Text Nodes → BGE-Small Embeddings [N × 384]
│
├─── Cosine Similarity → PyG Graph ──► 2-layer GCN
│    (edges where sim > 0.70)               │
│                                Structural Embeddings [N × 384]
│                                           │
└────────────── GNNHybridRetriever ─────────┘
       α × semantic + (1-α) × structural
                     │
               Top-K Context Nodes
                     │
          Qwen3.5-4B-Neo (GGUF, CPU)
                     │
                   Answer

Repository Structure

├── app.py                  # Main Gradio application
├── requirements.txt        # Python dependencies
├── README.md               # This file (also serves as the lab report)
└── utils/
    ├── __init__.py
    ├── gnn_model.py        # MedicalGNN (2-layer GCN) + get_structural_embeddings()
    ├── graph_builder.py    # Dataset loading, PropertyGraphIndex, PyG conversion
    └── retriever.py        # GNNHybridRetriever (BaseRetriever subclass)

References

  1. Lewis, P. et al. (2020). Retrieval Augmented Generation for Knowledge-Intensive NLP Tasks. NeurIPS 2020, 33, 9459–9474.

  2. Edge, D. et al. (2024). From Local to Global: A Graph RAG Approach to Query-Focused Summarization. arXiv:2404.16130.

  3. Kipf, T. N. & Welling, M. (2017). Semi-Supervised Classification with Graph Convolutional Networks. ICLR 2017.