ritvik360/nl2sql-bench
0
1"""2data_factory/augmentor.py3==========================4Rule-based Natural Language augmentation.5 6These transformations operate ONLY on NL question strings.7SQL is NEVER modified — it always comes from the verified template library.8 9Three augmentation strategies:10 1. Synonym replacement — swaps domain words with semantically equivalent ones11 2. Condition reordering — shuffles conjunctive phrases (preserves meaning)12 3. Date normalisation — expresses dates in different formats when applicable13"""14 15from __future__ import annotations16 17import random18import re19from copy import deepcopy20from typing import Iterator21 22 23# ─────────────────────────────────────────────────────────────────────────────24# SYNONYM DICTIONARIES25# ─────────────────────────────────────────────────────────────────────────────26 27# Format: "canonical_term": ["synonym1", "synonym2", ...]28# All synonyms are semantically equivalent in a business context.29 30_SYNONYMS: dict[str, list[str]] = {31 32 # Verbs / action starters33 "list": ["show", "display", "return", "give me", "find", "retrieve"],34 "show": ["list", "display", "return", "get", "retrieve"],35 "find": ["identify", "locate", "get", "show", "retrieve", "look up"],36 "return": ["show", "give", "list", "retrieve", "output"],37 "retrieve": ["fetch", "get", "return", "pull"],38 "get": ["retrieve", "fetch", "return", "give me"],39 40 # Aggregation words41 "total": ["sum", "aggregate", "overall", "cumulative", "combined"],42 "average": ["mean", "avg", "typical"],43 "count": ["number of", "quantity of", "how many"],44 "highest": ["largest", "maximum", "top", "greatest"],45 "lowest": ["smallest", "minimum", "least"],46 47 # Business / domain48 "customer": ["client", "buyer", "user", "account holder", "shopper"],49 "customers": ["clients", "buyers", "users", "account holders", "shoppers"],50 "product": ["item", "SKU", "article", "goods"],51 "products": ["items", "SKUs", "articles", "goods"],52 "order": ["purchase", "transaction", "sale"],53 "orders": ["purchases", "transactions", "sales"],54 "revenue": ["income", "earnings", "sales amount", "money earned"],55 "spending": ["expenditure", "spend", "purchases"],56 "amount": ["value", "sum", "total", "figure"],57 "price": ["cost", "rate", "charge", "fee"],58 59 # Healthcare60 "patient": ["person", "individual", "case"],61 "patients": ["persons", "individuals", "cases"],62 "doctor": ["physician", "clinician", "practitioner", "specialist"],63 "doctors": ["physicians", "clinicians", "practitioners"],64 "appointment": ["visit", "consultation", "session"],65 "appointments": ["visits", "consultations", "sessions"],66 "medication": ["drug", "medicine", "pharmaceutical", "prescription drug"],67 "medications": ["drugs", "medicines", "pharmaceuticals"],68 "diagnosis": ["condition", "finding", "medical finding"],69 70 # Finance71 "account": ["bank account", "profile", "portfolio entry"],72 "accounts": ["bank accounts", "profiles"],73 "loan": ["credit", "borrowing", "debt instrument"],74 "loans": ["credits", "borrowings", "debt instruments"],75 "transaction": ["transfer", "payment", "operation", "activity"],76 "transactions": ["transfers", "payments", "operations"],77 "balance": ["funds", "available amount", "account balance"],78 79 # HR80 "employee": ["staff member", "worker", "team member", "headcount"],81 "employees": ["staff", "workers", "team members", "workforce"],82 "department": ["team", "division", "unit", "group"],83 "departments": ["teams", "divisions", "units"],84 "salary": ["pay", "compensation", "remuneration", "earnings"],85 "project": ["initiative", "program", "assignment", "engagement"],86 "projects": ["initiatives", "programs", "assignments"],87 88 # Adjectives / Qualifiers89 "active": ["current", "ongoing", "live", "existing"],90 "delivered": ["completed", "fulfilled", "received"],91 "cancelled": ["voided", "aborted", "terminated"],92 "alphabetically": ["by name", "in alphabetical order", "A to Z"],93 "descending": ["from highest to lowest", "in decreasing order", "largest first"],94 "ascending": ["from lowest to highest", "in increasing order", "smallest first"],95 "distinct": ["unique", "different"],96 "in stock": ["available", "with available inventory", "not out of stock"],97}98 99 100# ─────────────────────────────────────────────────────────────────────────────101# DATE PHRASE PATTERNS102# These will be replaced with alternative date expressions.103# ─────────────────────────────────────────────────────────────────────────────104 105_DATE_ALTERNATES: list[tuple[str, list[str]]] = [106 # ISO partial107 ("2024-01-01", ["January 1st 2024", "Jan 1, 2024", "the start of 2024", "2024 start"]),108 ("2023-01-01", ["January 1st 2023", "Jan 1, 2023", "the start of 2023"]),109 ("2025-01-01", ["January 1st 2025", "the start of 2025"]),110 # Quarter references111 ("Q1", ["the first quarter", "January through March", "Jan-Mar"]),112 ("Q2", ["the second quarter", "April through June", "Apr-Jun"]),113 ("Q3", ["the third quarter", "July through September", "Jul-Sep"]),114 ("Q4", ["the fourth quarter", "October through December", "Oct-Dec"]),115 # Year references116 ("in 2024", ["during 2024", "throughout 2024", "for the year 2024"]),117 ("in 2023", ["during 2023", "throughout 2023", "for the year 2023"]),118]119 120 121# ─────────────────────────────────────────────────────────────────────────────122# CONDITION REORDERING123# Splits on "and" between two conditions and reverses them.124# ─────────────────────────────────────────────────────────────────────────────125 126def _reorder_conditions(text: str, rng: random.Random) -> str:127 """128 If the text contains ' and ' connecting two distinct clauses,129 randomly swap their order 50% of the time.130 131 Example:132 "active employees earning above $100,000"133 → "employees earning above $100,000 that are active"134 """135 # Only attempt if "and" is present as a clause connector136 matches = list(re.finditer(r'\b(?:and|who are|that are|with)\b', text, re.IGNORECASE))137 if not matches or rng.random() > 0.5:138 return text139 140 # Take the first match and swap text around it141 m = matches[0]142 before = text[:m.start()].strip()143 after = text[m.end():].strip()144 connector = m.group(0).lower()145 146 # Build swapped version147 if connector in ("and",):148 swapped = f"{after} and {before}"149 else:150 swapped = f"{after} {connector} {before}"151 152 # Return swapped only if it doesn't break grammar badly153 # (heuristic: swapped should not start with a verb)154 if swapped and not swapped[0].isupper():155 swapped = swapped[0].upper() + swapped[1:]156 return swapped157 158 159# ─────────────────────────────────────────────────────────────────────────────160# SYNONYM REPLACEMENT161# ─────────────────────────────────────────────────────────────────────────────162 163def _apply_synonyms(text: str, rng: random.Random, max_replacements: int = 3) -> str:164 """165 Replace up to `max_replacements` words/phrases with synonyms.166 Replacement is probabilistic (50% chance per match) to maintain diversity.167 """168 result = text169 replacements_done = 0170 171 # Shuffle the synonym keys to get different replacement targets each call172 keys = list(_SYNONYMS.keys())173 rng.shuffle(keys)174 175 for canonical in keys:176 if replacements_done >= max_replacements:177 break178 synonyms = _SYNONYMS[canonical]179 # Case-insensitive match on word boundary180 pattern = re.compile(r'\b' + re.escape(canonical) + r'\b', re.IGNORECASE)181 if pattern.search(result) and rng.random() < 0.5:182 replacement = rng.choice(synonyms)183 # Preserve original casing for first character184 def _replace(m: re.Match) -> str:185 original = m.group(0)186 if original[0].isupper():187 return replacement[0].upper() + replacement[1:]188 return replacement189 result = pattern.sub(_replace, result, count=1)190 replacements_done += 1191 192 return result193 194 195# ─────────────────────────────────────────────────────────────────────────────196# DATE FORMAT VARIATION197# ─────────────────────────────────────────────────────────────────────────────198 199def _vary_dates(text: str, rng: random.Random) -> str:200 """Replace date phrases with alternate representations."""201 result = text202 for phrase, alternates in _DATE_ALTERNATES:203 if phrase.lower() in result.lower() and rng.random() < 0.6:204 alt = rng.choice(alternates)205 result = re.sub(re.escape(phrase), alt, result, count=1, flags=re.IGNORECASE)206 return result207 208 209# ─────────────────────────────────────────────────────────────────────────────210# PUBLIC API211# ─────────────────────────────────────────────────────────────────────────────212 213def augment_nl(214 nl_question: str,215 n: int = 3,216 seed: int = 42,217) -> list[str]:218 """219 Generate `n` rule-based augmented variants of a natural language question.220 221 Each variant applies a different combination of:222 - synonym replacement223 - condition reordering224 - date format variation225 226 The original question is NOT included in the output.227 228 Parameters229 ----------230 nl_question : str231 The base NL question to augment.232 n : int233 Number of variants to generate.234 seed : int235 Random seed for reproducibility.236 237 Returns238 -------239 list[str]240 Up to `n` distinct augmented strings. May be fewer if the question241 is too short to vary meaningfully.242 """243 rng = random.Random(seed)244 variants: list[str] = []245 seen: set[str] = {nl_question}246 247 strategies = [248 # Strategy 1: synonym only249 lambda t, r: _apply_synonyms(t, r, max_replacements=2),250 # Strategy 2: synonym + date251 lambda t, r: _vary_dates(_apply_synonyms(t, r, max_replacements=2), r),252 # Strategy 3: condition reorder + synonym253 lambda t, r: _apply_synonyms(_reorder_conditions(t, r), r, max_replacements=1),254 # Strategy 4: heavy synonym255 lambda t, r: _apply_synonyms(t, r, max_replacements=4),256 # Strategy 5: date only257 lambda t, r: _vary_dates(t, r),258 ]259 260 for i in range(n * 3): # Over-generate, then deduplicate261 strategy = strategies[i % len(strategies)]262 # Use a different seed offset per variant attempt263 local_rng = random.Random(seed + i * 31)264 candidate = strategy(nl_question, local_rng).strip()265 266 # Normalise whitespace267 candidate = " ".join(candidate.split())268 269 if candidate and candidate not in seen:270 seen.add(candidate)271 variants.append(candidate)272 273 if len(variants) >= n:274 break275 276 return variants277 278 279def generate_all_augmentations(280 nl_question: str,281 seed: int = 42,282 n_per_template: int = 3,283) -> Iterator[str]:284 """285 Yield augmented NL variants one at a time (generator).286 Suitable for streaming into a large dataset without memory pressure.287 """288 yield from augment_nl(nl_question, n=n_per_template, seed=seed)289 