"""
均衡FOF策略 - 最优版本 (V3 Final)
============================================================
基于方正金工《如何构建均衡的FOF组合？》研报框架，
在LOF基金中实现并增强的FOF策略。

核心改进（vs原始复现V2）：
1. 因子增强: 增加动量因子(63日窗口,剔除近5日短期反转)、
   波动调整动量、LOF换手率因子、重仓股动量因子
2. 因子权重优化: 基础因子40% + 动量因子40% + 附加因子20%

关键参数：
- 动量窗口: 63个交易日
- 短期反转剔除: 5个交易日
- 风格分类: 两步独立RBSA (大/小盘 × 成长/价值)
- 每风格选Top5, 共20只基金等权
- 调仓频率: 季度 (2/5/8/11月首个交易日)
- 风控: 无主动风控叠加(实证显示MA/VolTarget均损夏普)

回测绩效 (2019-2025):
- 年化收益: 17.49%
- 夏普比率: 1.15
- 最大回撤: 24.64%
- 卡玛比率: 0.71
- 月度胜率: 59.5%

运行: python3 fof_strategy_v3_final.py
============================================================
"""
# %%
import numpy as np
import pandas as pd
import dai
import bigtrader
from datetime import timedelta
import warnings
warnings.filterwarnings('ignore')

# ============================================================
# 策略参数
# ============================================================
START_DATE = '2017-01-01'
END_DATE = '2026-07-20'
LOOKBACK_DAYS = 252           # 因子回望期（约1年交易日）
N_PER_STYLE = 5              # 每风格选基数量
REBALANCE_MONTHS = [2, 5, 8, 11]  # 调仓月份
MOM_WINDOW = 63              # 动量窗口（3个月）
SKIP_RECENT = 5              # 剔除近5日（短期反转效应）

# 风格代理指数
IDX_LARGE = '399314.SZ'      # 巨潮大盘
IDX_MID = '399315.SZ'        # 巨潮中盘
IDX_SMALL = '399316.SZ'      # 巨潮小盘
IDX_VALUE = '000016.SH'      # 上证50（价值代理）
IDX_GROWTH = '399006.SZ'     # 创业板指（成长代理）
IDX_HS300 = '000300.SH'      # 沪深300（基准）
ALL_INDICES = [IDX_LARGE, IDX_MID, IDX_SMALL, IDX_VALUE, IDX_GROWTH, IDX_HS300]

# 费用
REDEMPTION_FEE = 0.005       # 赎回费0.5%
SUBSCRIPTION_FEE = 0.0       # 申购费0%

print("=" * 70)
print("  均衡FOF策略 V3 Final - BigTrader回测")
print("=" * 70)
print(f"  回测区间: {START_DATE} ~ {END_DATE}")
print(f"  动量窗口: {MOM_WINDOW}日, 剔除近{SKIP_RECENT}日")
print(f"  每风格: Top{N_PER_STYLE}, 季度调仓")
print("=" * 70)

# ============================================================
# 数据获取
# ============================================================

def get_trading_dates(start, end):
    """获取交易日历"""
    sql = "SELECT DISTINCT date FROM cn_stock_index_bar1d WHERE instrument='000300.SH' ORDER BY date"
    df = dai.query(sql, filters={'date': [start, end]}).df()
    return sorted(pd.to_datetime(df['date']).tolist())


def get_lof_universe(date_str):
    """获取主动管理型LOF基金池（排除指数/债/货币/QDII/FOF等）"""
    sql = """
    SELECT DISTINCT instrument, name FROM cn_fund_instruments
    WHERE instrument LIKE '1%'
       AND name NOT LIKE '%ETF%' AND name NOT LIKE '%指数%'
       AND name NOT LIKE '%分级%' AND name NOT LIKE '%债%'
       AND name NOT LIKE '%货币%' AND name NOT LIKE '%QDII%'
       AND name NOT LIKE '%FOF%' AND name NOT LIKE '%定开%'
       AND name NOT LIKE '%定期%' AND name NOT LIKE '%封闭%'
       AND name NOT LIKE '%港股通%' AND name NOT LIKE '%联接%'
       AND name NOT LIKE '%养老%' AND name NOT LIKE '%保本%'
    ORDER BY instrument
    """
    try:
        df = dai.query(sql, filters={'date': [date_str, date_str]}).df()
        df = df[~df['instrument'].str.startswith('159')]  # 排除场内ETF
    except Exception:
        df = pd.DataFrame()
    return df


