Team Ai
Apppublic

joelgilbert/NL2SQL

sourceHugging Facemitupdated 11mo agoView on Hugging Face
0likes
embeddings.py142 linesDownload Raw Back to vector_store
1"""2Embedding generation utilities for vector search.3 4Note: Upstash Vector provides built-in embedding generation,5so this module provides helper functions for text preparation.6"""7 8import logging9import re10from typing import Dict, List11 12logger = logging.getLogger(__name__)13 14 15class EmbeddingHelper:16    """Helper functions for preparing text for embedding."""17    18    @staticmethod19    def normalize_text(text: str) -> str:20        """21        Normalize text before embedding generation.22        23        Args:24            text: Raw text to normalize25            26        Returns:27            Normalized text28        """29        # Convert to lowercase30        text = text.lower()31        32        # Remove extra whitespace33        text = re.sub(r'\s+', ' ', text)34        35        # Strip leading/trailing whitespace36        text = text.strip()37        38        return text39    40    @staticmethod41    def embed_schema(table_structure: Dict) -> str:42        """43        Format table schema for embedding.44        45        Args:46            table_structure: Table structure dictionary from schema manager47            48        Returns:49            Formatted text suitable for embedding50        """51        table_name = table_structure.get('table_name', '')52        columns = table_structure.get('columns', [])53        54        # Create descriptive text55        parts = [f"Table {table_name}"]56        57        # Add column information58        if columns:59            column_desc = []60            for col in columns:61                col_name = col.get('column_name', '')62                col_type = col.get('data_type', '')63                column_desc.append(f"{col_name} {col_type}")64            65            parts.append(f"with columns: {', '.join(column_desc)}")66        67        text = " ".join(parts)68        return EmbeddingHelper.normalize_text(text)69    70    @staticmethod71    def embed_question(question: str) -> str:72        """73        Normalize and prepare user question for embedding.74        75        Args:76            question: User's natural language question77            78        Returns:79            Normalized question80        """81        # Basic normalization82        question = EmbeddingHelper.normalize_text(question)83        84        # Remove question marks and punctuation for better matching85        question = re.sub(r'[?!.,;:]', '', question)86        87        return question.strip()88    89    @staticmethod90    def format_example_queries(similar_queries: List[Dict]) -> str:91        """92        Format similar queries into a examples string for prompt.93        94        Args:95            similar_queries: List of similar query dictionaries96            97        Returns:98            Formatted examples text99        """100        if not similar_queries:101            return "No similar examples found."102        103        examples = []104        for i, query in enumerate(similar_queries, 1):105            question = query.get('question', 'N/A')106            sql = query.get('sql', 'N/A')107            examples.append(f"Example {i}:\nQuestion: {question}\nSQL: {sql}\n")108        109        return "\n".join(examples)110    111    @staticmethod112    def expand_abbreviations(text: str) -> str:113        """114        Expand common abbreviations in queries.115        116        Args:117            text: Text with potential abbreviations118            119        Returns:120            Text with expanded abbreviations121        """122        abbreviations = {123            r'\bqty\b': 'quantity',124            r'\bamt\b': 'amount',125            r'\btot\b': 'total',126            r'\bavg\b': 'average',127            r'\bmax\b': 'maximum',128            r'\bmin\b': 'minimum',129            r'\bnum\b': 'number',130            r'\bpct\b': 'percent',131            r'\bid\b': 'identifier',132        }133        134        for abbrev, full in abbreviations.items():135            text = re.sub(abbrev, full, text, flags=re.IGNORECASE)136        137        return text138 139 140# Global embedding helper instance141embedding_helper = EmbeddingHelper()142