Skip to content

Commit 5a661c0

Browse files
Merge pull request #184 from akfamily/dev
feat(backtest): 增加策略参数严格校验并优化并行网格搜索
2 parents c681fdc + 891a7ad commit 5a661c0

7 files changed

Lines changed: 280 additions & 16 deletions

File tree

Cargo.lock

Lines changed: 1 addition & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

Cargo.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
[package]
22
name = "akquant"
3-
version = "0.1.93"
3+
version = "0.1.94"
44
edition = "2024"
55
description = "High-performance quantitative trading framework based on Rust and Python"
66
license = "MIT"

pyproject.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@ build-backend = "maturin"
44

55
[project]
66
name = "akquant"
7-
version = "0.1.93"
7+
version = "0.1.94"
88
description = "High-performance quantitative trading framework based on Rust and Python"
99
readme = "README.md"
1010
license = {text = "MIT License"}

python/akquant/backtest/__init__.pyi

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -108,6 +108,7 @@ def run_backtest(
108108
timer_execution_policy: Literal["same_cycle", "next_event"] = ...,
109109
fill_policy: Optional[FillPolicy] = ...,
110110
stream_mode: Literal["observability", "audit"] = ...,
111+
strict_strategy_params: bool = True,
111112
**kwargs: Any,
112113
) -> BacktestResult: ...
113114
def run_warm_start(

python/akquant/backtest/engine.py

Lines changed: 81 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -221,6 +221,45 @@ def _accepts_strategy_kwarg(
221221
)
222222

223223

224+
def _split_strategy_kwargs(
225+
strategy_input: Union[Type[Strategy], Strategy, Callable[[Any, Bar], None], None],
226+
strategy_kwargs: Dict[str, Any],
227+
) -> Tuple[Dict[str, Any], List[str]]:
228+
"""Split kwargs into constructor-accepted kwargs and unknown keys."""
229+
if not isinstance(strategy_input, type) or not issubclass(strategy_input, Strategy):
230+
return strategy_kwargs, []
231+
232+
try:
233+
signature = inspect.signature(strategy_input.__init__)
234+
except (TypeError, ValueError):
235+
return strategy_kwargs, []
236+
237+
supports_var_kwargs = any(
238+
parameter.kind == inspect.Parameter.VAR_KEYWORD
239+
for parameter in signature.parameters.values()
240+
)
241+
if supports_var_kwargs:
242+
return strategy_kwargs, []
243+
244+
accepted_names = {
245+
parameter_name
246+
for parameter_name, parameter in signature.parameters.items()
247+
if parameter_name != "self"
248+
and parameter.kind
249+
in {
250+
inspect.Parameter.POSITIONAL_OR_KEYWORD,
251+
inspect.Parameter.KEYWORD_ONLY,
252+
}
253+
}
254+
accepted_kwargs = {
255+
key: value for key, value in strategy_kwargs.items() if key in accepted_names
256+
}
257+
unknown_keys = sorted(
258+
key for key in strategy_kwargs.keys() if key not in accepted_names
259+
)
260+
return accepted_kwargs, unknown_keys
261+
262+
224263
def _maybe_warn_deprecated_symbol_argument(
225264
*,
226265
symbol: Union[str, List[str]],
@@ -478,6 +517,7 @@ def _load_data_map_from_adapter(
478517
def _build_strategy_instance(
479518
strategy: Union[Type[Strategy], Strategy, Callable[[Any, Bar], None], None],
480519
strategy_kwargs: Dict[str, Any],
520+
strict_strategy_params: bool,
481521
logger: Any,
482522
initialize: Optional[Callable[[Any], None]],
483523
on_start: Optional[Callable[[Any], None]],
@@ -489,9 +529,32 @@ def _build_strategy_instance(
489529
context: Optional[Dict[str, Any]],
490530
) -> Strategy:
491531
if isinstance(strategy, type) and issubclass(strategy, Strategy):
532+
accepted_kwargs, unknown_keys = _split_strategy_kwargs(
533+
strategy, strategy_kwargs
534+
)
535+
if unknown_keys:
536+
unknown_keys_text = ", ".join(unknown_keys)
537+
if strict_strategy_params:
538+
raise TypeError(
539+
"Unknown strategy constructor parameter(s): "
540+
f"{unknown_keys_text}. Strategy={strategy.__module__}."
541+
f"{strategy.__name__}"
542+
)
543+
logger.warning(
544+
"Ignoring unknown strategy constructor parameter(s): %s. "
545+
"Strategy=%s.%s",
546+
unknown_keys_text,
547+
strategy.__module__,
548+
strategy.__name__,
549+
)
492550
try:
493-
return cast(Strategy, strategy(**strategy_kwargs))
551+
return cast(Strategy, strategy(**accepted_kwargs))
494552
except TypeError as e:
553+
if strict_strategy_params:
554+
raise TypeError(
555+
"Failed to instantiate strategy with provided parameters: "
556+
f"{e}. Strategy={strategy.__module__}.{strategy.__name__}"
557+
) from e
495558
logger.warning(
496559
f"Failed to instantiate strategy with provided parameters: {e}. "
497560
"Falling back to default constructor (no arguments)."
@@ -752,6 +815,7 @@ def run_backtest(
752815
broker_profile: Optional[str] = None,
753816
timer_execution_policy: str = "same_cycle",
754817
fill_policy: Optional[FillPolicy] = None,
818+
strict_strategy_params: bool = True,
755819
**kwargs: Any,
756820
) -> BacktestResult:
757821
"""
@@ -793,6 +857,9 @@ def run_backtest(
793857
"temporal": "same_cycle|next_event"}。
794858
预留未实现 price_basis: mid_quote、vwap_window、twap_window。
795859
若提供该参数,则其语义优先于 execution_mode 与 timer_execution_policy。
860+
:param strict_strategy_params: 是否严格校验策略构造参数。True 时若参数不匹配将抛错;
861+
False 时保持兼容行为(忽略未知参数并在失败时
862+
回退无参构造)。
796863
:param timezone: 时区名称 (默认 "Asia/Shanghai")
797864
:param t_plus_one: 是否启用 T+1 交易规则 (默认 False)
798865
:param initialize: 初始化回调函数 (仅当 strategy 为函数时使用)
@@ -1191,18 +1258,6 @@ def wrapped_stream_on_event(event: BacktestStreamEvent) -> None:
11911258
if config and config.end_time:
11921259
end_time = config.end_time
11931260

1194-
# Update kwargs if needed by strategy (optional, can be removed if strategies
1195-
# don't need it)
1196-
if start_time:
1197-
kwargs["start_time"] = start_time
1198-
if end_time:
1199-
kwargs["end_time"] = end_time
1200-
1201-
# 注意: initial_cash, commission_rate, timezone, show_progress, history_depth
1202-
# 已经在上方通过优先级逻辑处理过了,这里不需要再覆盖
1203-
1204-
# Risk Config injection handled later
1205-
12061261
# Handle strategy_params explicitly
12071262
if "strategy_params" in kwargs:
12081263
s_params = kwargs.pop("strategy_params")
@@ -1233,6 +1288,10 @@ def wrapped_stream_on_event(event: BacktestStreamEvent) -> None:
12331288
strategy_loader_options=strategy_loader_options,
12341289
)
12351290
strategy_kwargs = dict(kwargs)
1291+
if start_time and _accepts_strategy_kwarg(strategy_input, "start_time"):
1292+
strategy_kwargs["start_time"] = start_time
1293+
if end_time and _accepts_strategy_kwarg(strategy_input, "end_time"):
1294+
strategy_kwargs["end_time"] = end_time
12361295
if (
12371296
symbols is not None
12381297
and "symbols" not in strategy_kwargs
@@ -1242,6 +1301,7 @@ def wrapped_stream_on_event(event: BacktestStreamEvent) -> None:
12421301
strategy_instance = _build_strategy_instance(
12431302
strategy_input,
12441303
strategy_kwargs,
1304+
strict_strategy_params,
12451305
logger,
12461306
initialize,
12471307
on_start,
@@ -1263,9 +1323,16 @@ def wrapped_stream_on_event(event: BacktestStreamEvent) -> None:
12631323
slot_strategy_input, "symbols"
12641324
):
12651325
slot_strategy_kwargs["symbols"] = symbols
1326+
if start_time and _accepts_strategy_kwarg(
1327+
slot_strategy_input, "start_time"
1328+
):
1329+
slot_strategy_kwargs["start_time"] = start_time
1330+
if end_time and _accepts_strategy_kwarg(slot_strategy_input, "end_time"):
1331+
slot_strategy_kwargs["end_time"] = end_time
12661332
slot_strategy_instances[slot_key_str] = _build_strategy_instance(
12671333
slot_strategy_input,
12681334
slot_strategy_kwargs,
1335+
strict_strategy_params,
12691336
logger,
12701337
initialize,
12711338
on_start,
@@ -3185,6 +3252,7 @@ def wrapped_stream_on_event(event: BacktestStreamEvent) -> None:
31853252
slot_strategy_instances[slot_key_str] = _build_strategy_instance(
31863253
slot_strategy_input,
31873254
{},
3255+
False,
31883256
logger,
31893257
None,
31903258
None,

python/akquant/optimize.py

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
提供类似 Backtrader optstrategy 的网格搜索功能.
55
"""
66

7+
import inspect
78
import itertools
89
import json
910
import multiprocessing
@@ -280,6 +281,43 @@ def _assert_parallel_pickleable(
280281
) from e
281282

282283

284+
def _validate_strategy_param_grid_keys(
285+
strategy: Type[Strategy], param_grid: Mapping[str, Sequence[Any]]
286+
) -> None:
287+
"""Validate that param_grid keys can be passed to strategy constructor."""
288+
try:
289+
signature = inspect.signature(strategy.__init__)
290+
except (TypeError, ValueError):
291+
return
292+
293+
supports_var_kwargs = any(
294+
parameter.kind == inspect.Parameter.VAR_KEYWORD
295+
for parameter in signature.parameters.values()
296+
)
297+
if supports_var_kwargs:
298+
return
299+
300+
accepted_names = {
301+
parameter_name
302+
for parameter_name, parameter in signature.parameters.items()
303+
if parameter_name != "self"
304+
and parameter.kind
305+
in {
306+
inspect.Parameter.POSITIONAL_OR_KEYWORD,
307+
inspect.Parameter.KEYWORD_ONLY,
308+
}
309+
}
310+
unknown_keys = sorted(
311+
key for key in param_grid.keys() if str(key) not in accepted_names
312+
)
313+
if unknown_keys:
314+
unknown_keys_text = ", ".join(str(key) for key in unknown_keys)
315+
raise TypeError(
316+
"Unknown strategy constructor parameter(s) in param_grid: "
317+
f"{unknown_keys_text}. Strategy={strategy.__module__}.{strategy.__name__}"
318+
)
319+
320+
283321
def _save_result_to_db(
284322
db_path: str, strategy_name: str, result: OptimizationResult
285323
) -> None:
@@ -352,6 +390,10 @@ def run_grid_search(
352390
:return: 优化结果 (DataFrame 或 List[OptimizationResult])
353391
"""
354392
backtest_kwargs = dict(kwargs)
393+
backtest_kwargs.setdefault("strict_strategy_params", True)
394+
strict_strategy_params = bool(backtest_kwargs.get("strict_strategy_params", False))
395+
if strict_strategy_params:
396+
_validate_strategy_param_grid_keys(strategy, param_grid)
355397
if "execution_mode" in backtest_kwargs:
356398
backtest_kwargs["execution_mode"] = _normalize_execution_mode_for_parallel(
357399
backtest_kwargs["execution_mode"]
@@ -483,6 +525,12 @@ def run_grid_search(
483525
if max_workers is None:
484526
max_workers = multiprocessing.cpu_count() or 1
485527

528+
if max_workers > 1:
529+
print(
530+
"Warning: max_workers>1 uses subprocess workers. Strategy self.log() "
531+
"output may not be visible in the main process console."
532+
)
533+
486534
# 如果只有一个任务或 worker=1,直接运行
487535
# (除非设置了 timeout,需要线程支持,仍走单线程逻辑)
488536
if max_workers == 1 or total_combinations == 1:

0 commit comments

Comments
 (0)