def get_fund_nav(instruments, start_date, end_date):
    """批量获取基金累计净值"""
    if not instruments:
        return pd.DataFrame()
    all_dfs = []
    for i in range(0, len(instruments), 50):
        batch = instruments[i:i+50]
        instr_str = "', '".join(batch)
        sql = f"SELECT date, instrument, accum_nav as nav FROM cn_fund_nav WHERE instrument IN ('{instr_str}') ORDER BY date"
        try:
            df = dai.query(sql, filters={'date': [start_date, end_date]}).df()
            if not df.empty:
                all_dfs.append(df)
        except Exception:
            pass
    return pd.concat(all_dfs, ignore_index=True) if all_dfs else pd.DataFrame()


def get_index_data(start_date, end_date):
    """获取风格代理指数日收益率"""
    instr_str = "', '".join(ALL_INDICES)
    sql = f"SELECT date, instrument, close FROM cn_stock_index_bar1d WHERE instrument IN ('{instr_str}') ORDER BY date"
    df = dai.query(sql, filters={'date': [start_date, end_date]}).df()
    if df.empty:
        return pd.DataFrame()
    pivot = df.pivot_table(index='date', columns='instrument', values='close').sort_index()
    return pivot.pct_change().dropna()


def get_fund_turnover(instruments, start_date, end_date):
    """获取LOF场内换手率"""
    if not instruments:
        return pd.DataFrame()
    all_dfs = []
    for i in range(0, len(instruments), 50):
        batch = instruments[i:i+50]
        instr_str = "', '".join(batch)
        sql = f"SELECT date, instrument, turn FROM cn_fund_bar1d WHERE instrument IN ('{instr_str}') ORDER BY date"
        try:
            df = dai.query(sql, filters={'date': [start_date, end_date]}).df()
            if not df.empty:
                all_dfs.append(df)
        except Exception:
            pass
    return pd.concat(all_dfs, ignore_index=True) if all_dfs else pd.DataFrame()


def get_fund_holdings(instruments, date_str):
    """获取最近一期前十大重仓股"""
    if not instruments:
        return pd.DataFrame()
    start_d = (pd.Timestamp(date_str) - timedelta(days=120)).strftime('%Y-%m-%d')
    all_dfs = []
    for i in range(0, len(instruments), 30):
        batch = instruments[i:i+30]
        instr_str = "', '".join(batch)
        sql = f"SELECT instrument, date, stock_instrument, holding_pct, rank FROM cn_fund_holding_detail WHERE instrument IN ('{instr_str}') ORDER BY instrument, rank"
        try:
            df = dai.query(sql, filters={'date': [start_d, date_str]}).df()
            if not df.empty:
                all_dfs.append(df)
        except Exception:
            pass
    if not all_dfs:
        return pd.DataFrame()
    result = pd.concat(all_dfs, ignore_index=True)
    result['date'] = pd.to_datetime(result['date'])
    latest = result.groupby('instrument')['date'].max().reset_index()
    latest.columns = ['instrument', 'latest_date']
    result = result.merge(latest, on='instrument')
    result = result[result['date'] == result['latest_date']]
    return result[['instrument', 'stock_instrument', 'holding_pct', 'rank']]


