"""Factor analysis with alphalens - lazy import to avoid ImportError.""" import sys import os import warnings _VNPY_SRC = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "vnpy_v4.4.0")) if _VNPY_SRC not in sys.path: sys.path.insert(0, _VNPY_SRC) from dataclasses import dataclass, field from typing import TYPE_CHECKING # Module-level imports for patch targets (with try/except guards for local importability) try: from alphalens.utils import get_clean_factor_and_forward_returns from alphalens.tears import create_full_tear_sheet from alphalens.performance import factor_information_coefficient except ImportError: # alphalens not available locally - set to None for patch targets get_clean_factor_and_forward_returns = None create_full_tear_sheet = None factor_information_coefficient = None try: from .alpha_lab import AlphaLabSession except ImportError: # AlphaLabSession not available - set to None for patch targets AlphaLabSession = None if TYPE_CHECKING: # Type hints only - not imported at runtime to avoid ImportError import polars as pl @dataclass class FactorReport: """Factor analysis report.""" factor_names: list[str] output_dir: str ic_summary: dict = field(default_factory=dict) report_paths: dict = field(default_factory=dict) def run_factor_analysis( symbols: list[str], factor_names: list[str], start: str, end: str, cfg, output_dir: str, periods: tuple = (1, 5, 10) ) -> FactorReport: """ Run factor analysis using AlphaLabSession and alphalens. Args: symbols: List of vt_symbols to analyze factor_names: List of factor names to compute start: Start date (YYYY-MM-DD) end: End date (YYYY-MM-DD) cfg: Database configuration object output_dir: Output directory for analysis results periods: Forward return periods for IC analysis (default: 1, 5, 10 days) Returns: FactorReport with analysis results including tears report and IC values """ from .registry import get_factor # API path passes cfg=None → load default data_platform.yaml (so read_db_daily # and AlphaLabSession can find the A-share DB). if cfg is None: from sanguo_data.config import load_config, find_config_path cfg = load_config(find_config_path()) # Check if alphalens is available if get_clean_factor_and_forward_returns is None or create_full_tear_sheet is None or factor_information_coefficient is None: return FactorReport( factor_names=factor_names, output_dir=output_dir, ic_summary={"error": "alphalens not installed"}, report_paths={} ) if AlphaLabSession is None: return FactorReport( factor_names=factor_names, output_dir=output_dir, ic_summary={"error": "AlphaLabSession not available"}, report_paths={} ) # Lazy imports for container environment try: import polars as pl import pandas as pd import matplotlib matplotlib.use("Agg") # Use non-interactive backend for headless operation import matplotlib.pyplot as plt except ImportError as e: return FactorReport( factor_names=factor_names, output_dir=output_dir, ic_summary={"error": f"Required import missing: {e}"}, report_paths={} ) # Create AlphaLab session and load symbols session = AlphaLabSession(lab_path=output_dir) session.load_symbols(symbols, start, end, cfg) # Calculate period split (simple deterministic split) from datetime import datetime start_dt = datetime.strptime(start, "%Y-%m-%d") end_dt = datetime.strptime(end, "%Y-%m-%d") total_days = (end_dt - start_dt).days # Simple split: train = first half, valid = empty, test = second half mid_point = start_dt + pd.Timedelta(days=total_days // 2) train_period = (start, mid_point.strftime("%Y-%m-%d")) valid_period = (mid_point.strftime("%Y-%m-%d"), mid_point.strftime("%Y-%m-%d")) test_period = (mid_point.strftime("%Y-%m-%d"), end) # Compute factors using AlphaLabSession factor_df = session.compute_factors(factor_names, train_period, valid_period, test_period) # Load close prices separately for tears computation # (factor_df only contains factor columns, not OHLCV) from sanguo_data.datareader import read_db_daily from datetime import datetime from zoneinfo import ZoneInfo # Convert string dates to datetime for database query _SH = ZoneInfo("Asia/Shanghai") start_dt = datetime.strptime(start, "%Y-%m-%d").replace(tzinfo=_SH) end_dt = datetime.strptime(end, "%Y-%m-%d").replace(tzinfo=_SH) # Load bars for close prices all_bars = [] for symbol in symbols: try: bars = read_db_daily(symbol, start_dt.strftime("%Y-%m-%d"), end_dt.strftime("%Y-%m-%d"), cfg) all_bars.extend(bars) except Exception as e: warnings.warn(f"Failed to load bars for {symbol}: {e}") continue # Create close price DataFrame if all_bars: close_df = pl.DataFrame({ "datetime": [b.datetime for b in all_bars], "vt_symbol": [b.vt_symbol for b in all_bars], "close": [b.close_price for b in all_bars] }) else: warnings.warn("No bars loaded for close prices - tears computation will fail") close_df = pl.DataFrame(schema={"datetime": pl.Datetime, "vt_symbol": pl.Utf8, "close": pl.Float64}) # Initialize IC summary and report paths ic_summary = {} report_paths = {} # Process each factor for factor_name in factor_names: try: # Convert polars DataFrame to pandas for alphalens factor_pd = factor_df.to_pandas() close_pd = close_df.to_pandas() # Check if factor column exists if factor_name not in factor_pd.columns: # If the specific factor name isn't found, use the last column # (compute_factors returns factors with their names as columns) factor_cols = [col for col in factor_pd.columns if col not in ["datetime", "vt_symbol"]] if factor_cols: factor_col = factor_cols[0] # Use first available factor column else: continue # No factor columns found else: factor_col = factor_name # Set MultiIndex (datetime, vt_symbol) as required by alphalens factor_pd["datetime"] = pd.to_datetime(factor_pd["datetime"]) factor_series = factor_pd.set_index(["datetime", "vt_symbol"])[factor_col] # Build prices DataFrame from separately loaded close prices # Localize close datetimes to Asia/Shanghai-aware to match factor_df's # aware datetimes (compute_factors localizes), else the date-alignment # filter (prices.index.isin(factor_dates)) empties prices → concat error. _close_dt = pd.to_datetime(close_pd["datetime"]) if _close_dt.dt.tz is None: _close_dt = _close_dt.dt.tz_localize("Asia/Shanghai") else: _close_dt = _close_dt.dt.tz_convert("Asia/Shanghai") close_pd["datetime"] = _close_dt prices_df = close_pd.pivot(index="datetime", columns="vt_symbol", values="close") # CRITICAL FIX: Align price data with factor data date range # Factor data only contains test period, but price data contains full range # Filter prices to only include dates that exist in factor data factor_dates = factor_series.index.get_level_values('datetime').unique() prices_df = prices_df[prices_df.index.isin(factor_dates)] # Ensure datetime index for prices prices_df.index = pd.to_datetime(prices_df.index) # Call get_clean_factor_and_forward_returns merged_data = get_clean_factor_and_forward_returns( factor=factor_series, prices=prices_df, periods=periods, # Use configurable periods max_loss=1.0 # TEMPORARY: Allow 100% loss to see IC data ) # Extract IC values using factor_information_coefficient ic_data = {} try: ic_df = factor_information_coefficient(merged_data) # Compute IC statistics for each period for period_col in ic_df.columns: period_name = f"{period_col}D" if period_col.isdigit() else period_col # Extract IC values for this period (drop NaN values) period_ic_values = ic_df[period_col].dropna() if len(period_ic_values) > 0: ic_mean = float(period_ic_values.mean()) ic_std = float(period_ic_values.std()) icir = ic_mean / ic_std if ic_std > 0 else 0.0 # Compute t-statistic if we have enough samples n = len(period_ic_values) t_stat = ic_mean / (ic_std / (n ** 0.5)) if ic_std > 0 and n > 1 else 0.0 ic_data[period_name] = { "mean": ic_mean, "std": ic_std, "icir": icir, "t_stat": t_stat, "count": n } else: ic_data[period_name] = { "error": "No valid IC values for period" } except Exception as ic_error: # Capture IC extraction error but continue with tears report ic_data = {"error": f"IC extraction failed: {type(ic_error).__name__}: {ic_error}"} # Generate tears sheet from io import StringIO import sys old_stdout = sys.stdout sys.stdout = StringIO() # Capture stdout to avoid display issues try: create_full_tear_sheet( merged_data, long_short=True, group_neutral=False, by_group=False ) finally: sys.stdout = old_stdout # Restore stdout # Save the tears report factor_report_path = os.path.join(output_dir, f"{factor_name}_tears.html") plt.savefig(factor_report_path.replace(".html", ".png")) # Save as PNG report_paths[factor_name] = factor_report_path.replace(".png", ".html") # Mark HTML as report # Store basic IC summary (simplified) — close prices sourced from DB (real) status = "success" ic_summary[factor_name] = { "status": status, "report": factor_report_path, "ic": ic_data # Add IC statistics } except Exception as e: err_type = type(e).__name__ ic_summary[factor_name] = { "status": "error", "error": f"{err_type}: {e}" } return FactorReport( factor_names=factor_names, output_dir=output_dir, ic_summary=ic_summary, report_paths=report_paths )