- evaluation.factor_correlation_report:多因子共同日期 ∩ 后逐日横截面 Spearman 相关 取均值 → FactorCorrelationReport(冗余剔除前置,v3 §12 Correlation→Redundancy) - ResearchService.run_factor_correlation + POST /api/factor-correlations (universe/factors/period;与其它研究同装配口径) - tests/test_factor_correlation.py:矩阵对角=1/近线性±相关符号/对称/无共同日期补零、 API 冒烟;全量 pytest 通过
106 lines
3.9 KiB
Python
106 lines
3.9 KiB
Python
"""C1 因子相关性分析测试:横截面 Spearman 相关矩阵(对角线=1、正/负相关符号)。"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from datetime import date
|
|
|
|
import numpy as np
|
|
import pandas as pd
|
|
import pytest
|
|
from app.api import deps
|
|
from app.infrastructure.persistence.sqlalchemy.base import Base
|
|
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
|
SqlAlchemyDailyBarRepository,
|
|
SqlAlchemyStockRepository,
|
|
)
|
|
from app.main import app
|
|
from app.quant.evaluation import factor_correlation_report
|
|
from fastapi.testclient import TestClient
|
|
from sqlalchemy import create_engine
|
|
from sqlalchemy.orm import sessionmaker
|
|
|
|
from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily
|
|
|
|
_SYMS = ["60000" + str(i) + ".SH" for i in range(5)]
|
|
|
|
|
|
def _panel(dates, symbols, drift) -> pd.DataFrame:
|
|
rng = np.random.default_rng(7)
|
|
arr = np.zeros((len(dates), len(symbols)))
|
|
for j in range(len(symbols)):
|
|
base = np.cumsum(rng.normal(0, 0.5, len(dates)))
|
|
arr[:, j] = base + drift * np.arange(len(dates)) / 100
|
|
return pd.DataFrame(arr, index=dates, columns=symbols)
|
|
|
|
|
|
class TestCorrelationReportUnit:
|
|
def test_diag_and_sign(self) -> None:
|
|
dates = pd.bdate_range("2024-01-01", periods=120)
|
|
pa = _panel(dates, _SYMS, 2.0)
|
|
pb = pa * 0.98 + 0.2 # 近线性正相关
|
|
pc = -pa # 完全负相关
|
|
rep = factor_correlation_report({"a": pa, "b": pb, "c": pc}, min_symbols=3)
|
|
assert rep.factors == ["a", "b", "c"]
|
|
assert rep.corr_matrix["a"]["a"] == 1.0
|
|
assert rep.corr_matrix["a"]["b"] > 0.9
|
|
assert rep.corr_matrix["a"]["c"] < -0.9
|
|
# 对称
|
|
assert abs(rep.corr_matrix["a"]["b"] - rep.corr_matrix["b"]["a"]) < 1e-3
|
|
assert rep.sample_days == 120
|
|
|
|
def test_no_common_dates(self) -> None:
|
|
d1 = pd.bdate_range("2024-01-01", periods=10)
|
|
d2 = pd.bdate_range("2023-01-01", periods=10)
|
|
rep = factor_correlation_report(
|
|
{"x": _panel(d1, _SYMS, 0), "y": _panel(d2, _SYMS, 0)}
|
|
)
|
|
assert rep.sample_days == 0 and rep.corr_matrix["x"]["y"] == 0.0
|
|
|
|
|
|
class TestCorrelationApi:
|
|
@pytest.fixture()
|
|
def client(self, tmp_path):
|
|
engine = create_engine(f"sqlite:///{tmp_path / 'corr.db'}", future=True)
|
|
Base.metadata.create_all(engine)
|
|
Session = sessionmaker(bind=engine, expire_on_commit=False)
|
|
df = synthetic_daily({s: 0.004 - 0.001 * i for i, s in enumerate(_SYMS)}, n=320)
|
|
with Session() as session:
|
|
SqlAlchemyStockRepository(session).upsert_many(
|
|
[
|
|
__import__("app.domain.entities.market", fromlist=["Stock"]).Stock(
|
|
symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1)
|
|
)
|
|
for i, s in enumerate(_SYMS)
|
|
]
|
|
)
|
|
SqlAlchemyDailyBarRepository(session).upsert_many(bars_dataframe_to_daily_bars(df))
|
|
session.commit()
|
|
|
|
def _override():
|
|
with Session() as s:
|
|
yield s
|
|
|
|
app.dependency_overrides[deps.get_session] = _override
|
|
with TestClient(app) as c:
|
|
yield c
|
|
app.dependency_overrides.clear()
|
|
|
|
def test_factor_correlations(self, client) -> None:
|
|
resp = client.post(
|
|
"/api/factor-correlations",
|
|
json={
|
|
"universe": {"exclude_st": False, "min_listing_days": 0},
|
|
"factors": [
|
|
{"name": "momentum_20", "weight": 1},
|
|
{"name": "momentum_60", "weight": 1},
|
|
{"name": "volatility_60", "weight": 1},
|
|
],
|
|
"period": ["2024-03-01", "2024-12-31"],
|
|
},
|
|
)
|
|
assert resp.status_code == 200
|
|
body = resp.json()
|
|
assert len(body["factors"]) == 3
|
|
assert body["corr_matrix"]["momentum_20"]["momentum_20"] == 1.0
|
|
assert "momentum_60" in body["corr_matrix"]
|