def get_stock_returns(stock_instruments, start_date, end_date):
    """获取股票日收益率"""
    if not stock_instruments:
        return pd.DataFrame()
    all_dfs = []
    for i in range(0, len(stock_instruments), 50):
        batch = stock_instruments[i:i+50]
        instr_str = "', '".join(batch)
        sql = f"SELECT date, instrument, close FROM cn_stock_bar1d WHERE instrument IN ('{instr_str}') ORDER BY date"
        try:
            df = dai.query(sql, filters={'date': [start_date, end_date]}).df()
            if not df.empty:
                all_dfs.append(df)
        except Exception:
            pass
    if not all_dfs:
        return pd.DataFrame()
    result = pd.concat(all_dfs, ignore_index=True)
    pivot = result.pivot_table(index='date', columns='instrument', values='close').sort_index()
    return pivot.pct_change().dropna()


# ============================================================
# 步骤1: 风格分类（两步独立RBSA）
# ============================================================

def classify_styles(nav_df, index_returns, min_days=60):
    """
    两步独立风格分类：
    Step1: 计算基金对大盘/小盘指数的beta差 → 判断大/中小盘
    Step2: 计算基金对价值/成长指数的beta差 → 判断价值/成长
    """
    if index_returns.empty:
        return {}
    needed = [IDX_LARGE, IDX_SMALL, IDX_VALUE, IDX_GROWTH]
    for col in needed:
        if col not in index_returns.columns:
            return {}

    fund_metrics = []
    for inst in nav_df['instrument'].unique():
        fdata = nav_df[nav_df['instrument'] == inst].sort_values('date').copy()
        if len(fdata) < min_days:
            continue
        fdata['ret'] = fdata['nav'].pct_change()
        fdata = fdata.dropna(subset=['ret'])
        fdata['date'] = pd.to_datetime(fdata['date'])
        merged = fdata[['date', 'ret']].merge(index_returns.reset_index(), on='date', how='inner')
        if len(merged) < min_days:
            continue
        y = merged['ret'].values
        y_dm = y - y.mean()

        def beta(col):
            x = merged[col].values
            x_dm = x - x.mean()
            d = (x_dm**2).sum()
            return (y_dm * x_dm).sum() / d if d > 0 else 0.0

        size_score = beta(IDX_LARGE) - beta(IDX_SMALL)
        style_score = beta(IDX_VALUE) - beta(IDX_GROWTH)
        fund_metrics.append({'instrument': inst, 'size_score': size_score, 'style_score': style_score})

    if not fund_metrics:
        return {}
    df = pd.DataFrame(fund_metrics)
    size_med = df['size_score'].median()
    df['is_large'] = df['size_score'] >= size_med

    fund_styles = {}
    for is_large, group in df.groupby('is_large'):
        style_med = group['style_score'].median()
        for _, row in group.iterrows():
            is_value = row['style_score'] >= style_med
            if is_large:
                style = '大盘价值' if is_value else '大盘成长'
            else:
                style = '中小盘价值' if is_value else '中小盘成长'
            fund_styles[row['instrument']] = style
    return fund_styles


# ============================================================
# 步骤2: 增强因子计算
# ============================================================

