simulate_test.py
#!/usr/bin/env python3
"""
模拟测试脚本
用于在没有账号的情况下进行模拟测试
"""
import sys
import os
import logging
import argparse
from datetime import datetime, timedelta
# 添加项目根目录到路径
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
# 配置日志
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
handlers=[
logging.FileHandler("simulate_test.log"),
logging.StreamHandler()
]
)
logger = logging.getLogger("AIQuant.simulate_test")
def main():
"""
主入口函数
"""
logger.info("AI Quant 模拟测试脚本启动")
try:
# 解析命令行参数
parser = argparse.ArgumentParser(
description="AI Quant 模拟测试脚本"
)
parser.add_argument(
"--stock-count",
type=int,
default=10,
help="测试的股票数量(默认:10)"
)
parser.add_argument(
"--data-days",
type=int,
default=365,
help="测试数据的天数(默认:365)"
)
parser.add_argument(
"--frequency",
choices=["d", "h", "m"],
default="d",
help="数据频率(d: 日线,h: 小时线,m: 分钟线,默认:d)"
)
parser.add_argument(
"--save-cache",
action="store_true",
help="保存测试数据到本地缓存"
)
parser.add_argument(
"--load-cache",
action="store_true",
help="从本地缓存加载测试数据"
)
parser.add_argument(
"--clear-cache",
action="store_true",
help="测试完成后清理缓存"
)
args = parser.parse_args()
logger.info(
f"测试参数: 股票数量={args.stock_count}, "
f"数据天数={args.data_days}, 频率={args.frequency}")
# 导入所需模块
from service.data.mock_data import (generate_random_kline_data,
generate_stock_pool,
generate_mock_account_data)
from service.data.data_cache import (save_data, load_data,
data_exists, clear_all_cache)
from app.strategy.single_stock import SingleStockAnalyzer
# 生成股票池
logger.info("生成股票池...")
stock_pool = generate_stock_pool(args.stock_count)
logger.info(f"股票池: {stock_pool}")
# 生成模拟账户数据
logger.info("生成模拟账户数据...")
account_data = generate_mock_account_data()
logger.info(f"账户余额: {account_data['balance']}")
# 初始化分析器
logger.info("初始化单票量化分析器...")
analyzer = SingleStockAnalyzer()
# 测试结果
test_results = []
# 遍历股票池,进行分析
for symbol in stock_pool:
logger.info(f"\n分析股票: {symbol}")
if args.load_cache and data_exists(symbol, args.frequency):
# 从缓存加载数据
logger.info("从缓存加载数据...")
data = load_data(symbol, args.frequency)
else:
# 生成模拟数据
logger.info("生成模拟数据...")
data = generate_random_kline_data(
symbol,
start_date="2023-01-01",
end_date=(datetime(2023, 1, 1) +
timedelta(days=args.data_days - 1)
).strftime("%Y-%m-%d"),
frequency=args.frequency
)
if args.save_cache:
logger.info("保存数据到缓存...")
save_data(symbol, data, args.frequency)
# 分析股票
result = analyzer.analyze(symbol, data)
logger.info(f"分析结果: {result}")
test_results.append(result)
# 统计测试结果
logger.info("\n测试结果统计:")
logger.info(f"总股票数量: {len(test_results)}")
# 按等级统计
rating_counts = {
"S": 0,
"A": 0,
"B": 0,
"C": 0
}
for result in test_results:
if result.rating in rating_counts:
rating_counts[result.rating] += 1
logger.info("标的等级分布:")
for rating, count in rating_counts.items():
logger.info(
f"{rating}级: {count}只 ({count/len(test_results)*100:.1f}%)")
# 按交易信号统计
signal_counts = {
"买入": 0,
"持有": 0,
"卖出": 0
}
for result in test_results:
if result.signal in signal_counts:
signal_counts[result.signal] += 1
logger.info("交易信号分布:")
for signal, count in signal_counts.items():
logger.info(
f"{signal}: {count}只 ({count/len(test_results)*100:.1f}%)")
# 计算平均评分
avg_score = sum(
result.score for result in test_results) / len(test_results)
logger.info(f"平均评分: {avg_score:.2f}")
# 清理缓存
if args.clear_cache:
logger.info("清理数据缓存...")
clear_all_cache()
logger.info("模拟测试完成")
return 0
except Exception as e:
logger.error(f"模拟测试过程中出错: {e}")
logger.exception("详细错误信息:")
return 1
if __name__ == "__main__":
try:
sys.exit(main())
except KeyboardInterrupt:
logger.info("程序被用户中断")
sys.exit(0)