perf(selection): 优化选股执行性能
This commit is contained in:
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user