def compute_factors(nav_df, index_returns, fund_styles,
                    turnover_df=None, holdings_df=None, stock_ret_df=None):
    """
    综合因子体系：
    - Alpha: CAPM截距项年化
    - IR: Alpha / 跟踪误差
    - Sharpe: (年化收益-无风险) / 年化波动
    - Calmar: 年化收益 / 最大回撤
    - 动量: MOM_WINDOW日累计收益(剔除近SKIP_RECENT日)
    - 波动调整动量: 动量 / 区间波动率
    - LOF换手率: 场内平均换手率（流动性/关注度代理）
    - 重仓股动量: 前十大重仓股近63日加权收益
    """
    style_bench = {
        '大盘成长': IDX_LARGE, '大盘价值': IDX_LARGE,
        '中小盘成长': IDX_SMALL, '中小盘价值': IDX_SMALL,
    }
    results = []

    for inst, style in fund_styles.items():
        fdata = nav_df[nav_df['instrument'] == inst].sort_values('date').copy()
        if len(fdata) < 120:
            continue
        fdata['ret'] = fdata['nav'].pct_change()
        fdata = fdata.dropna(subset=['ret'])
        fdata['date'] = pd.to_datetime(fdata['date'])
        if len(fdata) < 100:
            continue

        rets = fdata['ret'].values
        n = len(rets)
        total = (1 + pd.Series(rets)).prod() - 1
        ann_ret = (1 + total) ** (252/n) - 1
        ann_vol = rets.std() * np.sqrt(252)
        if ann_vol == 0:
            continue
        sharpe = (ann_ret - 0.025) / ann_vol
        cum = (1 + pd.Series(rets)).cumprod()
        max_dd = abs(((cum - cum.cummax()) / cum.cummax()).min())
        calmar = ann_ret / max_dd if max_dd > 0 else 0

        # Alpha & IR
        bench_col = style_bench.get(style, IDX_LARGE)
        alpha, ir = np.nan, np.nan
        if bench_col in index_returns.columns:
            merged = fdata[['date', 'ret']].merge(
                index_returns[[bench_col]].reset_index(), on='date', how='inner')
            if len(merged) >= 60:
                X = np.column_stack([np.ones(len(merged)), merged[bench_col].values])
                y_reg = merged['ret'].values
                try:
                    params = np.linalg.lstsq(X, y_reg, rcond=None)[0]
                    alpha = params[0] * 252
                    te = (y_reg - X @ params).std() * np.sqrt(252)
                    ir = alpha / te if te > 0 else 0
                except:
                    pass

        # 动量因子（63日窗口，剔除近5日）
        if n >= MOM_WINDOW + SKIP_RECENT:
            mom_rets = rets[-(MOM_WINDOW + SKIP_RECENT):-SKIP_RECENT] if SKIP_RECENT > 0 else rets[-MOM_WINDOW:]
            mom = (1 + pd.Series(mom_rets)).prod() - 1
            vol_w = mom_rets.std() * np.sqrt(252)
            adj_mom = mom / vol_w if vol_w > 0 else 0
        else:
            mom, adj_mom = 0, 0

        row = {
            'instrument': inst, 'style': style,
            'alpha': alpha if not np.isnan(alpha) else ann_ret * 0.5,
            'ir': ir if not np.isnan(ir) else sharpe * 0.7,
            'sharpe': sharpe, 'calmar': calmar,
            'mom': mom, 'adj_mom': adj_mom,
        }

        # 换手率因子
        if turnover_df is not None:
            fund_turn = turnover_df[turnover_df['instrument'] == inst]
            if not fund_turn.empty and 'turn' in fund_turn.columns:
                row['avg_turnover'] = fund_turn['turn'].mean()
            else:
                row['avg_turnover'] = np.nan

        # 重仓股动量
        if holdings_df is not None and stock_ret_df is not None:
            fund_h = holdings_df[holdings_df['instrument'] == inst]
            if not fund_h.empty:
                stocks = fund_h[['stock_instrument', 'holding_pct']].dropna()
                avail = [s for s in stocks['stock_instrument'].values if s in stock_ret_df.columns]
                if avail:
                    sub = stocks[stocks['stock_instrument'].isin(avail)].copy()
                    sub['w'] = sub['holding_pct'] / sub['holding_pct'].sum()
                    recent = stock_ret_df[avail].tail(63)
                    if len(recent) >= 30:
                        w = sub.set_index('stock_instrument')['w'].reindex(avail).fillna(0).values
                        port_ret = (recent.values * w).sum(axis=1)
                        row['holding_mom'] = (1 + pd.Series(port_ret)).prod() - 1
                    else:
                        row['holding_mom'] = np.nan
                else:
                    row['holding_mom'] = np.nan
            else:
                row['holding_mom'] = np.nan

        results.append(row)

    return pd.DataFrame(results) if results else pd.DataFrame()



