Team Ai
Apppublic

ritvik360/nl2sql-bench

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
augmentor.py289 linesDownload Raw Back to data_factory
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