Spaces:
Running
Running
File size: 6,894 Bytes
df42e17 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 | """ํน์ง ํ๋ก๊ทธ๋จ(feature program) DSL โ ๊ฒ์ฆ โ ์์ ์ ์ฝ์ ๋ถ๊ฐํ SQL โ DuckDB ์คํ.
ํนํ ๊ตฌ์ฑ (c)(d): ์๋ฏธ ๋ชจ๋ธ์ ์กฐ์ธ ๊ฒฝ๋ก๋ก ํ์ ๋ฒ์๋ฅผ ํ์ ํ๊ณ , ๋ชจ๋ ํ๋ก๊ทธ๋จ์
'๊ด์ธก ๊ธฐ์ค ์์ (seed_time) ์ด์ ๋ ์ฝ๋๋ง' ์ด๋ผ๋ ์์ ์ ์ฝ์ ์ปดํ์ผ๋ฌ๊ฐ ๊ฐ์ ๋ก ๋ถ๊ฐํ๋ค.
LLM ์ ์์ ์ ์ฝ์ ์ฐ์ง ์๋๋ค โ ์ธ ์ ์๋ค.
"""
from __future__ import annotations
import re
from dataclasses import dataclass
import duckdb
import pandas as pd
from semantic_model import SemanticModel
AGGS = {"count": "count", "sum": "sum", "mean": "avg", "min": "min", "max": "max",
"std": "stddev_samp", "nunique": "count(distinct", "last": "arg_max"}
WINDOWS = (30, 90, 365, 1095, 3650, None) # days; None = ์ ์ฒด ์ด๋ ฅ
# expr ์์ ํ์ฉ๋๋ ๋น-์ปฌ๋ผ ํ ํฐ (DuckDB ์ค์นผ๋ผ ํํ์ ๋ถ๋ถ์งํฉ). ์ด ๋ฐ์ ์๋ณ์๋ ์ ๋ถ ๊ฑฐ๋ถ.
SQL_WORDS = {
"case", "when", "then", "else", "end", "is", "null", "not", "and", "or", "in", "between", "like",
"true", "false", "as", "cast", "integer", "int", "double", "float", "bigint", "date", "timestamp",
"coalesce", "nullif", "abs", "round", "floor", "ceil", "ln", "log", "sqrt", "greatest", "least",
"date_diff", "datediff", "date_part", "extract", "year", "month", "day", "interval", "days",
"seed_time", "epoch", "sign", "power",
}
IDENT = re.compile(r"[A-Za-z_][A-Za-z0-9_.]*")
STRING = re.compile(r"'[^']*'")
@dataclass
class FeatureProgram:
name: str
path: list[str] # entity table ์์ ์์ํ๋ ์กฐ์ธ ๊ฒฝ๋ก
agg: str # AGGS ํค, ๋๋ "none" (์ํฐํฐ ์์ฒด ์์ฑ)
expr: str # ๋ฆฌํ/๊ฒฝ๋ก ํ
์ด๋ธ ์ปฌ๋ผ์ ๋ํ ์ค์นผ๋ผ์
window_days: int | None # ์์ ์ ์ฝ ์ฐฝ; None = ์ ์ฒด ์ด๋ ฅ
rationale: str = ""
@classmethod
def from_dict(cls, d: dict) -> "FeatureProgram":
return cls(name=str(d["name"]), path=list(d["path"]), agg=str(d.get("agg", "none")).lower(),
expr=str(d["expr"]).strip(), window_days=d.get("window_days"), rationale=str(d.get("rationale", "")))
def to_dict(self) -> dict:
return {"name": self.name, "path": self.path, "agg": self.agg, "expr": self.expr,
"window_days": self.window_days, "rationale": self.rationale}
def validate(fp: FeatureProgram, sm: SemanticModel, entity_table: str) -> str | None:
"""None = ํต๊ณผ. ๋ฌธ์์ด = ๊ฑฐ๋ถ ์ฌ์ . (์ ์ธ ์ ๋ ํ
์ด๋ธยท์ปฌ๋ผยท๊ฒฝ๋ก๋ ์ฐธ์กฐ ์์ฒด๋ฅผ ๊ฑฐ๋ถ)"""
if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]{0,63}", fp.name):
return f"bad feature name {fp.name!r}"
if fp.path[:1] != [entity_table]:
return f"path must start at entity table {entity_table}"
if (e := sm.path_error(fp.path)):
return e
if fp.agg == "none":
if len(fp.path) != 1:
return "agg=none is only for entity-table attributes (path length 1)"
elif fp.agg not in AGGS:
return f"unknown agg {fp.agg!r}; use one of {sorted(AGGS)}"
elif sm.time_table(fp.path) is None:
return "aggregation path has no timestamped table; as-of constraint impossible"
if fp.window_days is not None and fp.window_days not in WINDOWS:
return f"window_days must be one of {WINDOWS}"
if ";" in fp.expr or "--" in fp.expr or "/*" in fp.expr:
return "statement separators/comments are not allowed in expr"
allowed = sm.columns_on(fp.path)
dropped = sm.sensitive("drop")
# ๋ฌธ์์ด ๋ฆฌํฐ๋ด ์ ๊ฑฐ ํ ์๋ณ์ ๊ฒ์ฌ
for tok in IDENT.findall(STRING.sub("''", fp.expr)):
if tok.lower() in SQL_WORDS:
continue
if tok not in allowed:
return f"undeclared identifier {tok!r} in expr (allowed: columns of {fp.path})"
if tok in dropped or any(f"{t}.{tok}" in dropped for t in fp.path):
return f"column {tok!r} is classified PII; not usable"
return None
def compile_sql(fp: FeatureProgram, sm: SemanticModel, label_view: str = "L") -> str:
"""L(row_id, entity_id, seed_time) ์ ์กฐ์ธํ์ฌ row_id ๋ณ ํน์ง๊ฐ 1๊ฐ๋ฅผ ๋ด๋ SQL."""
ent = fp.path[0]
pk = sm.tables[ent]["pkey"]
joins = [f"JOIN {ent} ON {ent}.{pk} = {label_view}.entity_id"]
for a, b in zip(fp.path, fp.path[1:]):
joins.append(f"JOIN {b} ON {sm.join_condition(a, b)}")
expr = fp.expr.replace("seed_time", f"{label_view}.seed_time")
if fp.agg == "none":
return f"SELECT {label_view}.row_id, ({expr}) AS value FROM {label_view} {' '.join(joins)}"
# ์์ ์ ์ฝ: ๊ฒฝ๋ก ์์ ์๊ฐ์ด ๋ณด์ ํ
์ด๋ธ *์ ๋ถ* ์ ๋ถ๊ฐํ๋ค (ํ๋๋ง ๊ฑธ๋ฉด ํ๋ฅ ํ
์ด๋ธ๋ก ๋ฏธ๋๊ฐ ์๋ค)
timed = [(t, f"{t}.{sm.tables[t]['time_col']}") for t in fp.path if sm.tables[t]["time_col"]]
tcol = timed[0][1] # last/arg_max ์ ๊ธฐ์ค ์๊ฐ์ด
where = []
for _, c in timed:
where.append(f"{c} < {label_view}.seed_time")
if fp.window_days is not None:
where.append(f"{c} >= {label_view}.seed_time - INTERVAL {int(fp.window_days)} DAY")
if fp.agg == "last":
agg_sql = f"arg_max(({expr}), {tcol})"
elif fp.agg == "nunique":
agg_sql = f"count(distinct ({expr}))"
else:
agg_sql = f"{AGGS[fp.agg]}(({expr}))"
return (f"SELECT {label_view}.row_id, {agg_sql} AS value FROM {label_view} {' '.join(joins)} "
f"WHERE {' AND '.join(where)} GROUP BY {label_view}.row_id")
class Executor:
"""์ 1 ์ปดํจํ
ํ๊ฒฝ: DB ์์์ ํน์ง ํ๋ก๊ทธ๋จ์ ์คํํด ํน์ง ํ๋ ฌ์ ๋ง๋ ๋ค."""
def __init__(self, db_tables: dict[str, pd.DataFrame]):
self.con = duckdb.connect()
for name, df in db_tables.items():
self.con.register(name, df)
def run(self, programs: list[FeatureProgram], sm: SemanticModel, labels: pd.DataFrame,
entity_col: str, time_col: str) -> tuple[pd.DataFrame, dict[str, str]]:
"""labels: ์์ธก ์์ฒญ ํ (entity_col, time_col[, target]). ๋ฐํ: ํน์ง ํ๋ ฌ(row ์์ ์ ์ง), ์คํจ ์ฌ์ ."""
L = pd.DataFrame({"row_id": range(len(labels)), "entity_id": labels[entity_col].values,
"seed_time": pd.to_datetime(labels[time_col]).values})
self.con.register("L", L)
out = pd.DataFrame(index=L.row_id)
errors = {}
for fp in programs:
try:
r = self.con.execute(compile_sql(fp, sm)).df().set_index("row_id")["value"]
out[fp.name] = pd.to_numeric(r.reindex(out.index), errors="coerce").astype("float64")
except Exception as e: # ์คํ ์คํจ๋ ๊ฑฐ๋ถ ์ฌ์ ๋ก ๊ธฐ๋ก, ํ๋ ฌ์์ ์ ์ธ
errors[fp.name] = f"{type(e).__name__}: {str(e).splitlines()[0][:160]}"
return out.reset_index(drop=True), errors
|