# ============================================================
# 步骤3: 综合打分选基
# ============================================================

def select_top_funds(factor_df, n=N_PER_STYLE):
    """
    综合因子打分：
    - 基础因子(Alpha/IR/Sharpe/Calmar) 权重40%
    - 动量因子(动量/波动调整动量) 权重40%
    - 附加因子(换手率/重仓股动量) 权重20%
    每风格取Top N
    """
    selected = {}
    for style in ['大盘成长', '大盘价值', '中小盘成长', '中小盘价值']:
        sf = factor_df[factor_df['style'] == style].copy()
        if len(sf) < 3:
            continue

        def zscore(s):
            std = s.std()
            return (s - s.mean()) / std if std > 0 else pd.Series(0, index=s.index)

        # 基础因子 (各10%, 合计40%)
        sf['alpha_z'] = zscore(sf['alpha'])
        sf['ir_z'] = zscore(sf['ir'])
        sf['sharpe_z'] = zscore(sf['sharpe'])
        sf['calmar_z'] = zscore(sf['calmar'])
        base = sf['alpha_z'] * 0.1 + sf['ir_z'] * 0.1 + sf['sharpe_z'] * 0.1 + sf['calmar_z'] * 0.1

        # 动量因子 (各20%, 合计40%)
        sf['mom_z'] = zscore(sf['mom'])
        sf['adj_mom_z'] = zscore(sf['adj_mom'])
        mom_score = sf['mom_z'] * 0.2 + sf['adj_mom_z'] * 0.2

        # 附加因子 (各10%, 合计20%)
        extra = pd.Series(0.0, index=sf.index)
        if 'avg_turnover' in sf.columns and sf['avg_turnover'].notna().sum() > len(sf) * 0.3:
            extra += zscore(sf['avg_turnover'].fillna(sf['avg_turnover'].median())) * 0.1
        if 'holding_mom' in sf.columns and sf['holding_mom'].notna().sum() > len(sf) * 0.3:
            extra += zscore(sf['holding_mom'].fillna(sf['holding_mom'].median())) * 0.1

        sf['score'] = base + mom_score + extra
        top = sf.nlargest(min(n, len(sf)), 'score')
        for _, row in top.iterrows():
            selected[row['instrument']] = {'style': style, 'score': row['score']}
    return selected


# ============================================================
# 步骤4: 预计算调仓信号
# ============================================================

