# Python量化实战:构建本地沪深股票日线数据库的完整方案
每次想尝试一个新的量化策略,或者验证一个技术指标的有效性,最头疼的往往不是代码怎么写,而是数据从哪里来。网上的数据源要么收费昂贵,要么接口不稳定,要么数据格式混乱。几年前我开始接触Python量化时,也在这个问题上卡了很久,直到发现了pytdx这个宝藏库——它让我能够直接从通达信行情服务器获取数据,而且是免费的。
今天要分享的,不仅仅是如何用pytdx下载单只股票的数据,而是如何构建一个完整的本地股票数据库。我会带你从零开始,实现自动识别沪深市场、批量下载数百只股票、处理分页获取限制、建立规范的存储体系,并且加入完善的异常处理机制。这套方案我已经在实际项目中运行了两年多,稳定下载了A股全市场的数据,支撑了多个策略的回测需求。
如果你正在为量化数据源发愁,或者想要建立一个属于自己的、可随时调用的本地数据仓库,那么这篇文章正是为你准备的。我们不需要复杂的金融数据API,不需要昂贵的数据库服务,只需要Python和一些耐心,就能搭建起一个可靠的数据基础设施。
## 1. 环境搭建与核心工具解析
在开始编写代码之前,我们需要先理解整个技术栈的构成。很多人一提到Python量化,就想到pandas、numpy这些数据分析库,这没错,但数据获取才是整个链条的起点。没有高质量、稳定的数据源,再精妙的策略也只是空中楼阁。
### 1.1 为什么选择pytdx?
在众多的数据获取方案中,我最终选择了pytdx,主要是基于以下几个实际考量:
- **数据质量可靠**:数据直接来自通达信行情服务器,与主流交易软件同源,保证了数据的准确性和一致性
- **完全免费**:不需要注册账号,不需要申请API密钥,没有调用次数限制(当然要合理使用)
- **延迟较低**:相比一些第三方数据接口,pytdx的延迟通常更低,对于日线数据来说几乎是实时的
- **社区活跃**:GitHub上有持续的维护和更新,遇到问题比较容易找到解决方案
不过需要明确的是,pytdx获取的是行情数据,主要用于个人学习和研究。如果是商业用途,请务必了解相关的合规要求。
### 1.2 基础环境配置
让我们从最基础的安装开始。我建议使用conda创建独立的Python环境,避免包版本冲突:
```bash
# 创建新的conda环境
conda create -n quant_data python=3.9
conda activate quant_data
# 安装核心依赖
pip install pytdx pandas numpy
```
如果你不用conda,用venv或者直接pip安装也可以。这里我选择Python 3.9是因为它在稳定性和新特性之间取得了很好的平衡,而且大多数量化相关的库都对3.9有很好的支持。
安装完成后,可以通过一个简单的测试来验证环境是否正常:
```python
import pytdx
import pandas as pd
import numpy as np
print(f"pytdx版本: {pytdx.__version__}")
print(f"pandas版本: {pd.__version__}")
print(f"numpy版本: {np.__version__}")
```
> 提示:在实际项目中,我通常会创建一个requirements.txt文件来管理依赖版本,这样可以确保在不同的机器上环境一致。特别是pytdx,不同版本可能会有API的变化。
### 1.3 理解通达信的数据结构
在开始编码之前,有必要了解一下通达信的数据组织方式。这对于后续处理各种边界情况很有帮助:
**市场分类规则**
- 上海市场(代码以6开头):股票、指数、基金等
- 深圳市场(代码以0、3开头):主板、创业板、中小板等
**K线类型代码**
```python
KLINE_TYPE = {
0: "5分钟K线",
1: "15分钟K线",
2: "30分钟K线",
3: "1小时K线",
4: "日K线",
5: "周K线",
6: "月K线",
7: "1分钟K线",
8: "日K线", # 分笔成交
9: "日K线" # 我们主要使用的日线数据
}
```
**数据字段说明**
每次获取的数据包含7个字段,按顺序分别是:
1. `datetime` - 时间戳(整数格式)
2. `open` - 开盘价
3. `high` - 最高价
4. `low` - 最低价
5. `close` - 收盘价
6. `amount` - 成交额(元)
7. `volume` - 成交量(手)
理解这些基础概念后,我们就可以开始构建数据下载的核心逻辑了。
## 2. 单只股票数据下载的深度实现
很多人教程只教你怎么下载数据,但很少告诉你可能会遇到什么问题以及如何解决。在这一节,我会分享一个经过实战检验的完整下载函数,它包含了连接管理、错误重试、数据验证等关键特性。
### 2.1 基础下载函数的设计
先来看一个最基础的实现,我会逐行解释每个部分的设计考虑:
```python
import os
import time
from datetime import datetime
from pytdx.hq import TdxHq_API
import pandas as pd
def download_single_stock(stock_code, save_dir="./data", retry_times=3):
"""
下载单只股票的日线数据
参数:
stock_code: 股票代码,如'000001'或'600000'
save_dir: 数据保存目录
retry_times: 连接失败时的重试次数
返回:
bool: 下载是否成功
str: 保存的文件路径或错误信息
"""
# 创建保存目录
os.makedirs(save_dir, exist_ok=True)
# 确定市场类型
if stock_code.startswith('6'):
market = 1 # 上海市场
prefix = 'sh'
elif stock_code.startswith('0') or stock_code.startswith('3'):
market = 0 # 深圳市场
prefix = 'sz'
else:
return False, f"无法识别的股票代码格式: {stock_code}"
# 构建文件名
filename = f"{prefix}{stock_code}.csv"
filepath = os.path.join(save_dir, filename)
# 检查文件是否已存在(避免重复下载)
if os.path.exists(filepath):
print(f"文件已存在: {filename}")
return True, filepath
# 连接服务器并获取数据
api = TdxHq_API()
for attempt in range(retry_times):
try:
# 尝试连接服务器
with api.connect('124.71.163.106', 7709):
print(f"正在下载 {stock_code} 的数据...")
# 获取日线数据
data = api.get_security_bars(
category=9, # 9表示日线
market=market, # 市场类型
code=stock_code, # 股票代码
start=0, # 起始位置
count=800 # 获取数量
)
if not data:
return False, f"未获取到 {stock_code} 的数据"
# 转换为DataFrame
df = pd.DataFrame(data, columns=[
'datetime', 'open', 'high', 'low',
'close', 'amount', 'volume'
])
# 转换时间格式
df['datetime'] = pd.to_datetime(df['datetime'])
# 按时间排序(确保数据顺序正确)
df = df.sort_values('datetime')
# 保存到CSV
df.to_csv(filepath, index=False)
print(f"成功保存: {filename},共 {len(df)} 条记录")
return True, filepath
except Exception as e:
print(f"第 {attempt + 1} 次尝试失败: {str(e)}")
if attempt < retry_times - 1:
time.sleep(2) # 等待2秒后重试
else:
return False, f"下载 {stock_code} 失败: {str(e)}"
return False, "未知错误"
```
这个函数有几个关键设计点:
1. **自动识别市场**:根据股票代码前缀判断是上海还是深圳市场
2. **文件存在检查**:避免重复下载,节省时间和资源
3. **重试机制**:网络不稳定时的容错处理
4. **数据排序**:确保数据按时间顺序排列
5. **完整的错误处理**:提供清晰的错误信息
### 2.2 处理800条限制的智能分页
pytdx有一个重要的限制:单次最多只能获取800条日线数据。对于历史较长的股票,我们需要分页获取。下面是一个智能分页的实现:
```python
def download_stock_with_pagination(stock_code, save_dir="./data", max_pages=10):
"""
分页下载股票历史数据,突破800条限制
参数:
stock_code: 股票代码
save_dir: 保存目录
max_pages: 最大分页数(防止无限循环)
返回:
DataFrame: 合并后的数据
"""
# 确定市场类型
if stock_code.startswith('6'):
market = 1
prefix = 'sh'
else:
market = 0
prefix = 'sz'
all_data = []
api = TdxHq_API()
try:
with api.connect('124.71.163.106', 7709):
for page in range(max_pages):
start_index = page * 800
print(f"获取 {stock_code} 第 {page + 1} 页数据,起始位置: {start_index}")
data = api.get_security_bars(
category=9,
market=market,
code=stock_code,
start=start_index,
count=800
)
if not data:
print(f"第 {page + 1} 页无数据,停止获取")
break
# 添加到总数据列表
all_data.extend(data)
# 如果获取的数据少于800条,说明已经获取完所有数据
if len(data) < 800:
print(f"获取到最后一页,共 {len(all_data)} 条记录")
break
# 避免请求过快
time.sleep(0.5)
except Exception as e:
print(f"下载 {stock_code} 时出错: {str(e)}")
return None
if not all_data:
return None
# 转换为DataFrame并处理
df = pd.DataFrame(all_data, columns=[
'datetime', 'open', 'high', 'low',
'close', 'amount', 'volume'
])
# 去重并排序
df['datetime'] = pd.to_datetime(df['datetime'])
df = df.drop_duplicates('datetime')
df = df.sort_values('datetime')
# 保存文件
filename = f"{prefix}{stock_code}.csv"
filepath = os.path.join(save_dir, filename)
df.to_csv(filepath, index=False)
print(f"已保存 {len(df)} 条记录到 {filename}")
return df
```
这个分页函数有几个值得注意的地方:
- **自动判断数据结束**:当获取的数据少于800条时,认为已经获取完所有历史数据
- **数据去重**:分页获取时可能有重复数据,需要去重处理
- **请求间隔**:添加了0.5秒的间隔,避免对服务器造成过大压力
### 2.3 数据质量验证与清洗
下载到的数据并不总是完美的,我们需要进行质量检查。下面是一些常见的数据问题及其处理方法:
```python
def validate_and_clean_data(df, stock_code):
"""
验证和清洗股票数据
参数:
df: 原始数据DataFrame
stock_code: 股票代码(用于日志输出)
返回:
DataFrame: 清洗后的数据
"""
if df is None or df.empty:
print(f"{stock_code}: 数据为空")
return None
original_count = len(df)
# 1. 检查必要字段是否存在
required_columns = ['datetime', 'open', 'high', 'low', 'close', 'volume']
missing_columns = [col for col in required_columns if col not in df.columns]
if missing_columns:
print(f"{stock_code}: 缺少必要字段 {missing_columns}")
return None
# 2. 去除重复的日期
df = df.drop_duplicates(subset=['datetime'])
# 3. 按日期排序
df = df.sort_values('datetime')
# 4. 检查价格数据的合理性
# 价格应该为正数
price_columns = ['open', 'high', 'low', 'close']
for col in price_columns:
invalid_mask = df[col] <= 0
if invalid_mask.any():
print(f"{stock_code}: 发现 {invalid_mask.sum()} 条{col}价格异常记录")
# 可以选择删除或标记异常记录
df = df[~invalid_mask]
# 5. 检查高价是否>=低价
invalid_high_low = df['high'] < df['low']
if invalid_high_low.any():
print(f"{stock_code}: 发现 {invalid_high_low.sum()} 条高价低于低价的记录")
df = df[~invalid_high_low]
# 6. 检查收盘价是否在最高最低价之间
invalid_close = (df['close'] > df['high']) | (df['close'] < df['low'])
if invalid_close.any():
print(f"{stock_code}: 发现 {invalid_close.sum()} 条收盘价超出范围的记录")
df = df[~invalid_close]
# 7. 处理缺失值
df = df.dropna()
cleaned_count = len(df)
if cleaned_count < original_count:
print(f"{stock_code}: 清洗后剩余 {cleaned_count} 条记录(移除 {original_count - cleaned_count} 条)")
return df
```
数据验证是量化分析中至关重要但常被忽视的一环。我曾经因为没做充分的数据清洗,导致回测结果出现严重偏差。上面的验证函数涵盖了最常见的数据问题,你可以根据实际需求进行调整。
## 3. 批量下载与数据库构建
单只股票的下载只是开始,真正的价值在于构建完整的股票数据库。这一节我会分享如何高效、稳定地批量下载全市场数据。
### 3.1 股票代码列表的获取与管理
首先我们需要获取沪深两市的股票代码列表。这里有几个实用的方法:
**方法一:从本地文件读取**
如果你已经有股票代码列表,可以保存为文本文件:
```python
def load_stock_codes_from_file(filepath):
"""
从文本文件加载股票代码列表
文件格式:每行一个股票代码
示例:
000001
000002
600000
600036
"""
with open(filepath, 'r', encoding='utf-8') as f:
codes = [line.strip() for line in f if line.strip()]
return codes
# 使用示例
stock_codes = load_stock_codes_from_file('stock_codes.txt')
```
**方法二:动态获取当前上市股票**
我们可以从一些公开数据源获取当前上市的股票列表。这里提供一个简单的方法:
```python
import requests
from io import StringIO
def get_current_stock_list():
"""
获取当前上市的股票列表(示例方法)
注意:实际使用时需要确保数据源的稳定性和合规性
"""
try:
# 这里只是一个示例,实际需要替换为可靠的数据源
url = "http://example.com/stock_list.csv" # 替换为实际URL
response = requests.get(url, timeout=10)
if response.status_code == 200:
# 假设CSV文件包含'code'列
df = pd.read_csv(StringIO(response.text))
codes = df['code'].astype(str).str.zfill(6).tolist()
return codes
else:
print("获取股票列表失败")
return []
except Exception as e:
print(f"获取股票列表时出错: {str(e)}")
return []
```
**方法三:手动维护重点股票池**
对于大多数量化策略来说,并不需要全市场所有股票。我们可以维护一个重点关注的股票池:
```python
# 常见指数成分股(示例)
index_constituents = {
'沪深300': [
'000001', '000002', '000063', '000066', '000069',
'000100', '000157', '000166', '000333', '000338',
'000425', '000538', '000568', '000625', '000651',
'000656', '000661', '000671', '000703', '000708',
# ... 更多代码
],
'中证500': [
'000009', '000012', '000021', '000027', '000028',
'000031', '000039', '000040', '000048', '000050',
'000059', '000060', '000061', '000062', '000070',
'000078', '000088', '000089', '000090', '000096',
# ... 更多代码
]
}
def get_stock_pool(pool_name='沪深300'):
"""获取指定股票池的代码列表"""
return index_constituents.get(pool_name, [])
```
### 3.2 批量下载的完整实现
有了股票代码列表,我们就可以实现批量下载了。这里的关键是处理好并发控制和错误恢复:
```python
import concurrent.futures
import logging
from tqdm import tqdm # 进度条库,需要安装:pip install tqdm
def setup_logging(log_dir="./logs"):
"""设置日志记录"""
os.makedirs(log_dir, exist_ok=True)
log_file = os.path.join(log_dir, f"download_{datetime.now().strftime('%Y%m%d_%H%M%S')}.log")
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s - %(levelname)s - %(message)s',
handlers=[
logging.FileHandler(log_file, encoding='utf-8'),
logging.StreamHandler()
]
)
return logging.getLogger(__name__)
def download_stock_worker(stock_code, save_dir, logger):
"""单个股票下载的工作函数"""
try:
success, result = download_single_stock(stock_code, save_dir)
if success:
logger.info(f"成功下载: {stock_code}")
return stock_code, True, result
else:
logger.error(f"下载失败: {stock_code} - {result}")
return stock_code, False, result
except Exception as e:
logger.error(f"处理 {stock_code} 时发生异常: {str(e)}")
return stock_code, False, str(e)
def batch_download_stocks(stock_codes, save_dir="./data", max_workers=5, retry_failed=True):
"""
批量下载股票数据
参数:
stock_codes: 股票代码列表
save_dir: 数据保存目录
max_workers: 最大并发数
retry_failed: 是否重试失败的下载
返回:
dict: 下载结果统计
"""
# 设置日志
logger = setup_logging()
# 创建保存目录
os.makedirs(save_dir, exist_ok=True)
# 记录已下载的股票
downloaded = []
failed = []
# 使用线程池并发下载
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
# 提交所有任务
future_to_code = {
executor.submit(download_stock_worker, code, save_dir, logger): code
for code in stock_codes
}
# 使用tqdm显示进度
with tqdm(total=len(stock_codes), desc="下载进度") as pbar:
for future in concurrent.futures.as_completed(future_to_code):
code = future_to_code[future]
try:
stock_code, success, result = future.result()
if success:
downloaded.append(stock_code)
else:
failed.append((stock_code, result))
except Exception as e:
logger.error(f"任务异常: {code} - {str(e)}")
failed.append((code, str(e)))
pbar.update(1)
pbar.set_postfix({
'成功': len(downloaded),
'失败': len(failed)
})
# 如果需要,重试失败的下载
if retry_failed and failed:
logger.info(f"开始重试 {len(failed)} 个失败的下载...")
retry_results = []
for code, error in failed[:10]: # 只重试前10个,避免无限重试
logger.info(f"重试: {code}")
success, result = download_single_stock(code, save_dir)
retry_results.append((code, success, result))
time.sleep(1) # 重试时增加间隔
# 更新结果
for code, success, result in retry_results:
if success:
downloaded.append(code)
failed = [(c, e) for c, e in failed if c != code]
# 生成统计报告
stats = {
'total': len(stock_codes),
'success': len(downloaded),
'failed': len(failed),
'success_rate': len(downloaded) / len(stock_codes) * 100 if stock_codes else 0,
'failed_codes': failed
}
logger.info(f"批量下载完成: 总计{stats['total']}只,成功{stats['success']}只,"
f"失败{stats['failed']}只,成功率{stats['success_rate']:.2f}%")
# 保存统计结果
stats_file = os.path.join(save_dir, "download_stats.json")
import json
with open(stats_file, 'w', encoding='utf-8') as f:
json.dump(stats, f, ensure_ascii=False, indent=2)
return stats
```
这个批量下载函数有几个重要的特性:
1. **并发控制**:使用线程池提高下载效率,同时通过`max_workers`参数控制并发数
2. **进度显示**:使用tqdm库显示实时进度
3. **详细日志**:记录下载过程中的所有事件,便于问题排查
4. **错误重试**:对失败的下载进行自动重试
5. **结果统计**:生成详细的下载统计报告
### 3.3 数据库的维护与更新
构建数据库不是一次性的工作,我们需要定期更新数据。下面是一个增量更新的方案:
```python
def update_stock_data(stock_code, save_dir="./data", days_to_update=5):
"""
增量更新股票数据
参数:
stock_code: 股票代码
save_dir: 数据目录
days_to_update: 需要更新的天数
返回:
bool: 更新是否成功
"""
# 确定市场前缀
if stock_code.startswith('6'):
prefix = 'sh'
else:
prefix = 'sz'
filepath = os.path.join(save_dir, f"{prefix}{stock_code}.csv")
# 如果文件不存在,直接下载全部数据
if not os.path.exists(filepath):
print(f"{stock_code}: 数据文件不存在,执行完整下载")
return download_single_stock(stock_code, save_dir)[0]
# 读取现有数据
try:
existing_df = pd.read_csv(filepath)
existing_df['datetime'] = pd.to_datetime(existing_df['datetime'])
# 获取最新日期
latest_date = existing_df['datetime'].max()
print(f"{stock_code}: 现有数据最新日期为 {latest_date.date()}")
# 计算需要更新的起始位置
# 这里简化处理,实际应该根据交易日历计算
# 获取新数据
market = 1 if stock_code.startswith('6') else 0
api = TdxHq_API()
with api.connect('124.71.163.106', 7709):
# 获取最新数据
new_data = api.get_security_bars(9, market, stock_code, 0, days_to_update)
if not new_data:
print(f"{stock_code}: 未获取到新数据")
return True
# 转换为DataFrame
new_df = pd.DataFrame(new_data, columns=[
'datetime', 'open', 'high', 'low',
'close', 'amount', 'volume'
])
new_df['datetime'] = pd.to_datetime(new_df['datetime'])
# 过滤掉已存在的数据
existing_dates = set(existing_df['datetime'])
new_df = new_df[~new_df['datetime'].isin(existing_dates)]
if new_df.empty:
print(f"{stock_code}: 没有需要更新的新数据")
return True
# 合并数据
updated_df = pd.concat([existing_df, new_df], ignore_index=True)
updated_df = updated_df.sort_values('datetime')
updated_df = updated_df.drop_duplicates('datetime')
# 保存更新后的数据
updated_df.to_csv(filepath, index=False)
print(f"{stock_code}: 成功更新 {len(new_df)} 条新数据")
return True
except Exception as e:
print(f"{stock_code}: 更新数据时出错: {str(e)}")
return False
def batch_update_stocks(stock_codes, save_dir="./data", max_workers=3):
"""
批量更新股票数据
"""
logger = setup_logging()
updated = []
failed = []
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
future_to_code = {
executor.submit(update_stock_data, code, save_dir): code
for code in stock_codes
}
with tqdm(total=len(stock_codes), desc="更新进度") as pbar:
for future in concurrent.futures.as_completed(future_to_code):
code = future_to_code[future]
try:
success = future.result()
if success:
updated.append(code)
logger.info(f"更新成功: {code}")
else:
failed.append(code)
logger.error(f"更新失败: {code}")
except Exception as e:
logger.error(f"更新异常: {code} - {str(e)}")
failed.append(code)
pbar.update(1)
return {
'total': len(stock_codes),
'updated': len(updated),
'failed': len(failed),
'failed_codes': failed
}
```
增量更新可以大大减少数据下载的时间,特别是对于历史数据已经比较完整的情况。在实际使用中,你可以设置一个定时任务,每天收盘后自动更新数据。
## 4. 高级应用与性能优化
当数据量变大时,我们需要考虑存储效率、查询性能等问题。这一节分享一些高级技巧和优化方案。
### 4.1 数据存储格式优化
CSV文件虽然简单易用,但当数据量很大时,读写效率会成为瓶颈。下面介绍几种优化方案:
**方案一:使用Parquet格式**
```python
import pyarrow as pa
import pyarrow.parquet as pq
def save_as_parquet(df, filepath):
"""
将DataFrame保存为Parquet格式
Parquet格式的优势:
1. 列式存储,查询效率高
2. 支持压缩,节省磁盘空间
3. 保持数据类型,避免类型转换错误
"""
table = pa.Table.from_pandas(df)
pq.write_table(table, filepath, compression='snappy')
print(f"数据已保存为Parquet格式: {filepath}")
def load_from_parquet(filepath):
"""从Parquet文件加载数据"""
table = pq.read_table(filepath)
return table.to_pandas()
# 使用示例
df = pd.DataFrame(...) # 你的数据
save_as_parquet(df, "data/sh000001.parquet")
loaded_df = load_from_parquet("data/sh000001.parquet")
```
**方案二:使用HDF5格式**
```python
def save_as_hdf5(df, filepath, key='data'):
"""
将DataFrame保存为HDF5格式
HDF5的优势:
1. 支持快速随机访问
2. 可以存储多个数据集
3. 支持数据切片查询
"""
df.to_hdf(filepath, key=key, mode='w', complevel=9, complib='blosc')
print(f"数据已保存为HDF5格式: {filepath}")
def load_from_hdf5(filepath, key='data'):
"""从HDF5文件加载数据"""
return pd.read_hdf(filepath, key=key)
# 使用示例
save_as_hdf5(df, "data/sh000001.h5")
loaded_df = load_from_hdf5("data/sh000001.h5")
```
**方案三:数据库存储**
对于大规模数据,可以考虑使用数据库。这里以SQLite为例:
```python
import sqlite3
from sqlite3 import Error
def create_database(db_file="stock_data.db"):
"""创建SQLite数据库"""
conn = None
try:
conn = sqlite3.connect(db_file)
print(f"数据库连接成功: {db_file}")
return conn
except Error as e:
print(f"数据库连接失败: {e}")
return None
def init_stock_table(conn):
"""初始化股票数据表"""
create_table_sql = """
CREATE TABLE IF NOT EXISTS stock_daily (
code TEXT NOT NULL,
date DATE NOT NULL,
open REAL,
high REAL,
low REAL,
close REAL,
volume INTEGER,
amount REAL,
PRIMARY KEY (code, date)
);
"""
try:
cursor = conn.cursor()
cursor.execute(create_table_sql)
conn.commit()
print("数据表创建成功")
except Error as e:
print(f"创建数据表失败: {e}")
def save_to_database(df, stock_code, conn):
"""将数据保存到数据库"""
# 添加股票代码列
df['code'] = stock_code
# 重命名列以匹配数据库表结构
df = df.rename(columns={
'datetime': 'date',
'volume': 'volume',
'amount': 'amount'
})
# 选择需要的列
df = df[['code', 'date', 'open', 'high', 'low', 'close', 'volume', 'amount']]
# 保存到数据库
try:
df.to_sql('stock_daily', conn, if_exists='append', index=False)
print(f"{stock_code}: 数据已保存到数据库")
return True
except Exception as e:
print(f"{stock_code}: 保存到数据库失败: {str(e)}")
return False
```
### 4.2 查询性能优化
当数据量很大时,如何快速查询需要的数据就变得很重要。下面是一些优化技巧:
**建立索引**
```python
def create_indexes(conn):
"""创建查询索引"""
indexes = [
"CREATE INDEX IF NOT EXISTS idx_code ON stock_daily (code);",
"CREATE INDEX IF NOT EXISTS idx_date ON stock_daily (date);",
"CREATE INDEX IF NOT EXISTS idx_code_date ON stock_daily (code, date);"
]
cursor = conn.cursor()
for index_sql in indexes:
try:
cursor.execute(index_sql)
print(f"索引创建成功: {index_sql}")
except Error as e:
print(f"创建索引失败: {e}")
conn.commit()
# 使用示例
conn = create_database("stock_data.db")
init_stock_table(conn)
create_indexes(conn)
```
**批量查询优化**
```python
def query_multiple_stocks(codes, start_date, end_date, conn):
"""
批量查询多只股票在指定时间段的数据
参数:
codes: 股票代码列表
start_date: 开始日期(字符串,格式:YYYY-MM-DD)
end_date: 结束日期(字符串,格式:YYYY-MM-DD)
conn: 数据库连接
返回:
dict: 每只股票的DataFrame
"""
results = {}
# 使用参数化查询防止SQL注入
query_sql = """
SELECT * FROM stock_daily
WHERE code = ? AND date BETWEEN ? AND ?
ORDER BY date
"""
cursor = conn.cursor()
for code in codes:
cursor.execute(query_sql, (code, start_date, end_date))
rows = cursor.fetchall()
if rows:
# 获取列名
column_names = [description[0] for description in cursor.description]
df = pd.DataFrame(rows, columns=column_names)
results[code] = df
else:
results[code] = None
return results
```
**缓存常用查询结果**
```python
from functools import lru_cache
import hashlib
@lru_cache(maxsize=128)
def cached_stock_query(code, start_date, end_date, db_file="stock_data.db"):
"""
带缓存的股票查询
使用LRU缓存最近查询的结果,避免重复查询数据库
"""
cache_key = hashlib.md5(f"{code}{start_date}{end_date}".encode()).hexdigest()
conn = sqlite3.connect(db_file)
query_sql = """
SELECT * FROM stock_daily
WHERE code = ? AND date BETWEEN ? AND ?
ORDER BY date
"""
df = pd.read_sql_query(query_sql, conn, params=(code, start_date, end_date))
conn.close()
return df
```
### 4.3 数据质量监控
建立数据质量监控机制,确保数据的准确性和完整性:
```python
class DataQualityMonitor:
"""数据质量监控器"""
def __init__(self, data_dir="./data"):
self.data_dir = data_dir
self.quality_report = {}
def check_missing_dates(self, df, stock_code):
"""检查缺失的交易日"""
if df is None or df.empty:
return []
# 假设我们已经有交易日历
# 这里简化处理,实际应该使用真实的交易日历
df['date'] = pd.to_datetime(df['datetime']).dt.date
date_range = pd.date_range(start=df['date'].min(), end=df['date'].max())
missing_dates = []
for date in date_range:
if date.date() not in df['date'].values:
missing_dates.append(date.date())
return missing_dates
def check_price_anomalies(self, df, stock_code):
"""检查价格异常"""
anomalies = []
if df is None or df.empty:
return anomalies
# 检查价格是否为0或负数
price_columns = ['open', 'high', 'low', 'close']
for col in price_columns:
zero_or_negative = df[df[col] <= 0]
if not zero_or_negative.empty:
for _, row in zero_or_negative.iterrows():
anomalies.append({
'date': row['datetime'],
'type': 'zero_or_negative_price',
'column': col,
'value': row[col]
})
# 检查最高价是否低于最低价
invalid_high_low = df[df['high'] < df['low']]
if not invalid_high_low.empty:
for _, row in invalid_high_low.iterrows():
anomalies.append({
'date': row['datetime'],
'type': 'high_low_invalid',
'high': row['high'],
'low': row['low']
})
return anomalies
def check_volume_anomalies(self, df, stock_code):
"""检查成交量异常"""
anomalies = []
if df is None or df.empty:
return anomalies
# 检查成交量为0但价格有变动的情况
zero_volume = df[(df['volume'] == 0) & (df['close'] != df['open'])]
if not zero_volume.empty:
for _, row in zero_volume.iterrows():
anomalies.append({
'date': row['datetime'],
'type': 'zero_volume_with_price_change',
'volume': row['volume'],
'price_change': row['close'] - row['open']
})
# 检查异常大的成交量(超过平均值的10倍)
avg_volume = df['volume'].mean()
if avg_volume > 0:
high_volume = df[df['volume'] > avg_volume * 10]
if not high_volume.empty:
for _, row in high_volume.iterrows():
anomalies.append({
'date': row['datetime'],
'type': 'extremely_high_volume',
'volume': row['volume'],
'avg_volume': avg_volume,
'ratio': row['volume'] / avg_volume
})
return anomalies
def run_quality_check(self, stock_codes):
"""运行完整的数据质量检查"""
all_anomalies = {}
for code in stock_codes:
# 确定文件路径
if code.startswith('6'):
filename = f"sh{code}.csv"
else:
filename = f"sz{code}.csv"
filepath = os.path.join(self.data_dir, filename)
if not os.path.exists(filepath):
print(f"警告: {code} 的数据文件不存在")
continue
try:
df = pd.read_csv(filepath)
df['datetime'] = pd.to_datetime(df['datetime'])
# 运行各项检查
missing_dates = self.check_missing_dates(df, code)
price_anomalies = self.check_price_anomalies(df, code)
volume_anomalies = self.check_volume_anomalies(df, code)
# 汇总结果
stock_anomalies = {
'missing_dates': missing_dates,
'price_anomalies': price_anomalies,
'volume_anomalies': volume_anomalies,
'total_anomalies': len(missing_dates) + len(price_anomalies) + len(volume_anomalies)
}
all_anomalies[code] = stock_anomalies
if stock_anomalies['total_anomalies'] > 0:
print(f"{code}: 发现 {stock_anomalies['total_anomalies']} 个数据质量问题")
except Exception as e:
print(f"{code}: 数据质量检查失败: {str(e)}")
# 生成质量报告
self.generate_quality_report(all_anomalies)
return all_anomalies
def generate_quality_report(self, anomalies):
"""生成数据质量报告"""
total_stocks = len(anomalies)
stocks_with_issues = sum(1 for v in anomalies.values() if v['total_anomalies'] > 0)
total_issues = sum(v['total_anomalies'] for v in anomalies.values())
report = {
'检查时间': datetime.now().strftime('%Y-%m-%d %H:%M:%S'),
'检查股票数量': total_stocks,
'存在问题股票数量': stocks_with_issues,
'问题股票比例': stocks_with_issues / total_stocks * 100 if total_stocks > 0 else 0,
'问题总数': total_issues,
'详细问题': anomalies
}
# 保存报告
report_file = os.path.join(self.data_dir, f"quality_report_{datetime.now().strftime('%Y%m%d')}.json")
import json
with open(report_file, 'w', encoding='utf-8') as f:
json.dump(report, f, ensure_ascii=False, indent=2)
print(f"数据质量报告已保存: {report_file}")
return report
# 使用示例
monitor = DataQualityMonitor("./data")
anomalies = monitor.run_quality_check(['000001', '600000', '000002'])
```
### 4.4 实战案例:构建完整的量化数据管道
最后,让我们把这些组件组合起来,构建一个完整的量化数据管道:
```python
class QuantDataPipeline:
"""量化数据管道"""
def __init__(self, config_file="config.json"):
self.config = self.load_config(config_file)
self.data_dir = self.config.get('data_dir', './data')
self.log_dir = self.config.get('log_dir', './logs')
# 创建必要的目录
os.makedirs(self.data_dir, exist_ok=True)
os.makedirs(self.log_dir, exist_ok=True)
# 设置日志
self.logger = self.setup_logger()
def load_config(self, config_file):
"""加载配置文件"""
default_config = {
'data_dir': './data',
'log_dir': './logs',
'max_workers': 5,
'retry_times': 3,
'update_frequency': 'daily', # daily, weekly, monthly
'storage_format': 'parquet', # csv, parquet, hdf5, database
'quality_check': True,
'notify_on_error': False
}
if os.path.exists(config_file):
import json
with open(config_file, 'r', encoding='utf-8') as f:
user_config = json.load(f)
default_config.update(user_config)
return default_config
def setup_logger(self):
"""设置日志记录器"""
log_file = os.path.join(
self.log_dir,
f"pipeline_{datetime.now().strftime('%Y%m%d')}.log"
)
logger = logging.getLogger('QuantDataPipeline')
logger.setLevel(logging.INFO)
# 文件处理器
file_handler = logging.FileHandler(log_file, encoding='utf-8')
file_handler.setLevel(logging.INFO)
# 控制台处理器
console_handler = logging.StreamHandler()
console_handler.setLevel(logging.INFO)
# 格式化器
formatter = logging.Formatter(
'%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
file_handler.setFormatter(formatter)
console_handler.setFormatter(formatter)
logger.addHandler(file_handler)
logger.addHandler(console_handler)
return logger
def run_daily_update(self):
"""运行每日数据更新"""
self.logger.info("开始每日数据更新流程")
# 1. 获取需要更新的股票列表
stock_codes = self.get_update_list()
self.logger.info(f"需要更新 {len(stock_codes)} 只股票的数据")
# 2. 批量更新数据
update_stats = batch_update_stocks(
stock_codes,
self.data_dir,
max_workers=self.config['max_workers']
)
# 3. 数据质量检查
if self.config['quality_check']:
self.logger.info("开始数据质量检查")
monitor = DataQualityMonitor(self.data_dir)
quality_report = monitor.run_quality_check(stock_codes)
# 记录质量问题
issues_count = sum(r['total_anomalies'] for r in quality_report.values())
self.logger.info(f"数据质量检查完成,发现 {issues_count} 个问题")
# 4. 生成更新报告
report = self.generate_update_report(update_stats)
self.logger.info("每日数据更新流程完成")
return report
def get_update_list(self):
"""获取需要更新的股票列表"""
# 这里可以从配置文件、数据库或API获取
# 示例:从文件读取
list_file = os.path.join(self.data_dir, "watchlist.txt")
if os.path.exists(list_file):
with open(list_file, 'r', encoding='utf-8') as f:
codes = [line.strip() for line in f if line.strip()]
return codes
else:
# 如果没有监控列表,返回空列表
self.logger.warning("未找到股票监控列表文件")
return []
def generate_update_report(self, stats):
"""生成更新报告"""
report = {
'timestamp': datetime.now().isoformat(),
'total_stocks': stats['total'],
'updated': stats['updated'],
'failed': stats['failed'],
'success_rate': stats['updated'] / stats['total'] * 100 if stats['total'] > 0 else 0,
'failed_codes': stats['failed_codes']
}
# 保存报告
report_file = os.path.join(
self.log_dir,
f"update_report_{datetime.now().strftime('%Y%m%d_%H%M%S')}.json"
)
import json
with open(report_file, 'w', encoding='utf-8') as f:
json.dump(report, f, ensure_ascii=False, indent=2)
self.logger.info(f"更新报告已保存: {report_file}")
return report
def export_for_backtest(self, codes, start_date, end_date, output_format='csv'):
"""
导出回测所需数据
参数:
codes: 股票代码列表
start_date: 开始日期
end_date: 结束日期
output_format: 输出格式(csv, parquet, hdf5)
返回:
dict: 每只股票的数据
"""
self.logger.info(f"导出回测数据: {len(codes)}只股票,{start_date}到{end_date}")
all_data = {}
for code in tqdm(codes, desc="导出进度"):
# 根据存储格式读取数据
if self.config['storage_format'] == 'database':
# 从数据库读取
conn = sqlite3.connect(os.path.join(self.data_dir, "stock_data.db"))
query = """
SELECT * FROM stock_daily
WHERE code = ? AND date BETWEEN ? AND ?
ORDER BY date
"""
df = pd.read_sql_query(query, conn, params=(code, start_date, end_date))
conn.close()
else:
# 从文件读取
if code.startswith('6'):
filename = f"sh{code}.{self.config['storage_format']}"
else:
filename = f"sz{code}.{self.config['storage_format']}"
filepath = os.path.join(self.data_dir, filename)
if not os.path.exists(filepath):
self.logger.warning(f"数据文件不存在: {filename}")
continue
# 根据格式读取
if self.config['storage_format'] == 'parquet':
df = pd.read_parquet(filepath)
elif self.config['storage_format'] == 'hdf5':
df = pd.read_hdf(filepath, key='data')
else: # csv
df = pd.read_csv(filepath)
# 过滤日期范围
if 'datetime' in df.columns:
df['date'] = pd.to_datetime(df['datetime'])
elif 'date' in df.columns:
df['date'] = pd.to_datetime(df['date'])
mask = (df['date'] >= pd.Timestamp(start_date)) & (df['date'] <= pd.Timestamp(end_date))
df = df[mask]
if not df.empty:
all_data[code] = df
self.logger.info(f"数据导出完成,成功导出 {len(all_data)} 只股票的数据")
return all_data
# 使用示例
if __name__ == "__main__":
# 初始化管道
pipeline = QuantDataPipeline("config.json")
# 运行每日更新
report = pipeline.run_daily_update()
# 导出回测数据
backtest_data = pipeline.export_for_backtest(
codes=['000001', '600000', '000002'],
start_date='2023-01-01',
end_date='2023-12-31',
output_format='csv'
)
print(f"导出完成,共 {len(backtest_data)} 只股票的数据")
```
这个完整的管道包含了数据更新、质量检查、报告生成和导出功能,可以直接用于生产环境。你可以根据自己的需求调整配置,比如修改更新频率、存储格式、监控的股票列表等。
在实际使用中,我通常会设置一个定时任务,每天收盘后自动运行数据更新。对于需要实时监控的策略,还可以考虑增加实时数据获取的功能。不过对于大多数量化策略来说,日线数据已经足够用了,关键是保证数据的准确性和完整性。
数据是量化交易的基石,一个可靠的数据管道能为你节省大量时间,让你更专注于策略开发。希望这套方案能帮助你建立起自己的数据基础设施。如果在使用过程中遇到问题,或者有更好的优化建议,欢迎交流讨论。毕竟,在量化这条路上,我们都是不断学习和改进的同行者。