fix: 修复因子回测分层错误问题 (#63)

* fix: force LF for uv.lock

* Fix factor backtest grouping stability
This commit is contained in:
Lytem
2026-07-08 12:11:03 +08:00
committed by GitHub
parent 3a5f3e05b6
commit 9c730e4b69
3 changed files with 56 additions and 17 deletions
+1
View File
@@ -0,0 +1 @@
backend/uv.lock text eol=lf
+44 -14
View File
@@ -140,6 +140,8 @@ class FactorBacktestService:
if panel.is_empty(): if panel.is_empty():
return _err("过滤后无有效数据") return _err("过滤后无有效数据")
panel = panel.sort(["symbol", "date"])
n_symbols = panel["symbol"].n_unique() n_symbols = panel["symbol"].n_unique()
n_dates = panel["date"].n_unique() n_dates = panel["date"].n_unique()
@@ -218,8 +220,8 @@ class FactorBacktestService:
.group_by("date") .group_by("date")
.agg( .agg(
pl.corr( pl.corr(
pl.col(factor_col).rank(method="random"), pl.col(factor_col).rank(method="average"),
pl.col("_next_return").rank(method="random"), pl.col("_next_return").rank(method="average"),
).alias("ic") ).alias("ic")
) )
.sort("date") .sort("date")
@@ -231,8 +233,8 @@ class FactorBacktestService:
def _calc_period_return(panel: pl.DataFrame, rebalance: str) -> pl.DataFrame: def _calc_period_return(panel: pl.DataFrame, rebalance: str) -> pl.DataFrame:
"""计算到下个调仓日的收益。 """计算到下个调仓日的收益。
weekly: 下周一的 open / 今日 close - 1 weekly: 下个周调仓日 close / 今日 close - 1
monthly: 下月首个交易日的 open / 今日 close - 1 monthly: 下个月调仓日 close / 今日 close - 1
只在调仓日标记行有效,其他行为 null。 只在调仓日标记行有效,其他行为 null。
""" """
import datetime as _dt import datetime as _dt
@@ -273,7 +275,6 @@ class FactorBacktestService:
# 最后一个调仓日没有下一个,不计算收益 # 最后一个调仓日没有下一个,不计算收益
# 构建 (date, symbol) → next_rebalance_date 的 close 价格映射 # 构建 (date, symbol) → next_rebalance_date 的 close 价格映射
# 简化: 用下个调仓日的 close / 当前 close
panel = panel.sort(["symbol", "date"]) panel = panel.sort(["symbol", "date"])
dates_col = panel["date"].to_list() dates_col = panel["date"].to_list()
close_col = panel["close"].to_list() close_col = panel["close"].to_list()
@@ -311,13 +312,36 @@ class FactorBacktestService:
@staticmethod @staticmethod
def _add_groups(panel: pl.DataFrame, factor_col: str, n_groups: int) -> pl.DataFrame: def _add_groups(panel: pl.DataFrame, factor_col: str, n_groups: int) -> pl.DataFrame:
"""截面分位数分组""" """截面序号分桶,避免 qcut 在重复因子值截面上抛错"""
return panel.with_columns( return (
pl.col(factor_col) panel.sort(["date", factor_col, "symbol"])
.qcut(n_groups, labels=[f"Q{i+1}" for i in range(n_groups)]) .with_columns(
.over("date") (pl.cum_count("symbol").over("date") - 1).alias("_factor_ord"),
pl.len().over("date").alias("_factor_count"),
)
.with_columns(
(
((pl.col("_factor_ord") * n_groups) / pl.col("_factor_count"))
.floor()
.cast(pl.Int64)
+ 1
)
.clip(1, n_groups)
.cast(pl.Utf8)
.map_elements(lambda v: f"Q{v}", return_dtype=pl.Utf8)
.alias("_group") .alias("_group")
) )
.drop(["_factor_ord", "_factor_count"])
)
@staticmethod
def _group_sort_key(group: str) -> int:
if group.startswith("Q"):
try:
return int(group[1:])
except ValueError:
pass
return 0
# ── 分组净值 ── # ── 分组净值 ──
@@ -337,7 +361,7 @@ class FactorBacktestService:
if pivot.is_empty(): if pivot.is_empty():
return [] return []
group_cols = [c for c in pivot.columns if c != "date"] group_cols = sorted([c for c in pivot.columns if c != "date"], key=FactorBacktestService._group_sort_key)
# 累乘净值曲线 # 累乘净值曲线
result: list[dict] = [] result: list[dict] = []
@@ -362,12 +386,15 @@ class FactorBacktestService:
if not group_nav: if not group_nav:
return [] return []
group_cols = [k for k in group_nav[0] if k != "date"] group_cols = sorted(
[k for k in group_nav[0] if k != "date"],
key=FactorBacktestService._group_sort_key,
)
n_days = max((end - start).days, 1) n_days = max((end - start).days, 1)
years = n_days / 365.25 years = n_days / 365.25
stats = [] stats = []
for i, c in enumerate(sorted(group_cols)): for i, c in enumerate(group_cols):
values = [r[c] for r in group_nav if r.get(c) is not None] values = [r[c] for r in group_nav if r.get(c) is not None]
if not values: if not values:
continue continue
@@ -419,7 +446,10 @@ class FactorBacktestService:
if not group_nav: if not group_nav:
return [], {} return [], {}
group_cols = sorted([k for k in group_nav[0] if k != "date"]) group_cols = sorted(
[k for k in group_nav[0] if k != "date"],
key=FactorBacktestService._group_sort_key,
)
if len(group_cols) < 2: if len(group_cols) < 2:
return [], {} return [], {}
@@ -20,13 +20,21 @@ interface Props {
result: FactorBacktestResult result: FactorBacktestResult
} }
function groupSortKey(group: string): number {
return group.startsWith('Q') ? Number(group.slice(1)) || 0 : 0
}
function getGroupCols(row: Record<string, any>): string[] {
return Object.keys(row).filter(k => k !== 'date').sort((a, b) => groupSortKey(a) - groupSortKey(b))
}
export function FactorGroupNavChart({ result }: Props) { export function FactorGroupNavChart({ result }: Props) {
const ct = useChartTheme() const ct = useChartTheme()
const option = useMemo(() => { const option = useMemo(() => {
if (!result.group_nav.length) return null if (!result.group_nav.length) return null
const dates = result.group_nav.map(r => (r.date as string).slice(0, 10)) const dates = result.group_nav.map(r => (r.date as string).slice(0, 10))
const groupCols = Object.keys(result.group_nav[0]).filter(k => k !== 'date').sort() const groupCols = getGroupCols(result.group_nav[0])
// 多空净值 // 多空净值
const lsNav = result.long_short_nav const lsNav = result.long_short_nav
@@ -101,7 +109,7 @@ export function FactorGroupNavChart({ result }: Props) {
// 图例 // 图例
const groupCols = result.group_nav.length > 0 const groupCols = result.group_nav.length > 0
? Object.keys(result.group_nav[0]).filter(k => k !== 'date').sort() ? getGroupCols(result.group_nav[0])
: [] : []
return ( return (