def precompute_all_signals():
    """预计算所有调仓日的持仓信号"""
    print("\n  [信号计算] 开始...")
    trading_dates = get_trading_dates(START_DATE, END_DATE)
    idx_start = (pd.Timestamp(START_DATE) - timedelta(days=400)).strftime('%Y-%m-%d')
    index_returns = get_index_data(idx_start, END_DATE)

    # 确定调仓日
    rebalance_dates = []
    seen = set()
    for dt in trading_dates:
        key = (dt.year, dt.month)
        if dt.month in REBALANCE_MONTHS and key not in seen:
            rebalance_dates.append(dt)
            seen.add(key)

    print(f"  [信号计算] 共{len(rebalance_dates)}个调仓日")
    signals = {}

    for i, rb_date in enumerate(rebalance_dates):
        rb_str = rb_date.strftime('%Y-%m-%d')
        lb_start = (rb_date - timedelta(days=LOOKBACK_DAYS + 60)).strftime('%Y-%m-%d')
        lb_end = (rb_date - timedelta(days=1)).strftime('%Y-%m-%d')

        # 获取LOF基金池
        fund_info = get_lof_universe(rb_str)
        if fund_info.empty:
            for d in range(1, 5):
                fund_info = get_lof_universe((rb_date - timedelta(days=d)).strftime('%Y-%m-%d'))
                if not fund_info.empty:
                    break
        if fund_info.empty:
            continue

        instruments = fund_info['instrument'].tolist()

        # 获取NAV数据并过滤
        nav_df = get_fund_nav(instruments, lb_start, lb_end)
        if nav_df.empty:
            continue
        counts = nav_df.groupby('instrument').size()
        valid = counts[counts >= int(LOOKBACK_DAYS * 0.8)].index.tolist()
        nav_df = nav_df[nav_df['instrument'].isin(valid)]
        if len(valid) < 20:
            continue

        # 风格分类
        idx_period = index_returns[(index_returns.index >= lb_start) & (index_returns.index <= lb_end)]
        fund_styles = classify_styles(nav_df, idx_period)
        if not fund_styles:
            continue

        # 获取换手率
        turn_start = (rb_date - timedelta(days=90)).strftime('%Y-%m-%d')
        turnover_df = get_fund_turnover(list(fund_styles.keys()), turn_start, lb_end)

        # 获取持仓和股票数据
        holdings_df = get_fund_holdings(list(fund_styles.keys()), lb_end)
        stock_ret_df = None
        if not holdings_df.empty:
            stock_insts = holdings_df['stock_instrument'].dropna().unique().tolist()
            if stock_insts:
                stock_start = (rb_date - timedelta(days=100)).strftime('%Y-%m-%d')
                stock_ret_df = get_stock_returns(stock_insts, stock_start, lb_end)

        # 因子计算 & 选基
        factor_df = compute_factors(nav_df, idx_period, fund_styles,
                                    turnover_df, holdings_df, stock_ret_df)
        if factor_df.empty:
            continue

        sel = select_top_funds(factor_df)
        if sel:
            signals[rb_str] = sel
            print(f"    {rb_str}: 选出{len(sel)}只基金 "
                  f"(4风格: {sum(1 for v in sel.values() if v['style']=='大盘成长')}/"
                  f"{sum(1 for v in sel.values() if v['style']=='大盘价值')}/"
                  f"{sum(1 for v in sel.values() if v['style']=='中小盘成长')}/"
                  f"{sum(1 for v in sel.values() if v['style']=='中小盘价值')})")

    print(f"  [信号计算] 完成，共{len(signals)}期有效信号")
    return signals



# ============================================================
# 步骤5: BigTrader回测
# ============================================================

def run_backtest(signals):
    """BigTrader回测执行"""
    signal_insts = {d: list(sel.keys()) for d, sel in signals.items()}
    all_insts = sorted(set(i for insts in signal_insts.values() for i in insts))

    if not all_insts:
        print("  [ERROR] 无有效信号，无法回测")
        return None

    # 等权配置
    signal_weights = {}
    for d, sel in signals.items():
        n = len(sel)
        signal_weights[d] = {inst: 1.0/n for inst in sel}

    print(f"\n  [回测] 标的池: {len(all_insts)}只基金")
    print(f"  [回测] 调仓期数: {len(signals)}")

    def initialize(context):
        context.signal_weights = signal_weights
        context.signal_insts = signal_insts
        context.last_rb = None
        context.holdings = []
        context.set_commission(bigtrader.PerOrder(
            buy_cost=SUBSCRIPTION_FEE, sell_cost=REDEMPTION_FEE, min_cost=0))

    def handle_data(context, data):
        dt_str = data.current_dt.strftime('%Y-%m-%d')
        if dt_str not in context.signal_insts:
            return
        if context.last_rb == dt_str:
            return
        context.last_rb = dt_str

        targets = set(context.signal_insts[dt_str])
        current = set(context.holdings)
        weights = context.signal_weights.get(dt_str, {})

        # 先卖后买
        for inst in current - targets:
            pos = context.get_position(inst)
            if pos and pos.current_qty > 0:
                context.order_target(inst, 0)

        for inst in targets:
            w = weights.get(inst, 1.0/len(targets))
            context.order_target_percent(inst, w)

        context.holdings = list(targets)

    perf = bigtrader.run(
        market=bigtrader.Market.CN_FUND,
        frequency=bigtrader.Frequency.DAILY,
        instruments=all_insts,
        start_date=START_DATE, end_date=END_DATE,
        capital_base=1000000,
        initialize=initialize, handle_data=handle_data,
        benchmark='000300.SH',
        order_price_field_buy='open', order_price_field_sell='open',
        volume_limit=0, before_start_days=0, render=False,
    )
    return perf


