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():
return _err("过滤后无有效数据")
panel = panel.sort(["symbol", "date"])
n_symbols = panel["symbol"].n_unique()
n_dates = panel["date"].n_unique()
@@ -218,8 +220,8 @@ class FactorBacktestService:
.group_by("date")
.agg(
pl.corr(
pl.col(factor_col).rank(method="random"),
pl.col("_next_return").rank(method="random"),
pl.col(factor_col).rank(method="average"),
pl.col("_next_return").rank(method="average"),
).alias("ic")
)
.sort("date")
@@ -231,8 +233,8 @@ class FactorBacktestService:
def _calc_period_return(panel: pl.DataFrame, rebalance: str) -> pl.DataFrame:
"""计算到下个调仓日的收益。
weekly: 下周一的 open / 今日 close - 1
monthly: 下月首个交易日的 open / 今日 close - 1
weekly: 下个周调仓日 close / 今日 close - 1
monthly: 下个月调仓日 close / 今日 close - 1
只在调仓日标记行有效,其他行为 null。
"""
import datetime as _dt
@@ -273,7 +275,6 @@ class FactorBacktestService:
# 最后一个调仓日没有下一个,不计算收益
# 构建 (date, symbol) → next_rebalance_date 的 close 价格映射
# 简化: 用下个调仓日的 close / 当前 close
panel = panel.sort(["symbol", "date"])
dates_col = panel["date"].to_list()
close_col = panel["close"].to_list()
@@ -311,13 +312,36 @@ class FactorBacktestService:
@staticmethod
def _add_groups(panel: pl.DataFrame, factor_col: str, n_groups: int) -> pl.DataFrame:
"""截面分位数分组"""
return panel.with_columns(
pl.col(factor_col)
.qcut(n_groups, labels=[f"Q{i+1}" for i in range(n_groups)])
.over("date")
"""截面序号分桶,避免 qcut 在重复因子值截面上抛错"""
return (
panel.sort(["date", factor_col, "symbol"])
.with_columns(
(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")
)
.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():
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] = []
@@ -362,12 +386,15 @@ class FactorBacktestService:
if not group_nav:
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)
years = n_days / 365.25
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]
if not values:
continue
@@ -419,7 +446,10 @@ class FactorBacktestService:
if not group_nav:
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:
return [], {}
@@ -20,13 +20,21 @@ interface Props {
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) {
const ct = useChartTheme()
const option = useMemo(() => {
if (!result.group_nav.length) return null
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
@@ -101,7 +109,7 @@ export function FactorGroupNavChart({ result }: Props) {
// 图例
const groupCols = result.group_nav.length > 0
? Object.keys(result.group_nav[0]).filter(k => k !== 'date').sort()
? getGroupCols(result.group_nav[0])
: []
return (