Files
ashare-data/benchmarks/bench_daily.py
T
2026-05-17 15:51:10 +08:00

109 lines
3.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""日线行情多源轮换 + 多 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()