CoolFace
Apppublic

kumar6591/data-quality-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
algorithm_portfolio.py136 linesDownload Raw Back to env
1from __future__ import annotations2 3import itertools4import re5from dataclasses import dataclass6from typing import Iterable7 8 9@dataclass(frozen=True)10class AlgoConfig:11    w_coverage: float12    w_stat: float13    w_risk: float14    w_novelty: float15    limit_bonus: float16    repeat_penalty: float17 18 19def _query_features(sql: str) -> dict[str, float]:20    s = (sql or "").lower()21    return {22        "coverage": float(any(k in s for k in ["count(", "sum(", "avg(", "group by", "distinct"])),23        "stat": float(any(k in s for k in ["avg(", "stddev", "variance", "percentile", "try_cast", "strptime"])),24        "risk": float(any(k in s for k in ["drop", "truncate", "delete", "insert", "update", "alter", "create"])),25        "novelty": float(any(k in s for k in ["left join", "except", "not in", "having", "case when"])),26        "has_limit": float("limit" in s),27    }28 29 30def _task_keywords(task_id: int) -> list[str]:31    if task_id == 1:32        return ["null", "email", "customer_id", "duplicate", "group by"]33    if task_id == 2:34        return ["quantity", "amount", "n/a", "try_cast", "order_date"]35    return ["transactions_baseline", "transactions_current", "category", "user_id", "avg(amount)"]36 37 38def _task_relevance(task_id: int, sql: str) -> float:39    s = (sql or "").lower()40    keys = _task_keywords(task_id)41    hits = sum(1 for k in keys if k in s)42    return hits / max(1, len(keys))43 44 45def _sql_shape_penalty(sql: str) -> float:46    # Penalize very long and likely redundant SQL in a constrained step budget.47    length = len(sql or "")48    if length < 120:49        return 0.050    if length < 300:51        return 0.0252    return 0.0553 54 55def algorithm_config_stream() -> Iterable[AlgoConfig]:56    # 11^4 * 7^2 = 717,409 total algorithm configurations.57    grid_a = [i / 10 for i in range(0, 11)]58    grid_b = [i / 20 for i in range(0, 7)]59    for a, b, c, d, e, f in itertools.product(grid_a, grid_a, grid_a, grid_a, grid_b, grid_b):60        yield AlgoConfig(61            w_coverage=a,62            w_stat=b,63            w_risk=c,64            w_novelty=d,65            limit_bonus=e,66            repeat_penalty=f,67        )68 69 70def _config_query_score(task_id: int, sql: str, cfg: AlgoConfig, q_prior: float) -> float:71    f = _query_features(sql)72    relevance = _task_relevance(task_id, sql)73    penalty_len = _sql_shape_penalty(sql)74    score = (75        cfg.w_coverage * f["coverage"]76        + cfg.w_stat * f["stat"]77        + cfg.w_novelty * f["novelty"]78        + cfg.limit_bonus * f["has_limit"]79        + 0.6 * relevance80        + 0.4 * q_prior81        - cfg.w_risk * f["risk"]82        - penalty_len83    )84    return score85 86 87def _ranking_for_config(task_id: int, queries: list[str], cfg: AlgoConfig, priors: list[float]) -> list[int]:88    pairs = []89    for i, q in enumerate(queries):90        pairs.append((i, _config_query_score(task_id, q, cfg, priors[i])))91    pairs.sort(key=lambda x: x[1], reverse=True)92    return [i for i, _ in pairs]93 94 95def select_best_config(task_id: int, queries: list[str], priors: list[float], max_configs: int = 100_000) -> AlgoConfig:96    best_cfg = None97    best_obj = -10**998 99    for idx, cfg in enumerate(algorithm_config_stream()):100        if idx >= max_configs:101            break102        ranking = _ranking_for_config(task_id, queries, cfg, priors)103 104        # Objective: prioritize top-2 quality and diversity in SQL intent.105        top = ranking[:2]106        top_score = sum(_config_query_score(task_id, queries[i], cfg, priors[i]) for i in top)107 108        intents = set()109        for i in top:110            s = queries[i].lower()111            intent = "join" if any(k in s for k in ["join", "except", "not in"]) else "agg"112            intents.add(intent)113        diversity_bonus = 0.05 if len(intents) > 1 else 0.0114 115        obj = top_score + diversity_bonus116        if obj > best_obj:117            best_obj = obj118            best_cfg = cfg119 120    return best_cfg if best_cfg is not None else AlgoConfig(0.5, 0.5, 1.0, 0.5, 0.0, 0.0)121 122 123def ensemble_order(task_id: int, queries: list[str], priors: list[float], max_configs: int = 100_000) -> list[str]:124    cfg = select_best_config(task_id, queries, priors, max_configs=max_configs)125    ranking = _ranking_for_config(task_id, queries, cfg, priors)126 127    # De-prioritize unsafe SQL just in case external user-provided probes are included.128    safe = []129    unsafe = []130    for i in ranking:131        if re.search(r"\b(drop|truncate|delete|insert|update|alter|create)\b", queries[i], re.IGNORECASE):132            unsafe.append(queries[i])133        else:134            safe.append(queries[i])135    return safe + unsafe136