CoolFace
Apppublic

prazy1208/text2sql

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
build_vector_store.py323 linesDownload Raw Back to root
1"""2Metadata-to-FAISS pipeline for PostgreSQL schemas (healthcare, retail, finance).3Extracts table/column metadata and comments, builds embeddings with sentence-transformers,4and stores them in per-schema FAISS indexes and metadata JSON.5 6Outputs per schema:7- Table-level: faiss_indexes/{schema}.index + metadata_store/{schema}_metadata.json8- Column-level: faiss_indexes/{schema}_columns.index + metadata_store/{schema}_columns_metadata.json9  (flat list order matches FAISS row indices 0..n-1).10 11Run: python build_vector_store.py12"""13 14import json15import logging16import os17from pathlib import Path18 19from dotenv import load_dotenv20from sqlalchemy import create_engine, text21 22# ---------------------------------------------------------------------------23# Logging24# ---------------------------------------------------------------------------25logging.basicConfig(26    level=logging.INFO,27    format="%(asctime)s [%(levelname)s] %(message)s",28    datefmt="%Y-%m-%d %H:%M:%S",29)30logger = logging.getLogger(__name__)31 32load_dotenv()33 34# Output directories35FAISS_INDEX_DIR = Path("faiss_indexes")36METADATA_STORE_DIR = Path("metadata_store")37SCHEMAS = ["healthcare_schema", "retail_schema", "finance_schema"]38 39 40def get_engine():41    """Build SQLAlchemy engine from .env (DATABASE_URL or DB_* variables)."""42    database_url = os.getenv("DATABASE_URL")43    if database_url:44        return create_engine(database_url)45    host = os.getenv("DB_HOST", "localhost")46    port = os.getenv("DB_PORT", "5432")47    user = os.getenv("DB_USER", "postgres")48    password = os.getenv("DB_PASSWORD", "")49    dbname = os.getenv("DB_NAME", "postgres")50    url = f"postgresql://{user}:{password}@{host}:{port}/{dbname}"51    return create_engine(url)52 53 54def extract_metadata(engine, schema_name: str) -> list[dict]:55    """56    Extract schema/table/column metadata and COMMENT ON TABLE / COMMENT ON COLUMN.57    Column types come from pg_catalog.format_type (Postgres-native, includes length/precision).58    Returns a list of dicts, one per table, with keys: schema_name, table_name,59    table_description, columns (list of {name, description, data_type}).60    """61    # Tables with comments (pg_catalog)62    tables_sql = text("""63        SELECT64            c.relname AS table_name,65            COALESCE(pg_catalog.obj_description(c.oid, 'pg_class'), '') AS table_comment66        FROM pg_catalog.pg_class c67        JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace68        WHERE n.nspname = :schema_name69          AND c.relkind = 'r'70          AND c.relname <> 'table_relationships'71          AND c.relname NOT LIKE '%\\_business\\_rules' ESCAPE '\\'72        ORDER BY c.relname73    """)74    # Columns: comments + data type (format_type matches psql / CREATE TABLE spelling)75    columns_sql = text("""76        SELECT77            c.relname AS table_name,78            a.attname AS column_name,79            pg_catalog.format_type(a.atttypid, a.atttypmod) AS data_type,80            COALESCE(pg_catalog.col_description(c.oid, a.attnum), '') AS column_comment81        FROM pg_catalog.pg_attribute a82        JOIN pg_catalog.pg_class c ON c.oid = a.attrelid83        JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace84        WHERE n.nspname = :schema_name85          AND c.relkind = 'r'86          AND c.relname <> 'table_relationships'87          AND c.relname NOT LIKE '%\\_business\\_rules' ESCAPE '\\'88          AND a.attnum > 089          AND NOT a.attisdropped90        ORDER BY c.relname, a.attnum91    """)92 93    with engine.connect() as conn:94        tables = conn.execute(tables_sql, {"schema_name": schema_name}).fetchall()95        columns = conn.execute(columns_sql, {"schema_name": schema_name}).fetchall()96 97    # Group columns by table98    cols_by_table: dict[str, list[dict]] = {}99    for row in columns:100        tbl = row.table_name101        if tbl not in cols_by_table:102            cols_by_table[tbl] = []103        dtype = (row.data_type or "").strip() if row.data_type is not None else ""104        cols_by_table[tbl].append({105            "name": row.column_name,106            "description": (row.column_comment or "").strip(),107            "data_type": dtype,108        })109 110    # Build one record per table111    result = []112    for row in tables:113        table_name = row.table_name114        table_comment = (row.table_comment or "").strip()115        result.append({116            "schema_name": schema_name,117            "table_name": table_name,118            "table_description": table_comment,119            "columns": cols_by_table.get(table_name, []),120        })121 122    logger.info("Extracted metadata for %d tables in %s", len(result), schema_name)123    return result124 125 126def nested_metadata_to_flat_columns(metadata_list: list[dict]) -> list[dict]:127    """128    Flatten nested table metadata to one dict per column for column-level FAISS.129 130    Stable order (documented for rebuild alignment with vectors):131    tables sorted by table_name; within each table, columns sorted by name.132    Each row: schema_name, table_name, column_name, data_type, description,133    table_description, fqn (schema.table.column).134    """135    flat: list[dict] = []136    for m in sorted(metadata_list, key=lambda x: x["table_name"]):137        schema = m["schema_name"]138        table = m["table_name"]139        table_desc = (m.get("table_description") or "").strip()140        for col in sorted(m.get("columns") or [], key=lambda c: c["name"]):141            name = col["name"]142            desc = (col.get("description") or "").strip()143            dtype = (col.get("data_type") or "").strip()144            fqn = f"{schema}.{table}.{name}"145            flat.append({146                "schema_name": schema,147                "table_name": table,148                "column_name": name,149                "data_type": dtype,150                "description": desc,151                "table_description": table_desc,152                "fqn": fqn,153            })154    return flat155 156 157def column_flat_rows_to_texts(flat_rows: list[dict]) -> list[str]:158    """One embedding string per column (metadata only; no row data)."""159    texts = []160    for r in flat_rows:161        tdesc = r["table_description"] or "(no table description)"162        cdesc = r["description"] or "(no description)"163        dtype = r["data_type"] or "(unknown type)"164        text = "\n".join(165            [166                f"Schema: {r['schema_name']}",167                f"Table: {r['table_name']}",168                f"Table description: {tdesc}",169                f"Column: {r['column_name']}",170                f"Type: {dtype}",171                f"Column description: {cdesc}",172            ]173        )174        texts.append(text)175    return texts176 177 178def metadata_to_texts(metadata_list: list[dict]) -> list[str]:179    """180    Turn per-table metadata into structured text chunks suitable for embedding.181    One string per table.182    """183    texts = []184    for m in metadata_list:185        lines = [186            f"Schema: {m['schema_name']}",187            f"Table: {m['table_name']}",188            f"Description: {m['table_description'] or '(no description)'}",189            "Columns:",190        ]191        for col in m["columns"]:192            desc = col["description"] or "(no description)"193            dtype = (col.get("data_type") or "").strip() or "(unknown type)"194            lines.append(f"  - {col['name']} ({dtype}): {desc}")195        texts.append("\n".join(lines))196    return texts197 198 199def get_embedding_model(model_name: str = "all-MiniLM-L6-v2"):200    """Load the sentence-transformers model once. Reuse the returned model for all schemas."""201    from sentence_transformers import SentenceTransformer202 203    logger.info("Loading embedding model: %s", model_name)204    return SentenceTransformer(model_name)205 206 207def build_embeddings(texts: list[str], model):208    """209    Generate embeddings for a list of texts using a pre-loaded sentence-transformers model.210    Returns a numpy array of shape (len(texts), embedding_dim).211    """212    logger.info("Encoding %d text(s)", len(texts))213    embeddings = model.encode(texts, show_progress_bar=len(texts) > 10)214    return embeddings215 216 217def build_faiss_index(schema_name: str, embeddings, metadata_list: list[dict]):218    """219    Create a FAISS IndexFlatL2 for the given embeddings, save the index to220    faiss_indexes/{schema}.index and metadata to metadata_store/{schema}_metadata.json.221    """222    import faiss223    import numpy as np224 225    FAISS_INDEX_DIR.mkdir(parents=True, exist_ok=True)226    METADATA_STORE_DIR.mkdir(parents=True, exist_ok=True)227 228    embeddings = np.asarray(embeddings, dtype=np.float32)229    if embeddings.ndim == 1:230        embeddings = embeddings.reshape(1, -1)231    d = embeddings.shape[1]232    index = faiss.IndexFlatL2(d)233    index.add(embeddings)234    n = index.ntotal235    logger.info("Built FAISS index for %s: %d vectors, dim=%d", schema_name, n, d)236 237    index_path = FAISS_INDEX_DIR / f"{schema_name}.index"238    faiss.write_index(index, str(index_path))239    logger.info("Saved FAISS index to %s", index_path)240 241    metadata_path = METADATA_STORE_DIR / f"{schema_name}_metadata.json"242    with open(metadata_path, "w", encoding="utf-8") as f:243        json.dump(metadata_list, f, indent=2, ensure_ascii=False)244    logger.info("Saved metadata to %s", metadata_path)245 246 247def build_column_faiss_index(schema_name: str, embeddings, flat_column_metadata: list[dict]):248    """249    Column-level FAISS: vector i corresponds to flat_column_metadata[i].250    Writes faiss_indexes/{schema}_columns.index and251    metadata_store/{schema}_columns_metadata.json.252    """253    import faiss254    import numpy as np255 256    if not flat_column_metadata:257        logger.warning("No columns to index for %s; skipping column FAISS.", schema_name)258        return259 260    FAISS_INDEX_DIR.mkdir(parents=True, exist_ok=True)261    METADATA_STORE_DIR.mkdir(parents=True, exist_ok=True)262 263    embeddings = np.asarray(embeddings, dtype=np.float32)264    if embeddings.ndim == 1:265        embeddings = embeddings.reshape(1, -1)266    n_vec = embeddings.shape[0]267    if n_vec != len(flat_column_metadata):268        raise ValueError(269            f"Column embedding count ({n_vec}) != flat metadata rows ({len(flat_column_metadata)}) "270            f"for {schema_name}"271        )272 273    d = embeddings.shape[1]274    index = faiss.IndexFlatL2(d)275    index.add(embeddings)276    index_path = FAISS_INDEX_DIR / f"{schema_name}_columns.index"277    faiss.write_index(index, str(index_path))278    logger.info(279        "Built column FAISS index for %s: %d vectors, dim=%d -> %s",280        schema_name,281        index.ntotal,282        d,283        index_path,284    )285 286    metadata_path = METADATA_STORE_DIR / f"{schema_name}_columns_metadata.json"287    with open(metadata_path, "w", encoding="utf-8") as f:288        json.dump(flat_column_metadata, f, indent=2, ensure_ascii=False)289    logger.info("Saved column flat metadata to %s", metadata_path)290 291 292def run_pipeline_for_schema(engine, schema_name: str, model):293    """Extract metadata, build embeddings, and write FAISS index + metadata for one schema."""294    logger.info("Processing schema: %s", schema_name)295    metadata_list = extract_metadata(engine, schema_name)296    if not metadata_list:297        logger.warning("No tables found in %s, skipping.", schema_name)298        return299    texts = metadata_to_texts(metadata_list)300    embeddings = build_embeddings(texts, model)301    build_faiss_index(schema_name, embeddings, metadata_list)302 303    flat_columns = nested_metadata_to_flat_columns(metadata_list)304    if flat_columns:305        col_texts = column_flat_rows_to_texts(flat_columns)306        col_embeddings = build_embeddings(col_texts, model)307        build_column_faiss_index(schema_name, col_embeddings, flat_columns)308    else:309        logger.warning("No column rows after flatten for %s; skipping column index.", schema_name)310 311 312def main():313    logger.info("Starting metadata-to-FAISS pipeline.")314    engine = get_engine()315    model = get_embedding_model()316    for schema_name in SCHEMAS:317        run_pipeline_for_schema(engine, schema_name, model)318    logger.info("Pipeline finished.")319 320 321if __name__ == "__main__":322    main()323