multimodal_rag_demo.py

script

← Back to skill

Content hash: 1bf700156b3ca51c417d0a1ea5f20a86a24966a54047008d1635dc5a95cb6ec6
#!/usr/bin/env python3
"""Multimodal RAG pipeline skeleton: image/pdf ingestion, embedding, retrieval.

Choose strategy: text-summarize (VLM caption) or native multimodal embeddings.
This demo uses simulated embeddings for zero-dependency execution.
"""

from __future__ import annotations

import hashlib
import uuid
from dataclasses import dataclass, field
from enum import Enum
from typing import Optional

import numpy as np


# ── Chunk types ────────────────────────────────────────────────────────────


class MediaType(str, Enum):
    TEXT = "text"
    IMAGE = "image"
    PDF_PAGE = "pdf_page"


@dataclass
class Chunk:
    id: str
    media_type: MediaType
    content: str  # text or VLM-generated caption
    source_path: str
    embedding: Optional[list[float]] = None


# ── Simulated embedder (swap for real model in prod) ──────────────────────


def simulate_embed(text: str, dim: int = 128) -> list[float]:
    """Deterministic hash-based embedding. Replace with:
    - CLIP (image+text)
    - Jina v4, Cohere Embed 4, voyage-multimodal
    - ColPali patch-level embeddings (for PDF pages)
    """
    rng = np.random.default_rng(int(hashlib.md5(text.encode()).hexdigest(), 16) % (2**31))
    vec = rng.normal(size=dim).astype(float)
    return (vec / np.linalg.norm(vec)).tolist()


# ── Vector index ──────────────────────────────────────────────────────────


class SimpleVectorIndex:
    """In-memory cosine similarity index. Swap for pgvector/Qdrant/Milvus."""

    def __init__(self):
        self.chunks: list[Chunk] = []

    def add(self, chunk: Chunk) -> None:
        self.chunks.append(chunk)

    def search(self, query_embedding: list[float], k: int = 5) -> list[tuple[Chunk, float]]:
        scored = []
        q = np.array(query_embedding)
        for chunk in self.chunks:
            if chunk.embedding is not None:
                c = np.array(chunk.embedding)
                sim = float(np.dot(q, c))
                scored.append((chunk, sim))
        scored.sort(key=lambda x: x[1], reverse=True)
        return scored[:k]


# ── Pipeline ──────────────────────────────────────────────────────────────


class MultimodalRAGPipeline:
    """Strategy-switchable multimodal RAG."""

    def __init__(self, strategy: str = "text-summarize"):
        if strategy not in ("text-summarize", "native-multimodal"):
            raise ValueError("Strategy must be 'text-summarize' or 'native-multimodal'")
        self.strategy = strategy
        self.index = SimpleVectorIndex()

    def ingest(self, media_type: MediaType, content: str, source_path: str) -> None:
        """Ingest one media item."""
        chunk = Chunk(
            id=uuid.uuid4().hex[:8],
            media_type=media_type,
            content=content,
            source_path=source_path,
        )

        if self.strategy == "text-summarize":
            # Embed the VLM-generated text caption
            chunk.embedding = simulate_embed(chunk.content)
        else:
            # Native multimodal: embed the content + media_type context
            chunk.embedding = simulate_embed(f"{media_type.value}: {chunk.content}")

        self.index.add(chunk)

    def query(self, text: str, k: int = 3) -> list[dict]:
        """Search: embed query with SAME model, retrieve top-k."""
        q_emb = simulate_embed(text)
        results = self.index.search(q_emb, k=k)
        return [{"source": c.source_path, "type": c.media_type.value,
                 "content": c.content[:100]} for c, _ in results]


# ── Demo ───────────────────────────────────────────────────────────────────


def main() -> None:
    pipeline = MultimodalRAGPipeline(strategy="text-summarize")

    # Simulate ingesting document chunks
    pipeline.ingest(MediaType.TEXT, "The REST API uses JWT for authentication on all endpoints.", "api-docs.md")
    pipeline.ingest(MediaType.IMAGE, "A bar chart showing Q4 revenue growth: $2.1M -> $3.4M -> $5.2M", "reports/q4-chart.png")
    pipeline.ingest(MediaType.PDF_PAGE, "Experimental results show 94.2% accuracy on the test set.", "paper.pdf#p3")
    pipeline.ingest(MediaType.TEXT, "The Skill Vault supports tools, resources, and prompts via MCP.", "skill-vault-readme.md")

    queries = [
        "What authentication method does the API use?",
        "Show me Q4 revenue figures",
        "What was the accuracy of the experiment?",
    ]

    print("=== Multimodal RAG (text-summarize strategy) ===\n")
    for q in queries:
        print(f"Q: {q}")
        for r in pipeline.query(q, k=2):
            print(f"  [{r['type']:8s}] {r['source']}")
            print(f"           {r['content']}")
        print()

    print("Rules for production:")
    print("  1. Query embedder MUST match index embedder")
    print("  2. Text-summarize: cheap, but loses visual fidelity")
    print("  3. Native multimodal: preserves details, heavier/costlier")
    print("  4. ColPali: best for documents where layout IS the meaning")


if __name__ == "__main__":
    main()