feat(factor): expression_match 防换皮模块——Python ast 交换律归一化+最大公共子树占比(QuantaAlpha 算法思想方言内实现) [nas]
This commit is contained in:
@@ -0,0 +1,92 @@
|
||||
# sanguo_factor/expression_match.py
|
||||
"""表达式结构相似度(AST 防换皮,2026-10-10 spec §4.2).
|
||||
|
||||
QuantaAlpha factor_ast 算法思想(最大公共子树+交换律)的方言内实现:
|
||||
不移植其 qlib 式解析器,直接用 Python ast——与 factor_guard 白名单
|
||||
同方言,注册的新因子表达式必然可解析.判定永远人做:本模块只产提示.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
|
||||
_COMMUTATIVE = (ast.Add, ast.Mult)
|
||||
|
||||
|
||||
def _is_meta(node: ast.AST) -> bool:
|
||||
"""元数据子节点:expr_context(Load/Store)与算子(Div/Mult...)——
|
||||
前者是 Name 的语境标记,后者的类型已编码进 BinOp 的 key.
|
||||
不跳过则 ? 兜底签名全局撞车+虚增计数(Name 算 2/整树多 1)."""
|
||||
return isinstance(node, (ast.expr_context, ast.operator))
|
||||
|
||||
|
||||
def _node_size(node: ast.AST) -> int:
|
||||
return 1 + sum(_node_size(c) for c in ast.iter_child_nodes(node)
|
||||
if not _is_meta(c))
|
||||
|
||||
|
||||
def _key(node: ast.AST) -> str:
|
||||
"""规范化结构签名:交换律 binop 左右子树排序后拼接."""
|
||||
if isinstance(node, ast.BinOp) and isinstance(node.op, _COMMUTATIVE):
|
||||
lk, rk = _key(node.left), _key(node.right)
|
||||
a, b = sorted((lk, rk))
|
||||
return f"({a}|{b}|{type(node.op).__name__})"
|
||||
if isinstance(node, ast.BinOp):
|
||||
return f"({_key(node.left)}>{_key(node.right)}|{type(node.op).__name__})"
|
||||
if isinstance(node, ast.Call) and isinstance(node.func, ast.Name):
|
||||
args = ",".join(_key(a) for a in node.args)
|
||||
return f"{node.func.id}({args})"
|
||||
if isinstance(node, ast.Name):
|
||||
return f"#{node.id}"
|
||||
if isinstance(node, ast.Constant):
|
||||
return f"#{node.value!r}"
|
||||
return f"?{type(node).__name__}"
|
||||
|
||||
|
||||
def _all_subtree_keys(node: ast.AST) -> dict[str, int]:
|
||||
"""子树签名→节点数(同签名取最大)."""
|
||||
out: dict[str, int] = {}
|
||||
for sub in ast.walk(node):
|
||||
if _is_meta(sub):
|
||||
continue
|
||||
k = _key(sub)
|
||||
out[k] = max(out.get(k, 0), _node_size(sub))
|
||||
return out
|
||||
|
||||
|
||||
def similarity(expr_a: str, expr_b: str) -> float | None:
|
||||
"""最大公共子树节点数 / 较小表达式节点数;任一解析失败=None.
|
||||
|
||||
取 .body 剥掉 Expression 包装节点——否则其签名落 ? 兜底桶,
|
||||
任意两表达式的根都会撞签(子树=全树,相似度恒 1.0).
|
||||
"""
|
||||
try:
|
||||
ta = ast.parse(expr_a, mode="eval").body
|
||||
tb = ast.parse(expr_b, mode="eval").body
|
||||
except (SyntaxError, ValueError):
|
||||
return None
|
||||
sa, sb = _node_size(ta), _node_size(tb)
|
||||
if sa == 0 or sb == 0:
|
||||
return None
|
||||
keys_a = _all_subtree_keys(ta)
|
||||
best = 0
|
||||
for sub in ast.walk(tb):
|
||||
if _is_meta(sub):
|
||||
continue
|
||||
k = _key(sub)
|
||||
if k in keys_a:
|
||||
best = max(best, min(keys_a[k], _node_size(sub)))
|
||||
return round(best / min(sa, sb), 4)
|
||||
|
||||
|
||||
def top_similar(expression: str, candidates: dict[str, str],
|
||||
floor: float = 0.6) -> list[dict]:
|
||||
"""对候选池按相似度排序,过滤低于 floor 的(提示非硬拒)."""
|
||||
out: list[dict] = []
|
||||
for name, expr in candidates.items():
|
||||
if not expression or not expr:
|
||||
continue
|
||||
r = similarity(expression, expr)
|
||||
if r is not None and r >= floor:
|
||||
out.append({"name": name, "ratio": r})
|
||||
out.sort(key=lambda h: h["ratio"], reverse=True)
|
||||
return out
|
||||
@@ -0,0 +1,34 @@
|
||||
# tests/factor/test_expression_match.py
|
||||
"""AST 防换皮(spec §4.2 2026-10-10):交换律归一化+最大公共子树占比."""
|
||||
from sanguo_factor.expression_match import similarity, top_similar
|
||||
|
||||
|
||||
def test_identical_and_commutative():
|
||||
assert similarity("close/open", "close/open") == 1.0
|
||||
# 交换律:a+b ≡ b+a(归一化后同构)
|
||||
assert similarity("close + open", "open + close") == 1.0
|
||||
assert similarity("ts_mean(close, 5) * volume",
|
||||
"volume * ts_mean(close, 5)") == 1.0
|
||||
|
||||
|
||||
def test_partial_common_subtree():
|
||||
# 公共子树=ts_mean(close,5)(5节点);分母=较小表达式(7节点)
|
||||
r = similarity("ts_mean(close, 5) / volume",
|
||||
"ts_mean(close, 5) * turnover")
|
||||
assert r is not None and 0.6 < r < 1.0
|
||||
|
||||
|
||||
def test_no_common_and_invalid():
|
||||
assert similarity("close / open", "volume * turnover") == 0.0
|
||||
assert similarity("close +++", "close/open") is None # 解析失败如实 None
|
||||
assert similarity("", "close") is None
|
||||
|
||||
|
||||
def test_top_similar_floor():
|
||||
cands = {"fa_a": "close + open", "fa_b": "open + close",
|
||||
"fa_c": "volume * turnover"}
|
||||
hits = top_similar("close + open", cands, floor=0.6)
|
||||
assert [h["name"] for h in hits] == ["fa_a", "fa_b"]
|
||||
assert all(h["ratio"] >= 0.6 for h in hits)
|
||||
# 空表达式/坏 candidates 不炸
|
||||
assert top_similar("close", {"fa_x": ""}) == []
|
||||
Reference in New Issue
Block a user