feat(factor): C1 因子相关性分析(横截面 Spearman 矩阵 + API)

- 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 通过
This commit is contained in:
Simon
2026-09-09 07:31:49 +08:00
parent 93e32f4e63
commit 0d05bfd187
5 changed files with 227 additions and 1 deletions
+42
View File
@@ -6,13 +6,20 @@ Phase 3 为同步执行(样本有限);Phase 4 将改为 Job + SSE 异步
from __future__ import annotations
from datetime import date
from fastapi import APIRouter, HTTPException
from pydantic import BaseModel, Field
from app.api.deps import ResearchServiceDep
from app.domain.entities.research import (
BacktestResult,
FactorCorrelationReport,
FactorSpec,
FactorTestReport,
ResearchSpec,
SelectionSpec,
UniverseSpec,
)
from app.quant.factors import FactorError
@@ -54,3 +61,38 @@ def last_backtest() -> BacktestResult:
if "default" not in _LAST_BACKTEST:
raise HTTPException(status_code=404, detail="尚无回测结果,请先 POST /api/backtests")
return _LAST_BACKTEST["default"]
class FactorCorrelationRequest(BaseModel):
"""因子相关性分析请求(v3 §12)。"""
universe: UniverseSpec = UniverseSpec()
factors: list[FactorSpec] = Field(min_length=2)
period: tuple[date, date]
price_adjustment: str = Field(default="none", pattern="^(none|qfq)$")
@router.post(
"/factor-correlations",
response_model=FactorCorrelationReport,
summary="多因子两两相关(横截面 Spearman)",
)
def run_factor_correlation(
req: FactorCorrelationRequest,
service: ResearchServiceDep,
) -> FactorCorrelationReport:
if req.period[0] >= req.period[1]:
raise HTTPException(status_code=400, detail="period 必须满足 start < end")
spec = ResearchSpec(
type="factor_test",
universe=req.universe,
price_adjustment=req.price_adjustment,
factors=req.factors,
selection=SelectionSpec(top_n=10),
rebalance="monthly",
period=req.period,
)
try:
return service.run_factor_correlation(spec)
except (ValueError, FactorError) as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
+16
View File
@@ -232,6 +232,22 @@ class FactorTestReport(BaseModel):
config_snapshot: dict = Field(default_factory=dict)
# ---------- 因子相关性 / 暴露分析(C1,v3 §12) ----------
class FactorCorrelationReport(BaseModel):
"""多因子两两相关(横截面相关逐日均值;v3 §12 冗余剔除前置)。"""
factors: list[str]
corr_matrix: dict[str, dict[str, float]] = Field(
default_factory=dict, description="{f1: {f2: spearman 相关系数}}(对角线=1)"
)
sample_days: int = 0
sample_min_symbols: int = 0
unimplemented: list[str] = Field(default_factory=list)
config_snapshot: dict = Field(default_factory=dict)
# ---------- 异步 Job 与 Experiment(Phase 4) ----------
+55 -1
View File
@@ -14,7 +14,7 @@ import math
import pandas as pd
from app.domain.entities.research import FactorTestReport, QuantileReturn
from app.domain.entities.research import FactorCorrelationReport, FactorTestReport, QuantileReturn
MIN_CROSS_SECTION = 5 # 少于该样本数的日期跳过(避免噪声 IC)
@@ -70,6 +70,60 @@ def quantile_returns(factor: pd.DataFrame, forward: pd.DataFrame, quantiles: int
return pd.Series(means)
def _cross_section_corr(pa: pd.Series, pb: pd.Series, min_n: int) -> float | None:
"""两股票列在某一日期截面值的 Spearman 相关(样本不足 → None)。"""
df = pd.concat([pa, pb], axis=1).dropna()
if len(df) < min_n or df.iloc[:, 0].nunique() < 2:
return None
return df.iloc[:, 0].corr(df.iloc[:, 1], method="spearman")
def factor_correlation_report(
panels: dict[str, pd.DataFrame],
*,
min_symbols: int = MIN_CROSS_SECTION,
max_days: int = 10000,
) -> FactorCorrelationReport:
"""两两因子相关:共同日期 ∩ 后,逐日横截面 Spearman 相关取均值。
用于冗余剔除与因子池管理(v3 §12 流程:Correlation → Redundancy Removal)。
"""
names = list(panels)
common = None
for panel in panels.values():
idx = set(panel.index)
common = idx if common is None else (common & idx)
if not common:
empty = {a: {b: (1.0 if a == b else 0.0) for b in names} for a in names}
return FactorCorrelationReport(
factors=names, corr_matrix=empty,
sample_days=0, sample_min_symbols=min_symbols,
)
days = sorted(common)[:max_days]
corr_matrix: dict[str, dict[str, float]] = {}
for a in names:
corr_matrix[a] = {}
for b in names:
if a == b:
corr_matrix[a][b] = 1.0
continue
acc: list[float] = []
for day in days:
if day not in panels[a].index or day not in panels[b].index:
continue
c = _cross_section_corr(panels[a].loc[day], panels[b].loc[day], min_symbols)
if c is not None:
acc.append(c)
corr_matrix[a][b] = round(float(pd.Series(acc).mean()), 4) if acc else 0.0
return FactorCorrelationReport(
factors=names,
corr_matrix=corr_matrix,
sample_days=len(days),
sample_min_symbols=min_symbols,
config_snapshot={"min_symbols": min_symbols},
)
def run_factor_test(
factor: pd.DataFrame,
forward: pd.DataFrame,
+9
View File
@@ -17,6 +17,7 @@ import pandas as pd
from app.domain.entities.research import (
BacktestResult,
FactorCorrelationReport,
FactorTestReport,
ResearchSpec,
)
@@ -24,7 +25,9 @@ from app.domain.repositories.market import (
DailyBarRepository,
StockRepository,
)
from app.quant.composite import build_factor_panels
from app.quant.engine import QuantEngine
from app.quant.evaluation import factor_correlation_report
from app.quant.universe import filter_stocks, resolve_members # noqa: F401 —— 范围过滤
# 流式路径每攒多少行落一个 DataFrame 分片(控制 concat 峰值)
@@ -129,6 +132,12 @@ class ResearchService:
daily = self._load_daily(spec)
return self._engine.run_backtest(daily, spec)
def run_factor_correlation(self, spec: ResearchSpec) -> FactorCorrelationReport:
"""多因子两两相关(v3 §12):同 universe/period 装配 → 横截面相关矩阵。"""
daily = self._load_daily(spec)
panels = {fs.name: build_factor_panels(daily, [fs])[0][1] for fs in spec.factors}
return factor_correlation_report(panels)
# ---- 数据装配 ----
def _load_daily(self, spec: ResearchSpec) -> pd.DataFrame:
+105
View File
@@ -0,0 +1,105 @@
"""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"]