feat(backtest): cta_optimizer run_optimization wrapper
This commit is contained in:
@@ -0,0 +1,181 @@
|
||||
"""CTA strategy parameter optimization wrapper using vnpy_ctastrategy.backtesting."""
|
||||
import sys
|
||||
import os
|
||||
import traceback
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import List, Any
|
||||
|
||||
# Add vnpy source to path for local development
|
||||
_VNPY_SRC = os.path.join(os.path.dirname(__file__), "..", "vnpy_v4.4.0")
|
||||
_VNPY_SRC = os.path.abspath(_VNPY_SRC)
|
||||
if _VNPY_SRC not in sys.path:
|
||||
sys.path.insert(0, _VNPY_SRC)
|
||||
|
||||
from sanguo_backtest.result_store import BacktestResult, save_result
|
||||
from sanguo_backtest.cta_engine import guess_exchange, Exchange
|
||||
|
||||
|
||||
def run_cta_optimization(
|
||||
strategy_class,
|
||||
symbol: str,
|
||||
grid: dict,
|
||||
start: str,
|
||||
end: str,
|
||||
cfg,
|
||||
db_path: str,
|
||||
max_workers: int = 2
|
||||
) -> List[BacktestResult]:
|
||||
"""
|
||||
Run CTA strategy parameter optimization using vnpy_ctastrategy BacktestingEngine.
|
||||
|
||||
Args:
|
||||
strategy_class: CTA strategy class to optimize
|
||||
symbol: Stock symbol (e.g., "600000")
|
||||
grid: Parameter grid dict {name: (start, end, step)}
|
||||
start: Backtest start date (YYYY-MM-DD format)
|
||||
end: Backtest end date (YYYY-MM-DD format)
|
||||
cfg: Configuration object (may contain data paths)
|
||||
db_path: SQLite database path for saving results
|
||||
max_workers: Maximum number of parallel optimization workers
|
||||
|
||||
Returns:
|
||||
List[BacktestResult]: List of result objects with optimization statistics
|
||||
"""
|
||||
# Generate unique task ID for this optimization run
|
||||
task_id = f"opt_{uuid.uuid4().hex[:8]}"
|
||||
|
||||
try:
|
||||
# Lazy import of BacktestingEngine and OptimizationSetting
|
||||
from vnpy_ctastrategy.backtesting import BacktestingEngine
|
||||
from vnpy.trader.optimize import OptimizationSetting
|
||||
|
||||
# Build vt_symbol for A-shares
|
||||
vt_symbol = f"{symbol}.{guess_exchange(symbol).value}"
|
||||
|
||||
# Convert date strings to datetime objects
|
||||
start_dt = datetime.strptime(start, "%Y-%m-%d")
|
||||
end_dt = datetime.strptime(end, "%Y-%m-%d") if end else None
|
||||
|
||||
# Create and configure backtesting engine
|
||||
engine = BacktestingEngine()
|
||||
|
||||
# Set parameters with A-share specific values (same as cta_engine)
|
||||
engine.set_parameters(
|
||||
vt_symbol=vt_symbol,
|
||||
interval="1d", # Daily interval for A-shares
|
||||
start=start_dt,
|
||||
end=end_dt,
|
||||
rate=0.001, # Commission rate (0.1% for A-shares)
|
||||
slippage=0, # No slippage for simplicity
|
||||
size=1, # Contract size (1 for stocks)
|
||||
pricetick=0.01, # Minimum price tick (0.01 yuan for A-shares)
|
||||
capital=0 # No initial capital limit
|
||||
)
|
||||
|
||||
# Add strategy without parameters (will be set by optimization)
|
||||
engine.add_strategy(strategy_class, {})
|
||||
|
||||
# Load historical data
|
||||
engine.load_data()
|
||||
|
||||
# Create optimization setting
|
||||
setting = OptimizationSetting()
|
||||
setting.set_target("sharpe_ratio") # Optimize for Sharpe ratio
|
||||
|
||||
# Add parameter ranges to optimization setting
|
||||
for name, (start_val, end_val, step) in grid.items():
|
||||
setting.add_parameter(name, start_val, end_val, step)
|
||||
|
||||
# Run optimization (headless, with parallel workers)
|
||||
optimization_results = engine.run_optimization(
|
||||
setting,
|
||||
output=False, # Headless mode
|
||||
max_workers=max_workers
|
||||
)
|
||||
|
||||
# Parse optimization results and convert to BacktestResult objects
|
||||
results = []
|
||||
for item in optimization_results:
|
||||
try:
|
||||
# Handle both tuple format (params, target_value, statistics)
|
||||
# and dict format ({params: ..., statistics: ...})
|
||||
if isinstance(item, tuple) and len(item) >= 3:
|
||||
params = item[0]
|
||||
statistics = item[2]
|
||||
elif isinstance(item, dict):
|
||||
params = item.get("params", {})
|
||||
statistics = item.get("statistics", {})
|
||||
else:
|
||||
# Unknown format, skip this result
|
||||
continue
|
||||
|
||||
# Create individual result for each optimization run
|
||||
result = BacktestResult(
|
||||
task_id=f"opt_{uuid.uuid4().hex[:8]}", # Unique ID per result
|
||||
type="optimize",
|
||||
status="done",
|
||||
strategy=strategy_class.__name__,
|
||||
symbol=symbol,
|
||||
params=params,
|
||||
start=start,
|
||||
end=end,
|
||||
statistics=statistics,
|
||||
equity_curve=None, # Not available in optimization results
|
||||
trades=None # Not available in optimization results
|
||||
)
|
||||
results.append(result)
|
||||
|
||||
# Save each result to database
|
||||
save_result(result, db_path=db_path)
|
||||
|
||||
except Exception as e:
|
||||
# Handle individual result parsing error
|
||||
error_result = BacktestResult(
|
||||
task_id=f"opt_{uuid.uuid4().hex[:8]}",
|
||||
type="optimize",
|
||||
status="failed",
|
||||
strategy=strategy_class.__name__,
|
||||
symbol=symbol,
|
||||
params={},
|
||||
start=start,
|
||||
end=end,
|
||||
statistics={},
|
||||
equity_curve=None,
|
||||
trades=None,
|
||||
error_msg=f"Result parsing error: {type(e).__name__}: {e}"
|
||||
)
|
||||
results.append(error_result)
|
||||
save_result(error_result, db_path=db_path)
|
||||
|
||||
return results
|
||||
|
||||
except Exception as e:
|
||||
# Handle any exceptions during optimization setup/execution
|
||||
error_msg = f"{type(e).__name__}: {e}\n{traceback.format_exc()}"
|
||||
|
||||
failed_result = BacktestResult(
|
||||
task_id=task_id,
|
||||
type="optimize",
|
||||
status="failed",
|
||||
strategy=strategy_class.__name__,
|
||||
symbol=symbol,
|
||||
params={},
|
||||
start=start,
|
||||
end=end,
|
||||
statistics={},
|
||||
equity_curve=None,
|
||||
trades=None,
|
||||
error_msg=error_msg
|
||||
)
|
||||
|
||||
# Save failed result to database
|
||||
save_result(failed_result, db_path=db_path)
|
||||
|
||||
return [failed_result]
|
||||
|
||||
|
||||
# Module-level reference for mocking in tests
|
||||
BacktestingEngine = None # Will be set when imported inside run_cta_optimization
|
||||
OptimizationSetting = None # Will be set when imported inside run_cta_optimization
|
||||
@@ -0,0 +1,263 @@
|
||||
"""Tests for sanguo_backtest.cta_optimizer module."""
|
||||
import pytest
|
||||
from unittest.mock import Mock, patch, MagicMock
|
||||
from sanguo_backtest.cta_optimizer import run_cta_optimization
|
||||
|
||||
|
||||
class TestRunCtaOptimization:
|
||||
"""Test suite for run_cta_optimization function."""
|
||||
|
||||
def test_run_optimization_returns_results(self, temp_db_path):
|
||||
"""Test successful optimization execution returns proper results."""
|
||||
# Mock strategy class
|
||||
mock_strategy_class = Mock()
|
||||
mock_strategy_class.__name__ = "TestStrategy"
|
||||
|
||||
# Mock engine
|
||||
mock_engine = MagicMock()
|
||||
|
||||
# Mock run_optimization to return list of results (tuple format)
|
||||
mock_engine.run_optimization.return_value = [
|
||||
({"n": 5}, 1.0, {"sharpe_ratio": 1.0, "total_return": 0.10}),
|
||||
({"n": 10}, 1.5, {"sharpe_ratio": 1.5, "total_return": 0.15}),
|
||||
]
|
||||
|
||||
# Mock OptimizationSetting
|
||||
mock_setting = MagicMock()
|
||||
mock_setting.set_target.return_value = None
|
||||
mock_setting.add_parameter.return_value = (True, "")
|
||||
|
||||
# Mock config
|
||||
mock_cfg = Mock()
|
||||
|
||||
# Create mock modules
|
||||
mock_ctastrategy = MagicMock()
|
||||
mock_ctastrategy.BacktestingEngine = Mock(return_value=mock_engine)
|
||||
|
||||
mock_trader = MagicMock()
|
||||
mock_trader.OptimizationSetting = Mock(return_value=mock_setting)
|
||||
|
||||
with patch.dict("sys.modules", {
|
||||
"vnpy_ctastrategy.backtesting": mock_ctastrategy,
|
||||
"vnpy.trader.optimize": mock_trader
|
||||
}):
|
||||
results = run_cta_optimization(
|
||||
strategy_class=mock_strategy_class,
|
||||
symbol="600000",
|
||||
grid={"n": (5, 20, 5)},
|
||||
start="2024-01-01",
|
||||
end="2024-03-31",
|
||||
cfg=mock_cfg,
|
||||
db_path=temp_db_path,
|
||||
max_workers=2
|
||||
)
|
||||
|
||||
# Verify results
|
||||
assert len(results) == 2
|
||||
assert all(result.type == "optimize" for result in results)
|
||||
assert all(result.status == "done" for result in results)
|
||||
assert all(result.strategy == "TestStrategy" for result in results)
|
||||
assert all(result.symbol == "600000" for result in results)
|
||||
assert all(result.start == "2024-01-01" for result in results)
|
||||
assert all(result.end == "2024-03-31" for result in results)
|
||||
|
||||
# Verify first result
|
||||
assert results[0].params == {"n": 5}
|
||||
assert results[0].statistics == {"sharpe_ratio": 1.0, "total_return": 0.10}
|
||||
|
||||
# Verify second result
|
||||
assert results[1].params == {"n": 10}
|
||||
assert results[1].statistics == {"sharpe_ratio": 1.5, "total_return": 0.15}
|
||||
|
||||
# Verify engine methods were called
|
||||
mock_engine.set_parameters.assert_called_once()
|
||||
mock_engine.add_strategy.assert_called_once_with(mock_strategy_class, {})
|
||||
mock_engine.run_optimization.assert_called_once()
|
||||
|
||||
# Verify OptimizationSetting configuration
|
||||
mock_setting.set_target.assert_called_once_with("sharpe_ratio")
|
||||
mock_setting.add_parameter.assert_called_once_with("n", 5, 20, 5)
|
||||
|
||||
def test_run_optimization_handles_dict_format_results(self, temp_db_path):
|
||||
"""Test optimization with dict format results (alternative format)."""
|
||||
mock_strategy_class = Mock()
|
||||
mock_strategy_class.__name__ = "DictStrategy"
|
||||
|
||||
mock_engine = MagicMock()
|
||||
|
||||
# Mock run_optimization to return list of dict results
|
||||
mock_engine.run_optimization.return_value = [
|
||||
{
|
||||
"params": {"window": 20},
|
||||
"statistics": {"sharpe_ratio": 1.2}
|
||||
},
|
||||
{
|
||||
"params": {"window": 30},
|
||||
"statistics": {"sharpe_ratio": 1.8}
|
||||
}
|
||||
]
|
||||
|
||||
mock_setting = MagicMock()
|
||||
mock_setting.set_target.return_value = None
|
||||
mock_setting.add_parameter.return_value = (True, "")
|
||||
|
||||
mock_cfg = Mock()
|
||||
|
||||
mock_ctastrategy = MagicMock()
|
||||
mock_ctastrategy.BacktestingEngine = Mock(return_value=mock_engine)
|
||||
|
||||
mock_trader = MagicMock()
|
||||
mock_trader.OptimizationSetting = Mock(return_value=mock_setting)
|
||||
|
||||
with patch.dict("sys.modules", {
|
||||
"vnpy_ctastrategy.backtesting": mock_ctastrategy,
|
||||
"vnpy.trader.optimize": mock_trader
|
||||
}):
|
||||
results = run_cta_optimization(
|
||||
strategy_class=mock_strategy_class,
|
||||
symbol="000001",
|
||||
grid={"window": (20, 30, 10)},
|
||||
start="2024-01-01",
|
||||
end="2024-03-31",
|
||||
cfg=mock_cfg,
|
||||
db_path=temp_db_path,
|
||||
max_workers=2
|
||||
)
|
||||
|
||||
# Verify dict format results were parsed correctly
|
||||
assert len(results) == 2
|
||||
assert results[0].params == {"window": 20}
|
||||
assert results[0].statistics == {"sharpe_ratio": 1.2}
|
||||
assert results[1].params == {"window": 30}
|
||||
assert results[1].statistics == {"sharpe_ratio": 1.8}
|
||||
|
||||
def test_optimization_handles_failure(self, temp_db_path):
|
||||
"""Test that exceptions during optimization are handled properly."""
|
||||
mock_strategy_class = Mock()
|
||||
mock_strategy_class.__name__ = "FailingStrategy"
|
||||
|
||||
mock_engine = MagicMock()
|
||||
mock_engine.run_optimization.side_effect = RuntimeError("Optimization failed")
|
||||
|
||||
mock_setting = MagicMock()
|
||||
|
||||
mock_cfg = Mock()
|
||||
|
||||
mock_ctastrategy = MagicMock()
|
||||
mock_ctastrategy.BacktestingEngine = Mock(return_value=mock_engine)
|
||||
|
||||
mock_trader = MagicMock()
|
||||
mock_trader.OptimizationSetting = Mock(return_value=mock_setting)
|
||||
|
||||
with patch.dict("sys.modules", {
|
||||
"vnpy_ctastrategy.backtesting": mock_ctastrategy,
|
||||
"vnpy.trader.optimize": mock_trader
|
||||
}):
|
||||
results = run_cta_optimization(
|
||||
strategy_class=mock_strategy_class,
|
||||
symbol="600000",
|
||||
grid={"n": (5, 20, 5)},
|
||||
start="2024-01-01",
|
||||
end="2024-03-31",
|
||||
cfg=mock_cfg,
|
||||
db_path=temp_db_path,
|
||||
max_workers=2
|
||||
)
|
||||
|
||||
# Verify failed result
|
||||
assert len(results) == 1
|
||||
assert results[0].type == "optimize"
|
||||
assert results[0].status == "failed"
|
||||
assert results[0].strategy == "FailingStrategy"
|
||||
assert results[0].symbol == "600000"
|
||||
assert results[0].error_msg is not None
|
||||
assert "RuntimeError" in results[0].error_msg
|
||||
assert "Optimization failed" in results[0].error_msg
|
||||
|
||||
def test_optimization_with_multiple_parameters(self, temp_db_path):
|
||||
"""Test optimization with multiple parameters in grid."""
|
||||
mock_strategy_class = Mock()
|
||||
mock_strategy_class.__name__ = "MultiParamStrategy"
|
||||
|
||||
mock_engine = MagicMock()
|
||||
mock_engine.run_optimization.return_value = [
|
||||
({"n": 5, "window": 10}, 1.0, {"sharpe_ratio": 1.0}),
|
||||
]
|
||||
|
||||
mock_setting = MagicMock()
|
||||
mock_setting.set_target.return_value = None
|
||||
mock_setting.add_parameter.return_value = (True, "")
|
||||
|
||||
mock_cfg = Mock()
|
||||
|
||||
mock_ctastrategy = MagicMock()
|
||||
mock_ctastrategy.BacktestingEngine = Mock(return_value=mock_engine)
|
||||
|
||||
mock_trader = MagicMock()
|
||||
mock_trader.OptimizationSetting = Mock(return_value=mock_setting)
|
||||
|
||||
with patch.dict("sys.modules", {
|
||||
"vnpy_ctastrategy.backtesting": mock_ctastrategy,
|
||||
"vnpy.trader.optimize": mock_trader
|
||||
}):
|
||||
results = run_cta_optimization(
|
||||
strategy_class=mock_strategy_class,
|
||||
symbol="600000",
|
||||
grid={
|
||||
"n": (5, 20, 5),
|
||||
"window": (10, 30, 10)
|
||||
},
|
||||
start="2024-01-01",
|
||||
end="2024-03-31",
|
||||
cfg=mock_cfg,
|
||||
db_path=temp_db_path,
|
||||
max_workers=4
|
||||
)
|
||||
|
||||
# Verify multiple parameters were added to setting
|
||||
assert mock_setting.add_parameter.call_count == 2
|
||||
assert len(results) == 1
|
||||
assert results[0].params == {"n": 5, "window": 10}
|
||||
|
||||
def test_optimization_generates_unique_task_ids(self, temp_db_path):
|
||||
"""Test that each optimization result gets unique task ID."""
|
||||
mock_strategy_class = Mock()
|
||||
mock_strategy_class.__name__ = "IdStrategy"
|
||||
|
||||
mock_engine = MagicMock()
|
||||
mock_engine.run_optimization.return_value = [
|
||||
({"n": 5}, 1.0, {"sharpe_ratio": 1.0}),
|
||||
({"n": 10}, 1.5, {"sharpe_ratio": 1.5}),
|
||||
]
|
||||
|
||||
mock_setting = MagicMock()
|
||||
mock_setting.set_target.return_value = None
|
||||
mock_setting.add_parameter.return_value = (True, "")
|
||||
|
||||
mock_cfg = Mock()
|
||||
|
||||
mock_ctastrategy = MagicMock()
|
||||
mock_ctastrategy.BacktestingEngine = Mock(return_value=mock_engine)
|
||||
|
||||
mock_trader = MagicMock()
|
||||
mock_trader.OptimizationSetting = Mock(return_value=mock_setting)
|
||||
|
||||
with patch.dict("sys.modules", {
|
||||
"vnpy_ctastrategy.backtesting": mock_ctastrategy,
|
||||
"vnpy.trader.optimize": mock_trader
|
||||
}):
|
||||
results = run_cta_optimization(
|
||||
strategy_class=mock_strategy_class,
|
||||
symbol="600000",
|
||||
grid={"n": (5, 20, 5)},
|
||||
start="2024-01-01",
|
||||
end="2024-03-31",
|
||||
cfg=mock_cfg,
|
||||
db_path=temp_db_path,
|
||||
max_workers=2
|
||||
)
|
||||
|
||||
# Verify unique task IDs
|
||||
assert results[0].task_id != results[1].task_id
|
||||
assert results[0].task_id.startswith("opt_")
|
||||
assert results[1].task_id.startswith("opt_")
|
||||
Reference in New Issue
Block a user