Avinashsappati/DBMS-NL2SQL
0
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 