From 9c730e4b69d798c549b29fead2bdc08b99a759f9 Mon Sep 17 00:00:00 2001 From: Lytem Date: Wed, 8 Jul 2026 12:11:03 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E4=BF=AE=E5=A4=8D=E5=9B=A0=E5=AD=90?= =?UTF-8?q?=E5=9B=9E=E6=B5=8B=E5=88=86=E5=B1=82=E9=94=99=E8=AF=AF=E9=97=AE?= =?UTF-8?q?=E9=A2=98=20(#63)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: force LF for uv.lock * Fix factor backtest grouping stability --- .gitattributes | 1 + backend/app/backtest/factor.py | 60 ++++++++++++++----- .../backtest/charts/FactorGroupNavChart.tsx | 12 +++- 3 files changed, 56 insertions(+), 17 deletions(-) create mode 100644 .gitattributes diff --git a/.gitattributes b/.gitattributes new file mode 100644 index 0000000..3309f39 --- /dev/null +++ b/.gitattributes @@ -0,0 +1 @@ +backend/uv.lock text eol=lf \ No newline at end of file diff --git a/backend/app/backtest/factor.py b/backend/app/backtest/factor.py index 740582a..03f70dd 100644 --- a/backend/app/backtest/factor.py +++ b/backend/app/backtest/factor.py @@ -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,14 +312,37 @@ 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") - .alias("_group") + """截面序号分桶,避免 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 + # ── 分组净值 ── @staticmethod @@ -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 [], {} diff --git a/frontend/src/pages/backtest/charts/FactorGroupNavChart.tsx b/frontend/src/pages/backtest/charts/FactorGroupNavChart.tsx index 5bae619..21e7614 100644 --- a/frontend/src/pages/backtest/charts/FactorGroupNavChart.tsx +++ b/frontend/src/pages/backtest/charts/FactorGroupNavChart.tsx @@ -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[] { + 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 (