SHA256
109 lines
3.7 KiB
Python
109 lines
3.7 KiB
Python
"""日线行情多源轮换 + 多 worker 并发压测
|
||
|
||
用法:
|
||
python -m benchmarks.bench_daily --codes 200 --days 30 --workers 1,2,4,8
|
||
|
||
行为:
|
||
- 从 stock_info 随机/顺序取若干只股票(必须先 --stock-info);
|
||
- 估算它们最近 N 天的真实缺口;
|
||
- 顺次以指定的 workers 配置跑一遍真实抓取,记录耗时、成功率、无数据数。
|
||
- 不会清空已有数据,跑完即可被认为是常规增量。
|
||
|
||
说明:
|
||
- 此脚本会真实命中外部数据源(baostock / 新浪 / 腾讯 / 东方财富),
|
||
请在配置允许的窗口内运行,避免被限流。
|
||
- 脚本不修改 config.yaml;要切换 workers 是通过传参,并临时覆盖 fetch.workers。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import random
|
||
import time
|
||
from datetime import datetime, timedelta
|
||
|
||
from src.config import load_config, get_fetch_config
|
||
from src.db import init_db, get_stock_codes
|
||
from src.fetchers.daily import fetch_daily
|
||
from src.log import get_logger
|
||
|
||
_logger = get_logger("bench")
|
||
|
||
|
||
def _parse_workers(arg: str) -> list[int]:
|
||
return [max(1, int(x.strip())) for x in arg.split(",") if x.strip()]
|
||
|
||
|
||
def _pick_codes(n: int, seed: int | None) -> list[str]:
|
||
all_codes = get_stock_codes()
|
||
if not all_codes:
|
||
raise SystemExit("stock_info 为空,请先运行:python -m src.main --stock-info")
|
||
if n >= len(all_codes):
|
||
return all_codes
|
||
rng = random.Random(seed)
|
||
return rng.sample(all_codes, n)
|
||
|
||
|
||
def main():
|
||
parser = argparse.ArgumentParser(description="日线抓取并发度压测")
|
||
parser.add_argument("--codes", type=int, default=200, help="参与压测的股票数(默认 200)")
|
||
parser.add_argument("--days", type=int, default=30, help="时间窗口天数(默认近 30 天)")
|
||
parser.add_argument(
|
||
"--workers", type=str, default="1,2,4,8",
|
||
help="并发档位列表,逗号分隔(默认 1,2,4,8)",
|
||
)
|
||
parser.add_argument(
|
||
"--source", type=str, default="all",
|
||
choices=["baostock", "sina", "tencent", "eastmoney", "all"],
|
||
help="数据源(默认 all 多源轮换)",
|
||
)
|
||
parser.add_argument("--seed", type=int, default=42, help="随机种子,固定后多次跑可重复")
|
||
args = parser.parse_args()
|
||
|
||
load_config()
|
||
init_db()
|
||
|
||
workers_list = _parse_workers(args.workers)
|
||
end_date = datetime.now().strftime("%Y%m%d")
|
||
start_date = (datetime.now() - timedelta(days=args.days)).strftime("%Y%m%d")
|
||
|
||
_logger.info(
|
||
"压测计划: codes=%d days=%d source=%s workers=%s seed=%s",
|
||
args.codes, args.days, args.source, workers_list, args.seed,
|
||
)
|
||
_logger.info("窗口: %s ~ %s", start_date, end_date)
|
||
|
||
sample_codes = _pick_codes(args.codes, args.seed)
|
||
_logger.info("已选股票样本(前10): %s ...", sample_codes[:10])
|
||
|
||
cfg = get_fetch_config()
|
||
|
||
results = []
|
||
for w in workers_list:
|
||
_logger.info("==== workers=%d 开始 ====", w)
|
||
cfg["workers"] = w
|
||
|
||
from src.fetchers import daily as daily_mod
|
||
|
||
original_get_stock_codes = daily_mod.get_stock_codes
|
||
daily_mod.get_stock_codes = lambda: sample_codes
|
||
try:
|
||
t0 = time.time()
|
||
fetch_daily(start_date=start_date, end_date=end_date, source=args.source)
|
||
elapsed = time.time() - t0
|
||
finally:
|
||
daily_mod.get_stock_codes = original_get_stock_codes
|
||
|
||
_logger.info("==== workers=%d 完成: %.1fs ====", w, elapsed)
|
||
results.append((w, elapsed))
|
||
|
||
_logger.info("==== 汇总 ====")
|
||
base = results[0][1] if results else 0.0
|
||
for w, sec in results:
|
||
speedup = base / sec if sec > 0 else 0.0
|
||
_logger.info(" workers=%d %.1fs speedup=%.2fx", w, sec, speedup)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|