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
@@ -1,5 +1,7 @@
"""Application tests for persisted whole-universe selection runs."""
import threading
import time
from datetime import date
from decimal import Decimal
from typing import Literal
@@ -42,6 +44,30 @@ class FakeReader:
raise AssertionError("the fake evaluator should be used")
class BatchReader(FakeReader):
def __init__(self, source: SelectionExecutionSource) -> None:
super().__init__(source)
self.batch_calls: list[tuple[str, ...]] = []
def load_histories(
self,
stocks: tuple[SelectionStock, ...],
target_trade_date: date,
) -> tuple[StockHistory, ...]:
self.batch_calls.append(tuple(stock.ts_code for stock in stocks))
return tuple(StockHistory(ts_code=stock.ts_code, name=stock.name) for stock in stocks)
class PartialBatchReader(BatchReader):
def load_histories(
self,
stocks: tuple[SelectionStock, ...],
target_trade_date: date,
) -> tuple[StockHistory, ...]:
self.batch_calls.append(tuple(stock.ts_code for stock in stocks))
return ()
class FakeStore:
def __init__(self) -> None:
self.items: list[SelectionRunItem] = []
@@ -114,6 +140,26 @@ class FakeStore:
return None
class BatchStore(FakeStore):
def __init__(self) -> None:
super().__init__()
self.batches: list[tuple[SelectionRunItem, ...]] = []
def record_item(self, run_id: str, item: SelectionRunItem) -> None:
raise AssertionError("the batch path should use record_items")
def record_items(self, run_id: str, items: tuple[SelectionRunItem, ...]) -> None:
assert run_id == "run-1"
batch = tuple(items)
self.batches.append(batch)
self.items.extend(batch)
class FailingBatchStore(BatchStore):
def record_items(self, run_id: str, items: tuple[SelectionRunItem, ...]) -> None:
raise RuntimeError("batch write unavailable")
class FakeEvaluator:
def __init__(self, results: dict[str, SelectionEvaluation]) -> None:
self.results = results
@@ -129,6 +175,29 @@ class RaisingEvaluator:
return SelectionEvaluation(ts_code, target_trade_date, "no_signal")
class ConcurrentHistoryEvaluator:
def __init__(self) -> None:
self.active = 0
self.peak = 0
self._lock = threading.Lock()
def execute(self, ts_code: str, target_trade_date: date) -> SelectionEvaluation:
raise AssertionError("the batch evaluator path should be used")
def execute_history(
self,
history: StockHistory,
target_trade_date: date,
) -> SelectionEvaluation:
with self._lock:
self.active += 1
self.peak = max(self.peak, self.active)
time.sleep(0.02)
with self._lock:
self.active -= 1
return SelectionEvaluation(history.ts_code, target_trade_date, "no_signal")
def _source() -> SelectionExecutionSource:
return SelectionExecutionSource(
market_sync_batch_id="market-run-1",
@@ -233,3 +302,65 @@ def test_execute_isolates_unexpected_single_stock_failure() -> None:
assert store.finished[0:2] == ("run-1", "partial_success")
assert store.finished[2]["evaluated_count"] == 2
assert store.finished[2]["failed_count"] == 1
def test_execute_batches_history_reads_writes_and_limits_evaluation_workers() -> None:
source = SelectionExecutionSource(
market_sync_batch_id="market-run-1",
target_trade_date=TARGET,
target_count=8,
valid_count=8,
coverage=Decimal("1"),
stocks=tuple(SelectionStock(f"{index:06d}.SZ", f"stock-{index}") for index in range(8)),
)
reader = BatchReader(source)
store = BatchStore()
evaluator = ConcurrentHistoryEvaluator()
service = RunZhixingB1(
reader,
store,
evaluator,
max_workers=4,
batch_size=4,
)
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
assert reader.batch_calls == [
("000000.SZ", "000001.SZ", "000002.SZ", "000003.SZ"),
("000004.SZ", "000005.SZ", "000006.SZ", "000007.SZ"),
]
assert [len(batch) for batch in store.batches] == [4, 4]
assert evaluator.peak <= 4
assert evaluator.peak >= 2
assert store.finished is not None
assert store.finished[0:2] == ("run-1", "success")
def test_execute_maps_missing_batch_history_without_a_single_stock_read() -> None:
source = _source()
reader = PartialBatchReader(source)
store = FakeStore()
service = RunZhixingB1(reader, store)
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
assert reader.batch_calls == [("000001.SZ", "600000.SH")]
assert [item.status for item in store.items] == ["missing_target_bar", "missing_target_bar"]
assert store.finished is not None
assert store.finished[0:2] == ("run-1", "failed")
def test_execute_marks_batch_write_failure_as_failed() -> None:
source = _source()
reader = BatchReader(source)
store = FailingBatchStore()
evaluator = ConcurrentHistoryEvaluator()
service = RunZhixingB1(reader, store, evaluator)
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
assert store.finished is not None
assert store.finished[0:2] == ("run-1", "failed")
assert store.finished[2]["error_type"] == "batch_error"
assert store.finished[2]["failed_count"] == 1