Team Ai
Apppublic

Avinashsappati/DBMS-NL2SQL

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
test_model.py314 linesDownload Raw Back to root
1"""### TRANSFORMER And etc.."""2import re3import json4import numpy as np5import torch6import sqlparse,argparse7 8from sentence_transformers import SentenceTransformer9from sklearn.metrics.pairwise import cosine_similarity10from transformers import T5ForConditionalGeneration, T5Tokenizer11 12# Configuration13MODEL_NAME       = "yashwantk05/t5-finetuned"14BGE_MODEL_NAME   = "BAAI/bge-base-en-v1.5"15MAX_INPUT_LENGTH = 102416MAX_NEW_TOKENS   = 25617NUM_BEAMS        = 218 19# Retrieval thresholds20SIMILARITY_THRESHOLD  = 0.45   # min table score to include21VAGUE_THRESHOLD       = 0.30   # below this → ask for clarification22MAX_TABLES_FALLBACK   = 523TOP_K_COLUMNS         = 1224 25# Schema loader26def load_schemas(tables_path: str) -> dict:27    with open(tables_path) as f:28        raw = json.load(f)29 30    schemas = {}31    for db in raw:32        db_id   = db["db_id"]33        t_names = db["table_names_original"]34        c_names = db["column_names_original"]35        c_types = db["column_types"]36        pk_ids  = set(db.get("primary_keys", []))37        fk_pairs= db.get("foreign_keys", [])38 39        col_lookup = {}40        for cid, (tid, cname) in enumerate(c_names):41            if tid == -1:42                continue43            col_lookup[cid] = (t_names[tid], cname)44 45        # Crucial for JOIN Operations46        fk_map = {}47        for (src_cid, dst_cid) in fk_pairs:48            if dst_cid in col_lookup:49                ref_table, ref_col = col_lookup[dst_cid]50                fk_map[src_cid] = f"{ref_table}.{ref_col}"51 52        columns = {t: [] for t in t_names}53        pks     = {t: set() for t in t_names}54        fks     = {t: {} for t in t_names}55 56        for cid, (tid, cname) in enumerate(c_names):57            if tid == -1:58                continue59            tname = t_names[tid]60            ctype = c_types[cid]61            columns[tname].append((cname, ctype.upper()))62            if cid in pk_ids:63                pks[tname].add(cname)64            if cid in fk_map:65                fks[tname][cname] = fk_map[cid]66 67        schemas[db_id] = dict(tables=t_names, columns=columns, pks=pks, fks=fks)68 69    return schemas70 71# Schema Representations72 73# Understanding schema semantically .74def build_table_text(table_name, columns, pks, fks) -> str:75    parts = []76    for col_name, col_type in columns:77        tags = []78        if col_name in pks:79            tags.append("PK")80        if col_name in fks:81            tags.append(f"FK→{fks[col_name]}")82        tags.append(col_type)83        parts.append(f"{col_name} ({' '.join(tags)})")84    return f"{table_name}: {', '.join(parts)}"85 86# For Embedding similarity87def build_column_texts(schema) -> tuple:88    texts, labels = [], []89    for tname in schema["tables"]:90        pks = schema["pks"][tname]91        fks = schema["fks"][tname]92        for col_name, col_type in schema["columns"][tname]:93            tags = []94            if col_name in pks:95                tags.append("PK")96            if col_name in fks:97                tags.append(f"FK→{fks[col_name]}")98            tags.append(col_type)99            texts.append(f"{tname}.{col_name} ({' '.join(tags)})")100            labels.append(f"{tname}.{col_name}")101    return texts, labels102 103# BGE Retriever104class SchemaRetriever:105    def __init__(self, model_name: str = BGE_MODEL_NAME):106        print(f"Loading BGE model: {model_name} ...")107        self.model = SentenceTransformer(model_name)108        self.model.eval()109 110    def retrieve(111        self,112        question: str,113        schema: dict,114        threshold: float = SIMILARITY_THRESHOLD,115        max_tables: int  = MAX_TABLES_FALLBACK,116        top_k_cols: int  = TOP_K_COLUMNS,117    ) -> tuple[str, float]:118        col_texts, col_labels = build_column_texts(schema)119        if not col_texts:120            return "", 0.0121 122        q_emb   = self.model.encode([question], normalize_embeddings=True)123        col_emb = self.model.encode(col_texts,  normalize_embeddings=True)124        col_scores = cosine_similarity(q_emb, col_emb)[0] # finding relevant columns only125 126        k = min(top_k_cols, len(col_texts))127        top_col_idx = col_scores.argsort()[-k:][::-1]128        needed_tables = {col_labels[i].split(".")[0] for i in top_col_idx}129 130        table_texts = [131            build_table_text(t, schema["columns"][t], schema["pks"][t], schema["fks"][t])132            for t in schema["tables"]133        ]134        tbl_emb    = self.model.encode(table_texts, normalize_embeddings=True)135        tbl_scores = cosine_similarity(q_emb, tbl_emb)[0]136 137        max_score = float(tbl_scores.max())138 139        dynamic_tables = {140            schema["tables"][i]141            for i, s in enumerate(tbl_scores)142            if s >= threshold143        }144        if not dynamic_tables:145            dynamic_tables = {schema["tables"][tbl_scores.argmax()]}146 147        dynamic_tables  = set(list(dynamic_tables)[:max_tables])148        selected_tables = set(list(needed_tables | dynamic_tables)[:max_tables])149 150        create_stmts = []151        for tname in schema["tables"]:152            if tname not in selected_tables:153                continue154            pks  = schema["pks"][tname]155            fks  = schema["fks"][tname]156            cols = schema["columns"][tname]157            col_defs = []158            for col_name, col_type in cols:159                constraint = ""160                if col_name in pks:161                    constraint += " PRIMARY KEY"162                if col_name in fks:163                    constraint += f" REFERENCES {fks[col_name]}"164                col_defs.append(f"  {col_name} {col_type}{constraint}")165            stmt = f"CREATE TABLE {tname} (\n" + ",\n".join(col_defs) + "\n);"166            create_stmts.append(stmt)167 168        return "\n".join(create_stmts), max_score169 170# Prompt Builder171def build_prompt(question: str, schema_str: str) -> str:172    return f"translate English to SQL: {question} schema: {schema_str}"173 174 175# SQL Validator176def validate_sql(sql: str, schema: dict) -> tuple[bool, list[str]]:177 178    errors = []179 180    # Level 1: Syntax — sqlparse can parse it without errors181    try:182        parsed = sqlparse.parse(sql)183        if not parsed or not parsed[0].tokens:184            errors.append("Syntax error: could not parse the generated SQL.")185            return False, errors186    except Exception as e:187        errors.append(f"Syntax error: {e}")188        return False, errors189 190    #  Level 2: Semantic — all table/column names exist in schema191    all_tables  = {t.lower() for t in schema["tables"]}192    all_columns = set()193    for tname in schema["tables"]:194        for col_name, _ in schema["columns"][tname]:195            all_columns.add(col_name.lower())196            all_columns.add(f"{tname.lower()}.{col_name.lower()}")197 198    sql_upper = sql.upper()199 200    # Extract table names after FROM / JOIN201    table_pattern = re.findall(202        r'\b(?:FROM|JOIN)\s+([a-zA-Z_][a-zA-Z0-9_]*)', sql, re.IGNORECASE203    )204    for t in table_pattern:205        if t.lower() not in all_tables:206            errors.append(f"Semantic error: table '{t}' does not exist in the schema.")207 208    if errors:209        return False, errors210 211    return True, []212 213# Inference engine214class Text2SQLEngine:215 216    def __init__(self, tables_path: str):217 218        self.device = "cuda" if torch.cuda.is_available() else "cpu"219        print(f"Device: {self.device}")220 221        print("Loading schemas ...")222        self.schemas   = load_schemas(tables_path)223        self.retriever = SchemaRetriever()224        self.retriever.model.to(self.device)225 226        print(f"Loading T5 model: {MODEL_NAME} ...")227        self.tokenizer = T5Tokenizer.from_pretrained(MODEL_NAME)228        self.model     = T5ForConditionalGeneration.from_pretrained(MODEL_NAME)229        self.model.to(self.device)230        self.model.eval()231 232    # ENTRY POINT233    def generate(self, question: str, db_id: str) -> dict:234        """235        Returns a dictIonary with:236            status      : "ok" | "vague" | "invalid_db" | "validation_failed"237            sql         : generated SQL string (if status == "ok")238            message     : human-readable message239            confidence  : BGE similarity score240            warnings    : list of validation warnings (if any)241        """242 243        # Checking if db_id exists244        if db_id not in self.schemas:245            available = list(self.schemas.keys())[:10]246            return {247                "status":  "invalid_db",248                "sql":     None,249                "message": f"Database '{db_id}' not found.\nAvailable (first 10): {available}",250                "confidence": 0.0,251                "warnings": [],252            }253 254        schema = self.schemas[db_id]255 256        # BGE retrieval + confidence check257        schema_str, confidence = self.retriever.retrieve(question, schema)258 259        if confidence < VAGUE_THRESHOLD:260            table_list = ", ".join(schema["tables"])261            return {262                "status":  "vague",263                "sql":     None,264                "message": (265                    f"Your question didn't clearly match anything in the '{db_id}' database "266                    f"(confidence: {confidence:.2f}).\n"267                    f"Could you be more specific?\n\n"268                    f"Available tables: {table_list}\n\n"269                    f"Example: instead of 'show me data', try 'list all singers with age above 30'."270                ),271                "confidence": confidence,272                "warnings": [],273            }274 275        # Building prompt and generating SQL276        prompt = build_prompt(question, schema_str)277 278        inputs = self.tokenizer(279            prompt,280            return_tensors="pt",281            max_length=MAX_INPUT_LENGTH,282            truncation=True,283        ).to(self.device)284 285        with torch.no_grad():286            output_ids = self.model.generate(287                **inputs,288                max_new_tokens=MAX_NEW_TOKENS,289                num_beams=NUM_BEAMS,290                early_stopping=True,291            )292 293        sql = self.tokenizer.decode(output_ids[0], skip_special_tokens=True).strip()294 295        # Validate generated SQL296        is_valid, errors = validate_sql(sql, schema)297 298        if not is_valid:299            return {300                "status":   "validation_failed",301                "sql":      sql,302                "message":  "SQL was generated but failed validation. Use with caution.",303                "confidence": confidence,304                "warnings": errors,305            }306 307        return {308            "status":     "ok",309            "sql":        sql,310            "message":    "Query generated successfully.",311            "confidence": confidence,312            "warnings":   [],313        }314