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