feat(#35): FeatureSpec 声明式特征定义引擎(PRD 5.3 特征 spec 语义解释)
新增 core/model-framework/feature_spec.py:把超参包中特征 spec 字段从 「非空字符串存在性校验」(issue #39)升级为可解释的声明式特征定义引擎。 - 手写递归下降解析器(tokenizer+parser),安全无 eval/exec;支持算子调用、 裸点位、数值/窗口字面量、嵌套、中文点位、带符号/科学计数数值。 - 不可变 AST(TagRef/Number/Window/OpCall),to_dict/repr 可序列化往返。 - 语义校验:未知算子、arity、参数 kind;返回 SpecIssue 列表供配置台聚合展示。 - 依赖分析 resolve_inputs:去重、稳定顺序,供训练/推理流水线按 tag 拉取数据。 - materialize 执行:EMA/SMA/RollingStd/Max/Min/RateOfChange/Diff/Lag/Log/ Scale/Clip/Combine 共 12 个内置算子,缺数据 fail-fast;numpy 可选依赖。 - 插件注册 register_operator:对齐 PRD「新增结构走插件注册」(呼应 #34 Recipe)。 - 49 个单元测试全通过;零内核改动,纯新增。
This commit is contained in:
@@ -0,0 +1,16 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""测试引导:把 `core/model-framework` 以包名 `model_framework` 挂载到 sys.modules。
|
||||
|
||||
目录名 `model-framework` 含连字符,无法直接以包名 import;挂载后模块内相对导入
|
||||
(`from .feature_spec import ...`)在 unittest 发现机制下可正常解析。
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
|
||||
MF_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
sys.path.insert(0, MF_DIR)
|
||||
if "model_framework" not in sys.modules:
|
||||
pkg = types.ModuleType("model_framework")
|
||||
pkg.__path__ = [MF_DIR]
|
||||
sys.modules["model_framework"] = pkg
|
||||
@@ -0,0 +1,334 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""FeatureSpec 声明式特征定义引擎测试(issue #35)。
|
||||
|
||||
覆盖:
|
||||
1. 解析:算子调用、裸点位、数值/窗口字面量、嵌套、中文点位、带符号数值;
|
||||
2. 解析错误:空 spec、非法字符、括号不匹配、多余内容、参数缺失;
|
||||
3. 语义校验:未知算子、arity 不匹配、参数 kind 错误;
|
||||
4. 依赖分析:resolve_inputs 去重与顺序、嵌套算子依赖汇总;
|
||||
5. 执行:EMA/SMA/RollingStd/RateOfChange/Diff/Lag/Log/Scale/Clip/Combine 的
|
||||
数值正确性,缺失点位 fail-fast;
|
||||
6. 插件注册:register_operator 扩展新算子;
|
||||
7. 往返:to_dict/repr 稳定。
|
||||
"""
|
||||
import math
|
||||
import unittest
|
||||
|
||||
import _bootstrap # noqa: F401 挂载包名
|
||||
|
||||
from model_framework.feature_spec import (
|
||||
FeatureAST,
|
||||
Number,
|
||||
OpCall,
|
||||
OPERATORS,
|
||||
ParseError,
|
||||
SpecIssue,
|
||||
TagRef,
|
||||
Window,
|
||||
describe,
|
||||
materialize,
|
||||
parse,
|
||||
register_operator,
|
||||
resolve_inputs,
|
||||
validate,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 解析
|
||||
# ---------------------------------------------------------------------------
|
||||
class ParseTest(unittest.TestCase):
|
||||
def test_simple_op_with_window(self):
|
||||
ast = parse("EMA(CLF-01.TEMP, 5m)")
|
||||
self.assertEqual(
|
||||
ast, OpCall("EMA", (TagRef("CLF-01.TEMP"), Window(5.0, "m")))
|
||||
)
|
||||
|
||||
def test_simple_op_with_number_window(self):
|
||||
ast = parse("RollingStd(CLF-01.CL2, 10)")
|
||||
self.assertEqual(ast, OpCall("RollingStd", (TagRef("CLF-01.CL2"), Number(10))))
|
||||
|
||||
def test_bare_tag(self):
|
||||
self.assertEqual(parse("炉压"), TagRef("炉压"))
|
||||
|
||||
def test_tag_with_dots_and_dash(self):
|
||||
self.assertEqual(parse("A.B-C_01"), TagRef("A.B-C_01"))
|
||||
|
||||
def test_signed_and_scientific_number(self):
|
||||
ast = parse("Scale(A, -0.5)")
|
||||
self.assertEqual(ast, OpCall("Scale", (TagRef("A"), Number(-0.5))))
|
||||
ast2 = parse("Scale(A, 1e-3)")
|
||||
self.assertAlmostEqual(ast2.args[1].value, 0.001)
|
||||
|
||||
def test_nested_op(self):
|
||||
# 嵌套:外层 Scale,内层 EMA 作为第一个参数点位位置(语法合法,语义由算子判定)
|
||||
ast = parse("Combine(EMA(A, 5m), B)")
|
||||
self.assertEqual(ast.name, "Combine")
|
||||
self.assertEqual(len(ast.args), 2)
|
||||
self.assertEqual(ast.args[0].name, "EMA")
|
||||
|
||||
def test_no_arg_op(self):
|
||||
ast = parse("Diff()")
|
||||
self.assertEqual(ast, OpCall("Diff", ()))
|
||||
|
||||
def test_integer_window_vs_number(self):
|
||||
self.assertEqual(parse("Lag(A, 3)").args[1], Number(3))
|
||||
self.assertEqual(parse("Lag(A, 3m)").args[1], Window(3.0, "m"))
|
||||
|
||||
def test_repr_roundtrip(self):
|
||||
for spec in ["EMA(CLF-01.TEMP, 5m)", "RateOfChange(炉压)", "Clip(P, -1, 1)"]:
|
||||
self.assertEqual(repr(parse(spec)).replace(" ", ""), spec.replace(" ", ""))
|
||||
|
||||
# ---- 解析错误 ----
|
||||
def test_empty_raises(self):
|
||||
with self.assertRaises((ValueError, ParseError)):
|
||||
parse("")
|
||||
with self.assertRaises((ValueError, ParseError)):
|
||||
parse(" ")
|
||||
|
||||
def test_non_string_raises(self):
|
||||
with self.assertRaises(ValueError):
|
||||
parse(123) # type: ignore[arg-type]
|
||||
|
||||
def test_unrecognized_char(self):
|
||||
with self.assertRaises(ParseError) as cm:
|
||||
parse("EMA(A, 5m) @")
|
||||
self.assertIsNotNone(cm.exception.position)
|
||||
|
||||
def test_missing_rparen(self):
|
||||
with self.assertRaises(ParseError):
|
||||
parse("EMA(A, 5m")
|
||||
|
||||
def test_missing_rparen_inner(self):
|
||||
with self.assertRaises(ParseError):
|
||||
parse("EMA(A, (5m)")
|
||||
|
||||
def test_trailing_garbage(self):
|
||||
with self.assertRaises(ParseError):
|
||||
parse("EMA(A, 5m) B")
|
||||
|
||||
def test_missing_arg_after_comma(self):
|
||||
with self.assertRaises(ParseError):
|
||||
parse("EMA(A, )")
|
||||
|
||||
def test_starts_with_paren(self):
|
||||
with self.assertRaises(ParseError):
|
||||
parse("(A)")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 语义校验
|
||||
# ---------------------------------------------------------------------------
|
||||
class ValidateTest(unittest.TestCase):
|
||||
def test_known_op_valid(self):
|
||||
self.assertEqual(validate(parse("EMA(A, 5m)")), [])
|
||||
|
||||
def test_unknown_operator(self):
|
||||
issues = validate(parse("FooBar(A, 5m)"))
|
||||
self.assertEqual(len(issues), 1)
|
||||
self.assertEqual(issues[0].code, "unknown_operator")
|
||||
|
||||
def test_arity_too_few(self):
|
||||
issues = validate(parse("EMA(A)"))
|
||||
self.assertTrue(any(i.code == "arity" for i in issues))
|
||||
|
||||
def test_arity_too_many(self):
|
||||
issues = validate(parse("EMA(A, 5m, 7)"))
|
||||
self.assertTrue(any(i.code == "arity" for i in issues))
|
||||
|
||||
def test_bad_arg_kind_number_where_window(self):
|
||||
# EMA 第二参数允许 window/number,故合法
|
||||
self.assertEqual(validate(parse("EMA(A, 7)")), [])
|
||||
# 但 tag 位置传 number 非法
|
||||
issues = validate(parse("EMA(5, 7)"))
|
||||
self.assertTrue(any(i.code == "bad_arg" for i in issues))
|
||||
|
||||
def test_combine_varargs(self):
|
||||
self.assertEqual(validate(parse("Combine(A, B, C)")), [])
|
||||
issues = validate(parse("Combine(A)"))
|
||||
self.assertTrue(any(i.code == "arity" for i in issues))
|
||||
|
||||
def test_nested_unknown(self):
|
||||
issues = validate(parse("Combine(Foo(A), B)"))
|
||||
self.assertTrue(any(i.code == "unknown_operator" for i in issues))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 依赖分析
|
||||
# ---------------------------------------------------------------------------
|
||||
class ResolveInputsTest(unittest.TestCase):
|
||||
def test_single_tag(self):
|
||||
self.assertEqual(resolve_inputs(parse("炉压")), ["炉压"])
|
||||
|
||||
def test_dedup_order(self):
|
||||
# 同一点位重复出现,去重且保持首次出现顺序
|
||||
self.assertEqual(resolve_inputs(parse("Combine(A, A)")), ["A"])
|
||||
|
||||
def test_multiple_tags(self):
|
||||
self.assertEqual(resolve_inputs(parse("Combine(A.tank1, A.tank2)")), ["A.tank1", "A.tank2"])
|
||||
|
||||
def test_op_collects_input(self):
|
||||
self.assertEqual(resolve_inputs(parse("EMA(CLF-01.TEMP, 5m)")), ["CLF-01.TEMP"])
|
||||
|
||||
def test_number_window_no_inputs(self):
|
||||
# 裸数值/窗口虽不是合法特征根,但 resolve_inputs 不报错
|
||||
self.assertEqual(resolve_inputs(Number(3)), [])
|
||||
self.assertEqual(resolve_inputs(Window(5.0, "m")), [])
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 执行
|
||||
# ---------------------------------------------------------------------------
|
||||
class MaterializeTest(unittest.TestCase):
|
||||
def setUp(self):
|
||||
# 一个稳定的伪时序:1..10
|
||||
self.series = {"A": [float(i) for i in range(1, 11)]} # 1..10
|
||||
|
||||
def test_bare_tag(self):
|
||||
self.assertEqual(materialize(parse("A"), self.series), self.series["A"])
|
||||
|
||||
def test_number(self):
|
||||
self.assertEqual(materialize(Number(3), {}), 3)
|
||||
|
||||
def test_sma_window3(self):
|
||||
out = materialize(parse("SMA(A, 3)"), self.series)
|
||||
# 前 2 个 NaN,第 3 个 = (1+2+3)/3 = 2.0
|
||||
self.assertTrue(math.isnan(out[0]) and math.isnan(out[1]))
|
||||
self.assertAlmostEqual(out[2], 2.0)
|
||||
self.assertAlmostEqual(out[9], (8 + 9 + 10) / 3)
|
||||
|
||||
def test_ema_decreasing_weight(self):
|
||||
out = materialize(parse("EMA(A, 5)"), self.series)
|
||||
# EMA 单调(输入单调增),首值 = 首个观测
|
||||
self.assertAlmostEqual(out[0], 1.0)
|
||||
self.assertTrue(all(out[i] <= out[i + 1] for i in range(len(out) - 1)))
|
||||
|
||||
def test_rolling_std(self):
|
||||
out = materialize(parse("RollingStd(A, 2)"), self.series)
|
||||
self.assertTrue(math.isnan(out[0]))
|
||||
# std(1,2) 无偏 = 0.7071...
|
||||
self.assertAlmostEqual(out[1], math.sqrt(0.5))
|
||||
|
||||
def test_rolling_max_min(self):
|
||||
mx = materialize(parse("RollingMax(A, 3)"), self.series)
|
||||
mn = materialize(parse("RollingMin(A, 3)"), self.series)
|
||||
self.assertEqual(mx[2], 3.0)
|
||||
self.assertEqual(mn[2], 1.0)
|
||||
|
||||
def test_diff(self):
|
||||
out = materialize(parse("Diff(A)"), self.series)
|
||||
self.assertTrue(math.isnan(out[0]))
|
||||
self.assertTrue(all(out[i] == 1.0 for i in range(1, len(out))))
|
||||
|
||||
def test_lag(self):
|
||||
out = materialize(parse("Lag(A, 2)"), self.series)
|
||||
self.assertTrue(math.isnan(out[0]) and math.isnan(out[1]))
|
||||
self.assertEqual(out[2], 1.0)
|
||||
|
||||
def test_rate_of_change(self):
|
||||
# 常数序列 → 变化率为 0(非 NaN;NaN 仅出现在前 window 步预热)
|
||||
const = {"C": [5.0] * 6}
|
||||
out = materialize(parse("RateOfChange(C)"), const)
|
||||
self.assertTrue(math.isnan(out[0])) # 预热步 NaN
|
||||
self.assertEqual(out[1], 0.0)
|
||||
# 含 0 的序列 → 分母为 0 → NaN
|
||||
zero_denom = {"Z": [0.0, 1.0, 2.0]}
|
||||
outz = materialize(parse("RateOfChange(Z)"), zero_denom)
|
||||
self.assertTrue(math.isnan(outz[1]))
|
||||
# 线性序列 ROC 步长1 = 1/prev
|
||||
out2 = materialize(parse("RateOfChange(A)"), self.series)
|
||||
self.assertAlmostEqual(out2[1], 1.0 / 1.0)
|
||||
self.assertAlmostEqual(out2[5], 1.0 / 5.0)
|
||||
|
||||
def test_log_negative_nan(self):
|
||||
data = {"P": [1.0, -2.0, math.e]}
|
||||
out = materialize(parse("Log(P)"), data)
|
||||
self.assertAlmostEqual(out[0], 0.0)
|
||||
self.assertTrue(math.isnan(out[1]))
|
||||
self.assertAlmostEqual(out[2], 1.0)
|
||||
|
||||
def test_scale(self):
|
||||
out = materialize(parse("Scale(A, 10)"), self.series)
|
||||
self.assertEqual(out[0], 10.0)
|
||||
self.assertEqual(out[9], 100.0)
|
||||
|
||||
def test_clip(self):
|
||||
out = materialize(parse("Clip(A, 3, 7)"), self.series)
|
||||
self.assertEqual(out, [3.0, 3.0, 3.0, 4.0, 5.0, 6.0, 7.0, 7.0, 7.0, 7.0])
|
||||
|
||||
def test_combine(self):
|
||||
data = {"A": [1.0, 2.0, 3.0], "B": [10.0, 20.0, 30.0]}
|
||||
self.assertEqual(materialize(parse("Combine(A, B)"), data), [11.0, 22.0, 33.0])
|
||||
|
||||
def test_missing_input_fails_fast(self):
|
||||
with self.assertRaises(KeyError):
|
||||
materialize(parse("EMA(Missing, 3)"), {"A": [1.0, 2.0, 3.0]})
|
||||
|
||||
def test_unknown_op_fails_fast(self):
|
||||
with self.assertRaises(ValueError):
|
||||
materialize(OpCall("NoSuchOp", (TagRef("A"),)), self.series)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 插件注册
|
||||
# ---------------------------------------------------------------------------
|
||||
class RegisterOperatorTest(unittest.TestCase):
|
||||
def test_register_then_parse_and_run(self):
|
||||
def _double(series_map, args):
|
||||
tag = args[0]
|
||||
return [x * 2 for x in series_map[tag.name]]
|
||||
|
||||
register_operator(
|
||||
"Double",
|
||||
min_arity=1,
|
||||
max_arity=1,
|
||||
arg_kinds=(("tag",),),
|
||||
func=_double,
|
||||
doc="示例自定义算子:翻倍",
|
||||
)
|
||||
try:
|
||||
self.assertIn("Double", OPERATORS)
|
||||
self.assertEqual(validate(parse("Double(A)")), [])
|
||||
self.assertEqual(
|
||||
materialize(parse("Double(A)"), {"A": [1.0, 2.0]}), [2.0, 4.0]
|
||||
)
|
||||
finally:
|
||||
OPERATORS.pop("Double", None)
|
||||
|
||||
def test_register_overrides(self):
|
||||
register_operator(
|
||||
"Stub", min_arity=0, max_arity=0, arg_kinds=(), func=lambda s, a: 1, doc="v1"
|
||||
)
|
||||
register_operator(
|
||||
"Stub", min_arity=0, max_arity=0, arg_kinds=(), func=lambda s, a: 2, doc="v2"
|
||||
)
|
||||
try:
|
||||
self.assertEqual(OPERATORS["Stub"].doc, "v2")
|
||||
finally:
|
||||
OPERATORS.pop("Stub", None)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 描述 / 往返
|
||||
# ---------------------------------------------------------------------------
|
||||
class DescribeAndSerializeTest(unittest.TestCase):
|
||||
def test_describe_contains_inputs(self):
|
||||
d = describe(parse("EMA(CLF-01.TEMP, 5m)"))
|
||||
self.assertIn("CLF-01.TEMP", d)
|
||||
self.assertIn("EMA", d)
|
||||
|
||||
def test_to_dict_roundtrip_shape(self):
|
||||
ast = parse("RateOfChange(炉压)")
|
||||
d = ast.to_dict()
|
||||
self.assertEqual(d["kind"], "op")
|
||||
self.assertEqual(d["name"], "RateOfChange")
|
||||
self.assertEqual(d["args"][0], {"kind": "tag", "name": "炉压"})
|
||||
|
||||
def test_window_seconds(self):
|
||||
w = Window(5.0, "m")
|
||||
self.assertEqual(w.seconds, 300)
|
||||
self.assertEqual(Window(2.0, "h").seconds, 7200)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user