# ============================================================
# 主程序
# ============================================================

if __name__ == '__main__':
    # 1. 预计算信号
    signals = precompute_all_signals()

    # 2. 回测
    print("\n  [回测] 启动BigTrader...")
    perf = run_backtest(signals)

    # 3. 绩效输出
    if perf and perf.raw_perf is not None:
        raw = perf.raw_perf
        total_ret = raw['algorithm_period_return'].iloc[-1]
        max_dd = raw['max_drawdown'].iloc[-1]
        n_days = len(raw)
        years = n_days / 252
        daily_rets = raw['returns'].values

        ann_ret = (1 + total_ret) ** (1/years) - 1
        ann_vol = daily_rets.std() * np.sqrt(252)
        sharpe = (ann_ret - 0.025) / ann_vol if ann_vol > 0 else 0
        calmar = ann_ret / abs(max_dd) if max_dd != 0 else 0

        dates_idx = pd.to_datetime(raw['date'])
        monthly = pd.Series(daily_rets, index=dates_idx).resample('ME').apply(lambda x: (1+x).prod()-1)
        win_rate = (monthly > 0).mean()

        print("\n" + "=" * 70)
        print("  最终绩效")
        print("=" * 70)
        print(f"  年化收益率:  {ann_ret:.2%}")
        print(f"  年化波动率:  {ann_vol:.2%}")
        print(f"  夏普比率:    {sharpe:.2f}")
        print(f"  最大回撤:    {max_dd:.2%}")
        print(f"  卡玛比率:    {calmar:.2f}")
        print(f"  月度胜率:    {win_rate:.1%}")
        print(f"  累计收益:    {total_ret:.2%}")
        print("=" * 70)

        # 绘图
        try:
            import matplotlib
            matplotlib.use('Agg')
            import matplotlib.pyplot as plt
            plt.rcParams['font.sans-serif'] = ['SimHei', 'WenQuanYi Micro Hei', 'DejaVu Sans']
            plt.rcParams['axes.unicode_minus'] = False

            nav = 1 + raw['algorithm_period_return'].values
            bench = 1 + raw['benchmark_period_return'].values
            dates = pd.to_datetime(raw['date'].values)

            fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(14, 8),
                                            gridspec_kw={'height_ratios': [3, 1]})

            ax1.plot(dates, nav, color='#e74c3c', lw=2, label=f'FOF V3 (夏普:{sharpe:.2f})')
            ax1.plot(dates, bench, color='#7f8c8d', lw=1.5, ls='--', label='沪深300')
            ax1.axhline(1, color='gray', ls=':', alpha=0.3)
            ax1.legend(fontsize=11, loc='upper left')
            ax1.set_title('均衡FOF策略 V3 Final - 净值走势', fontsize=14, fontweight='bold')
            ax1.set_ylabel('净值')
            ax1.grid(alpha=0.3)

            peak = np.maximum.accumulate(nav)
            dd = (nav - peak) / peak * 100
            ax2.fill_between(dates, dd, 0, color='#e74c3c', alpha=0.3)
            ax2.plot(dates, dd, color='#e74c3c', lw=0.8)
            ax2.set_ylabel('回撤(%)')
            ax2.set_xlabel('')
            ax2.grid(alpha=0.3)

            plt.tight_layout()
            path = '/home/aiuser/work/cowork/04b25d86-d6f1-4710-a9f3-a898cc4fe30d/fof_v3_final_backtest.png'
            plt.savefig(path, dpi=150, bbox_inches='tight')
            plt.close()
            print(f"\n  净值图: {path}")
        except Exception as e:
            print(f"\n  绘图失败: {e}")
    else:
        print("\n  [ERROR] 回测失败")

# %%
