perf(selection): 优化选股执行性能

This commit is contained in:
yuxuanhui
2026-08-12 09:45:16 +08:00
parent dd04933d63
commit 8963c067b3
22 changed files with 1333 additions and 120 deletions
@@ -3,12 +3,17 @@
from __future__ import annotations
import logging
import time
from collections.abc import Callable, Sequence
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from datetime import date
from typing import Literal, Protocol
from typing import Literal, Protocol, cast
from ..domain.models import SelectionEvaluation
from ..domain.models import SelectionEvaluation, StockHistory
from ..domain.runs import (
BatchSelectionRunStore,
BatchSelectionUniverseReader,
SelectionExecutionSource,
SelectionRerunRequired,
SelectionResultQuery,
@@ -17,6 +22,7 @@ from ..domain.runs import (
SelectionRunItem,
SelectionRunStatus,
SelectionRunStore,
SelectionStock,
SelectionUniverseReader,
)
from .evaluate import EvaluateZhixingB1
@@ -48,12 +54,21 @@ class RunZhixingB1:
reader: SelectionUniverseReader,
store: SelectionRunStore,
evaluator: SelectionEvaluator | None = None,
*,
max_workers: int = 4,
batch_size: int = 200,
) -> None:
"""Inject storage ports and optionally a test evaluator."""
"""Inject storage ports and configure bounded chunk execution."""
if max_workers < 1:
raise ValueError("max_workers must be at least 1")
if batch_size < 1:
raise ValueError("batch_size must be at least 1")
self.reader = reader
self.store = store
self.evaluator = evaluator or EvaluateZhixingB1(reader)
self.max_workers = max_workers
self.batch_size = batch_size
def prepare(
self,
@@ -81,35 +96,57 @@ class RunZhixingB1:
returns so the UI never mistakes a lost worker exception for success.
"""
stocks = _unique_stocks(prepared.source.stocks)
evaluated_count = 0
selected_stock_count = 0
signal_count = 0
failed_count = 0
history_rows = 0
batch_count = _chunk_count(len(stocks), self.batch_size)
read_seconds = 0.0
evaluate_seconds = 0.0
persist_seconds = 0.0
try:
for stock in prepared.source.stocks:
try:
evaluation = self.evaluator.execute(
stock.ts_code,
with ThreadPoolExecutor(max_workers=self.max_workers) as executor:
for batch_stocks in _chunks(stocks, self.batch_size):
read_started = time.perf_counter()
histories = self._load_histories(
batch_stocks,
prepared.source.target_trade_date,
)
except Exception as exc: # noqa: BLE001 - isolate one stock from the batch
logger.exception(
"selection_item_failed run_id=%s ts_code=%s",
prepared.run.id,
stock.ts_code,
read_seconds += time.perf_counter() - read_started
history_rows += sum(
len(history.bars) for history in histories if history is not None
)
evaluation = SelectionEvaluation(
ts_code=stock.ts_code,
target_trade_date=prepared.source.target_trade_date,
status="data_error",
reason=_safe_item_error(exc),
evaluate_started = time.perf_counter()
items = tuple(
_to_item(
stock.ts_code,
stock.name,
evaluation,
)
for stock, evaluation in zip(
batch_stocks,
executor.map(
self._evaluate_stock,
batch_stocks,
histories,
[prepared.source.target_trade_date] * len(batch_stocks),
),
strict=True,
)
)
item = _to_item(stock.ts_code, stock.name, evaluation)
self.store.record_item(prepared.run.id, item)
evaluated_count += 1
selected_stock_count += evaluation.status == "selected"
signal_count += len(evaluation.signals)
failed_count += evaluation.status in _FAILURE_STATUSES
evaluate_seconds += time.perf_counter() - evaluate_started
evaluated_count += len(items)
selected_stock_count += sum(item.status == "selected" for item in items)
signal_count += sum(item.signal_count for item in items)
failed_count += sum(item.status in _FAILURE_STATUSES for item in items)
persist_started = time.perf_counter()
self._record_items(prepared.run.id, items)
persist_seconds += time.perf_counter() - persist_started
status = _run_status(evaluated_count, failed_count)
self.store.finish_run(
@@ -121,7 +158,12 @@ class RunZhixingB1:
failed_count=failed_count,
)
except Exception as exc: # noqa: BLE001 - worker boundary must persist failure state
logger.exception("selection_run_failed run_id=%s", prepared.run.id)
logger.error(
"selection_run_failed run_id=%s error_type=%s reason=%s",
prepared.run.id,
exc.__class__.__name__,
_safe_item_error(exc),
)
try:
self.store.finish_run(
prepared.run.id,
@@ -134,7 +176,94 @@ class RunZhixingB1:
error_message=str(exc),
)
except Exception: # noqa: BLE001 - preserve the original worker failure
logger.exception("selection_run_failure_persist_failed run_id=%s", prepared.run.id)
logger.error(
"selection_run_failure_persist_failed run_id=%s",
prepared.run.id,
)
finally:
logger.info(
"selection_run_summary run_id=%s stock_count=%d history_rows=%d "
"batch_count=%d worker_count=%d read_seconds=%.3f "
"evaluate_seconds=%.3f persist_seconds=%.3f",
prepared.run.id,
len(stocks),
history_rows,
batch_count,
self.max_workers,
read_seconds,
evaluate_seconds,
persist_seconds,
)
def _load_histories(
self,
stocks: Sequence[SelectionStock],
target_trade_date: date,
) -> tuple[StockHistory | None, ...]:
"""Load one chunk when the reader supports it, with old-path fallback."""
typed_stocks = tuple(stocks)
loader = getattr(self.reader, "load_histories", None)
if callable(loader):
batch_reader = cast(BatchSelectionUniverseReader, self.reader)
loaded = batch_reader.load_histories(typed_stocks, target_trade_date)
histories_by_code = {history.ts_code: history for history in loaded}
return tuple(
histories_by_code.get(
stock.ts_code,
StockHistory(ts_code=stock.ts_code, name=stock.name),
)
for stock in typed_stocks
)
if isinstance(self.evaluator, EvaluateZhixingB1):
return tuple(
self.reader.load_history(stock.ts_code, target_trade_date) for stock in typed_stocks
)
return (None,) * len(typed_stocks)
def _evaluate_stock(
self,
stock: SelectionStock,
history: StockHistory | None,
target_trade_date: date,
) -> SelectionEvaluation:
"""Evaluate one stock inside a worker and isolate its exception."""
ts_code = stock.ts_code
try:
execute_history: Callable[[StockHistory, date], SelectionEvaluation] | None = getattr(
self.evaluator,
"execute_history",
None,
)
if history is not None and execute_history is not None:
return execute_history(history, target_trade_date)
return self.evaluator.execute(ts_code, target_trade_date)
except Exception as exc: # noqa: BLE001 - isolate one stock from the batch
logger.warning(
"selection_item_failed ts_code=%s error_type=%s reason=%s",
ts_code,
exc.__class__.__name__,
_safe_item_error(exc),
)
return SelectionEvaluation(
ts_code=ts_code,
target_trade_date=target_trade_date,
status="data_error",
reason=_safe_item_error(exc),
)
def _record_items(self, run_id: str, items: Sequence[SelectionRunItem]) -> None:
"""Use batch persistence while retaining the old single-item seam."""
record_items = getattr(self.store, "record_items", None)
if callable(record_items):
batch_store = cast(BatchSelectionRunStore, self.store)
batch_store.record_items(run_id, tuple(items))
return
for item in items:
self.store.record_item(run_id, item)
def get_run(
self,
@@ -187,6 +316,35 @@ def _safe_item_error(error: Exception) -> str:
return " ".join(str(error).split())[:500] or error.__class__.__name__
def _chunks(
values: Sequence[SelectionStock],
size: int,
) -> tuple[tuple[SelectionStock, ...], ...]:
"""Split a stable stock sequence into bounded immutable chunks."""
return tuple(tuple(values[index : index + size]) for index in range(0, len(values), size))
def _chunk_count(value_count: int, size: int) -> int:
"""Return the number of chunks without materializing empty chunks."""
return (value_count + size - 1) // size
def _unique_stocks(stocks: Sequence[SelectionStock]) -> tuple[SelectionStock, ...]:
"""Keep the first source row for each stock so it is evaluated once."""
seen: set[str] = set()
unique: list[SelectionStock] = []
for stock in stocks:
ts_code = stock.ts_code
if ts_code in seen:
continue
seen.add(ts_code)
unique.append(stock)
return tuple(unique)
__all__ = [
"PreparedSelectionRun",
"RunZhixingB1",