feat(factor): alpha_lab compute_factors 完整化
添加 compute_factors 方法到 AlphaLabSession: - 在 __init__ 添加 _loaded_symbols 和 _loaded_bars 缓存 - load_symbols 现在缓存 bar 数据供 compute_factors 使用 - compute_factors 使用缓存的 bars 调用 AlphaDataset - 使用懒导入避免本地 Python 3.14 缺少 polars/vnpy.alpha 的 ImportError 测试 (container only): - test_compute_factors_calls_prepare_and_fetch 验证 AlphaDataset 调用 Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
@@ -8,23 +8,25 @@ if _VNPY_SRC not in sys.path:
|
||||
|
||||
class AlphaLabSession:
|
||||
"""Session manager for vnpy.alpha AlphaLab operations."""
|
||||
|
||||
|
||||
def __init__(self, lab_path: str):
|
||||
"""
|
||||
Initialize AlphaLab session.
|
||||
|
||||
|
||||
Args:
|
||||
lab_path: Path to AlphaLab directory
|
||||
"""
|
||||
from vnpy.alpha.lab import AlphaLab
|
||||
|
||||
|
||||
self.lab_path = lab_path
|
||||
self.lab = AlphaLab(lab_path)
|
||||
self._loaded_symbols: list[str] = []
|
||||
self._loaded_bars: dict[str, list] = {}
|
||||
|
||||
def load_symbols(self, symbols: list[str], start: str, end: str, cfg) -> None:
|
||||
"""
|
||||
Load symbol data from database and save to AlphaLab.
|
||||
|
||||
|
||||
Args:
|
||||
symbols: List of vt_symbols to load
|
||||
start: Start date (YYYY-MM-DD)
|
||||
@@ -33,8 +35,55 @@ class AlphaLabSession:
|
||||
"""
|
||||
from sanguo_data.datareader import read_db_daily
|
||||
from .data_adapter import save_alpha_lab_data
|
||||
|
||||
|
||||
for symbol in symbols:
|
||||
bars = read_db_daily(symbol, start, end, cfg)
|
||||
if bars:
|
||||
save_alpha_lab_data(bars, self.lab_path)
|
||||
# Cache bars for compute_factors
|
||||
if symbol not in self._loaded_symbols:
|
||||
self._loaded_symbols.append(symbol)
|
||||
self._loaded_bars[symbol] = bars
|
||||
|
||||
def compute_factors(self, factor_names: list[str], train_period: tuple, valid_period: tuple, test_period: tuple):
|
||||
"""
|
||||
Compute factors using cached bars and vnpy.alpha AlphaDataset.
|
||||
|
||||
Args:
|
||||
factor_names: List of factor names to compute
|
||||
train_period: Training period tuple (start, end)
|
||||
valid_period: Validation period tuple (start, end)
|
||||
test_period: Test period tuple (start, end)
|
||||
|
||||
Returns:
|
||||
polars DataFrame with computed factors for test period
|
||||
"""
|
||||
# Lazy imports to avoid ImportError on local Python 3.14 without polars/vnpy.alpha
|
||||
import polars as pl
|
||||
from vnpy.alpha.dataset import AlphaDataset, Segment
|
||||
from .registry import get_factor
|
||||
from .data_adapter import convert_bars_to_alpha_df
|
||||
|
||||
# Gather all cached bars across loaded symbols
|
||||
all_bars = []
|
||||
for symbol in self._loaded_symbols:
|
||||
all_bars.extend(self._loaded_bars.get(symbol, []))
|
||||
|
||||
# Convert bars to AlphaLab DataFrame format
|
||||
df = convert_bars_to_alpha_df(all_bars)
|
||||
|
||||
# Create AlphaDataset with the specified periods
|
||||
ds = AlphaDataset(df, train_period, valid_period, test_period)
|
||||
|
||||
# Add each factor to the dataset
|
||||
for name in factor_names:
|
||||
factor = get_factor(name)
|
||||
if factor is None:
|
||||
continue # Skip unknown factors
|
||||
ds.add_feature(name, factor["expression"])
|
||||
|
||||
# Prepare data (compute features)
|
||||
ds.prepare_data(max_workers=1)
|
||||
|
||||
# Return test period data
|
||||
return ds.fetch_raw(Segment.TEST)
|
||||
|
||||
@@ -35,3 +35,92 @@ def test_load_symbols_calls_read_db_daily():
|
||||
session = AlphaLabSession(lab_path=lab_path)
|
||||
# Verify load_symbols method exists
|
||||
assert hasattr(session, "load_symbols")
|
||||
|
||||
|
||||
def test_compute_factors_calls_prepare_and_fetch(tmp_path):
|
||||
"""Test compute_factors calls AlphaDataset methods correctly."""
|
||||
from pathlib import Path
|
||||
from datetime import datetime
|
||||
from unittest.mock import MagicMock
|
||||
from sanguo_factor.alpha_lab import AlphaLabSession
|
||||
import polars as pl
|
||||
|
||||
# Create mock bar data
|
||||
mock_bar = MagicMock()
|
||||
mock_bar.vt_symbol = "600000.SSE"
|
||||
mock_bar.datetime = datetime(2024, 1, 1)
|
||||
mock_bar.open_price = 1.0
|
||||
mock_bar.high_price = 1.0
|
||||
mock_bar.low_price = 1.0
|
||||
mock_bar.close_price = 1.0
|
||||
mock_bar.volume = 1
|
||||
mock_bar.turnover = 0
|
||||
mock_bar.open_interest = 0
|
||||
|
||||
# Mock AlphaLab session initialization
|
||||
with patch("vnpy.alpha.lab.AlphaLab"), \
|
||||
patch("sanguo_data.datareader.read_db_daily") as mock_read, \
|
||||
patch("sanguo_factor.data_adapter.save_alpha_lab_data"), \
|
||||
patch("sanguo_factor.data_adapter.convert_bars_to_alpha_df") as mock_convert, \
|
||||
patch("sanguo_factor.registry.get_factor") as mock_get_factor, \
|
||||
patch("vnpy.alpha.dataset.AlphaDataset") as mock_dataset_class:
|
||||
|
||||
# Setup mock returns
|
||||
mock_read.return_value = [mock_bar]
|
||||
mock_get_factor.return_value = {"expression": "ts_mean(close,5)"}
|
||||
|
||||
# Create mock AlphaDataset instance
|
||||
mock_dataset = MagicMock()
|
||||
mock_dataset.fetch_raw.return_value = pl.DataFrame({
|
||||
"datetime": [],
|
||||
"vt_symbol": [],
|
||||
"ma5": []
|
||||
})
|
||||
mock_dataset_class.return_value = mock_dataset
|
||||
|
||||
# Create DataFrame for convert_bars_to_alpha_df
|
||||
mock_df = pl.DataFrame({
|
||||
"vt_symbol": ["600000.SSE"],
|
||||
"datetime": [datetime(2024, 1, 1)],
|
||||
"open": [1.0],
|
||||
"high": [1.0],
|
||||
"low": [1.0],
|
||||
"close": [1.0],
|
||||
"volume": [1.0],
|
||||
"turnover": [0.0],
|
||||
"open_interest": [0.0]
|
||||
})
|
||||
mock_convert.return_value = mock_df
|
||||
|
||||
# Create session and load symbols
|
||||
lab_path = str(tmp_path / "alpha_lab")
|
||||
session = AlphaLabSession(lab_path=lab_path)
|
||||
session.load_symbols(["600000"], "2024-01-01", "2024-06-30", cfg=MagicMock())
|
||||
|
||||
# Compute factors
|
||||
df = session.compute_factors(
|
||||
["ma5"],
|
||||
("2024-01-01", "2024-04-30"),
|
||||
("2024-05-01", "2024-05-15"),
|
||||
("2024-05-16", "2024-06-30")
|
||||
)
|
||||
|
||||
# Verify AlphaDataset was created with correct periods
|
||||
mock_dataset_class.assert_called_once_with(
|
||||
mock_df,
|
||||
("2024-01-01", "2024-04-30"),
|
||||
("2024-05-01", "2024-05-15"),
|
||||
("2024-05-16", "2024-06-30")
|
||||
)
|
||||
|
||||
# Verify add_feature was called for the factor
|
||||
mock_dataset.add_feature.assert_called_once_with("ma5", "ts_mean(close,5)")
|
||||
|
||||
# Verify prepare_data was called
|
||||
mock_dataset.prepare_data.assert_called_once_with(max_workers=1)
|
||||
|
||||
# Verify fetch_raw was called
|
||||
assert mock_dataset.fetch_raw.called
|
||||
|
||||
# Verify return value is a DataFrame
|
||||
assert isinstance(df, pl.DataFrame)
|
||||
|
||||
Reference in New Issue
Block a user