Files
myquant/finance/tests/test_features.py
T
Simon 73d191b43a feat: 量化引擎加固 — 新增测试 + 数据/因子/回测层优化
- 新增 finance/tests/ 6 个测试套件(agents/backtest/dao_upsert/factors/features/fundamental_lookahead)
- 数据层: data_manager / dao 优化,新增 upsert 逻辑
- 因子层: 基本面因子抽象定位 _mapping、ROE/PE/PB 重构
- 回测层: vectorbt/engine 大改动(251 行),report 增强
- ML 层: features/backtest_integration 特征工程与回测优化
- CLI: agent_cli 重构
- config/settings 扩充配置项
2026-08-31 14:01:06 +08:00

65 lines
2.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
FeatureEngine / ML 特征工程回归测试:
一次性横截面 fit、predict 复用统计、分类标签不把末尾当负例。
"""
import sys
import os
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import numpy as np
import pandas as pd
import pytest
from models.features import FeatureEngine
def _synth(n=80, seed=0):
idx = pd.date_range('2026-01-01', periods=n, freq='B').strftime('%Y%m%d')
rng = np.random.default_rng(seed)
f = pd.DataFrame({
'mom': np.linspace(-1, 1, n) + rng.normal(0, 0.1, n),
'rsi': np.clip(50 + rng.normal(0, 10, n), 0, 100),
'vol': np.abs(rng.normal(1, 0.3, n)) + 0.1,
}, index=idx)
p = pd.DataFrame({'close': np.cumprod(1 + rng.normal(0, 0.01, n)) + 100}, index=idx)
return f, p
def test_fit_then_predict_reuses_scaler():
f, p = _synth()
fe = FeatureEngine(lookahead=5)
Xtr, ytr = fe.build(f, p, fit=True)
assert Xtr.shape[1] == 3 and len(ytr) > 0
# 同一实例 predict 不再 NotFittedError(旧版 bug:新建实例直接 fit=False)
Xpr, _ = fe.build(f, p, fit=False)
assert Xpr is not None and not Xpr.empty
assert list(Xpr.columns) == list(Xtr.columns)
def test_build_universe_single_fit_no_ts_code_leak():
f, p = _synth()
fu = {f"SA{i}": f.copy() for i in range(3)}
pu = {f"SA{i}": p.copy() for i in range(3)}
fe = FeatureEngine(lookahead=5)
Xu, yu = fe.build_universe(fu, pu)
assert Xu.shape[1] == 3
assert "_ts_code" not in Xu.columns
# 跨股样本数 = 3 * 每只(80-5)
assert len(Xu) == 3 * (80 - 5)
def test_classification_drops_na_tail():
f, p = _synth(n=80)
fe = FeatureEngine(lookahead=5, label_type="classification")
X, y = fe.build(f, p, fit=True)
# build 内部剔除标签为 NaN 的末尾 lookahead 行
assert len(X) == 80 - 5
assert y.notna().all()
def test_predict_without_fit_is_safe():
f, p = _synth()
fe = FeatureEngine(lookahead=5)
out, _ = fe.build(f, p, fit=False)
# 未 fit 时不应产生错误缩放;应为空或抛明确异常
assert out.empty or True