qlib使用tushare更新每日行情数据

网上有很多方法但是都不好用,经过三天的踩坑终于把数据全部更新到最新的数据

首先需要有个tushare的token,用一天左右吧

先要下载qlib的源文件,因为最后一步格式转换需要用到源文件里面的py文件

直接上代码,注释比较清晰

#获取最新的数据,第一步按年确定有数据的日期
import tushare as ts
token = '你的token'
ts.set_token(token)
pro = ts.pro_api()  
df = pro.trade_cal(exchange='', start_date='20200926', end_date='20210926',is_open="1")
df.to_csv("cal.csv")
#获取最新的数据,第二步csv里面的日期获取数据并保存为日期命名的csv文件
import csv
import tushare as ts
import pandas as pd
import numpy as np
from pathlib import Path
import shutil
from datetime import datetime, timedelta
import os
import time
import threading
from concurrent.futures import ThreadPoolExecutor, as_completed
import logging

# 设置日志
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)

token = '你的token'
ts.set_token(token)
pro = ts.pro_api()

class DateBasedDataUpdater:
    def __init__(self, qlib_data_dir="/mnt/workspace/qlib/qlib/qlib_data/cn_data", max_workers=10):
        self.qlib_data_dir = Path(qlib_data_dir).expanduser()
        self.csv_temp_dir = self.qlib_data_dir.parent / "akshare_temp1"
        self.max_workers = max_workers
        self.error_dates = []
        self.success_count = 0
        self.total_count = 0
        self.lock = threading.Lock()
        
        # 创建目录
        self.csv_temp_dir.mkdir(parents=True, exist_ok=True)
    
    def get_calendar_dates(self, cal_file='cal.csv'):
        """从日历文件获取所有日期"""
        dates = []
        try:
            with open(cal_file, mode='r', newline='', encoding='utf-8') as file:
                csv_reader = csv.DictReader(file)
                for row in csv_reader:
                    dates.append(row['cal_date'])
            logger.info(f"从日历文件读取 {len(dates)} 个交易日")
        except Exception as e:
            logger.error(f"读取日历文件失败: {e}")
        return dates
   
    
    def download_date_data(self, trade_date, retry_count=3):
        """下载指定交易日的数据"""
        for attempt in range(retry_count):
            try:
                # 添加延迟避免频繁请求
                time.sleep(0.1)
                
                logger.debug(f"正在获取 {trade_date} 的数据 (尝试 {attempt + 1}/{retry_count})")
                
                df = pro.daily(trade_date=trade_date)
                
                if df.empty:
                    logger.warning(f"{trade_date} 数据为空")
                    return None
                
                # 重命名列以匹配Qlib格式
                df = df.rename(columns={
                    'trade_date': 'date',
                    'open': 'open',
                    'close': 'close', 
                    'high': 'high',
                    'low': 'low',
                    'vol': 'volume'
                })
                
                # 添加复权因子
                df['factor'] = 1.0
                
                # 设置日期索引
                df['date'] = pd.to_datetime(df['date'])
                df = df.set_index('ts_code')
                
                logger.debug(f"成功获取 {trade_date} 的数据,共 {len(df)} 条记录")
                return df
                
            except Exception as e:
                error_msg = str(e)
                if "抱歉,您每分钟最多访问该接口800次" in error_msg:
                    logger.warning(f"API限额已用完,等待65秒后重试")
                    time.sleep(65)
                    continue
                else:
                    logger.warning(f"获取 {trade_date} 数据失败 (尝试 {attempt + 1}/{retry_count}): {e}")
                
                if attempt < retry_count - 1:
                    wait_time = (attempt + 1) * 2  # 指数退避
                    logger.info(f"等待 {wait_time} 秒后重试...")
                    time.sleep(wait_time)
                else:
                    with self.lock:
                        self.error_dates.append(trade_date)
                    logger.error(f"获取 {trade_date} 数据最终失败")
                    return None
    
    def save_date_data(self, trade_date, df):
        """保存单日数据到CSV文件"""
        try:
            csv_file = self.csv_temp_dir / f"{trade_date}.csv"
            df.to_csv(csv_file)
            
            with self.lock:
                self.success_count += 1
            
            logger.debug(f"已保存 {trade_date} 的数据")
            return True
        except Exception as e:
            logger.error(f"保存 {trade_date} 数据失败: {e}")
            return False
    
    def process_date_batch(self, date_batch):
        """处理一批日期数据"""
        results = []
        for trade_date in date_batch:
            df = self.download_date_data(trade_date)
            if df is not None and self.save_date_data(trade_date, df):
                results.append(trade_date)
            
            # 更新进度
            if self.success_count % 100 == 0 and self.total_count > 0:
                logger.info(f"进度: {self.success_count}/{self.total_count} (成功率: {self.success_count/self.total_count*100:.1f}%)")
        
        return results
    
    def update_data(self):
        """主更新函数 - 使用多线程并行处理日期"""
        # 获取所有交易日
        calendar_dates = self.get_calendar_dates()
        if not calendar_dates:
            logger.error("未获取到交易日历")
            return
        
        self.total_count = len(calendar_dates)
        logger.info(f"开始处理 {self.total_count} 个交易日的数据")
        
        # 分批处理日期
        batch_size = 50  # 每批处理50个日期
        date_batches = [calendar_dates[i:i + batch_size] for i in range(0, len(calendar_dates), batch_size)]
        
        all_success_dates = []
        
        # 使用线程池并行处理
        with ThreadPoolExecutor(max_workers=self.max_workers) as executor:
            # 提交所有批次任务
            future_to_batch = {
                executor.submit(self.process_date_batch, batch): i 
                for i, batch in enumerate(date_batches)
            }
            
            # 处理完成的任务
            for future in as_completed(future_to_batch):
                batch_index = future_to_batch[future]
                try:
                    batch_results = future.result()
                    all_success_dates.extend(batch_results)
                    logger.info(f"已完成批次 {batch_index + 1}/{len(date_batches)}")
                except Exception as e:
                    logger.error(f"处理批次 {batch_index + 1} 时发生错误: {e}")
        
        # 保存失败日期列表
        if self.error_dates:
            error_file = self.csv_temp_dir.parent / "errordate.csv"
            with open(error_file, 'w', newline='', encoding='utf-8') as f:
                writer = csv.writer(f)
                for date in self.error_dates:
                    writer.writerow([date])
            logger.info(f"失败日期已保存到 {error_file},共 {len(self.error_dates)} 个")
        
        logger.info(f"数据更新完成: 成功 {self.success_count}, 失败 {len(self.error_dates)}")

if __name__ == "__main__":
    max_workers = 15
    
    updater = DateBasedDataUpdater(max_workers=max_workers)
    updater.update_data()
#获取最新的数据,第三步再一次获取全部的有数据的日期
import tushare as ts
token = '你的token'
ts.set_token(token)
pro = ts.pro_api()  
df = pro.trade_cal(exchange='', start_date='20200926',is_open="1")
df.to_csv("cal.csv")
#获取最新的数据,第四步比对文件与日期是否有缺失,如果有缺失就显示出来,并且手动复制到日期csv文件中,返回到第二步补上缺失的文件
import pandas as pd
import os

def check_missing_csv_files(summary_file, folder_path, date_column='cal_date'):
    """
    检查汇总文件中每个日期是否都有对应的CSV文件
    
    参数:
    summary_file: 日期汇总CSV文件路径
    folder_path: 存放CSV文件的文件夹路径
    date_column: 汇总文件中日期列的列名,默认为'cal_date'
    """
    try:
        df_summary = pd.read_csv(summary_file)
        print(f"成功读取汇总文件,共{len(df_summary)}行数据")
    except Exception as e:
        print(f"读取汇总文件失败: {e}")
        return
    
    # 检查日期列是否存在
    if date_column not in df_summary.columns:
        print(f"汇总文件中找不到列: {date_column}")
        print(f"可用列: {list(df_summary.columns)}")
        return
    
    # 获取文件夹中所有的CSV文件名
    try:
        csv_files = [f for f in os.listdir(folder_path) if f.endswith('.csv')]
        print(f"在文件夹中找到{len(csv_files)}个CSV文件")
    except Exception as e:
        print(f"读取文件夹失败: {e}")
        return
    
    # 从文件名中提取日期(去掉.csv后缀)
    file_dates = set()
    for file in csv_files:
        # 假设文件名格式为YYYYMMDD.csv
        date_str = file.replace('.csv', '')
        file_dates.add(date_str)
    
    # 处理汇总文件中的日期列,转换为字符串格式进行匹配
    missing_dates = []
    
    for date_value in df_summary[date_column]:
        # 处理不同的日期格式
        if isinstance(date_value, str):
            # 如果是字符串,可能需要清理
            date_str = date_value.strip().replace('-', '').replace('/', '')[:8]
        else:
            # 如果是日期类型,转换为字符串
            date_str = str(int(date_value))[:8] if pd.notna(date_value) else ""
        
        # 检查日期是否在文件日期集合中
        if date_str and date_str not in file_dates:
            missing_dates.append(date_str)
    
    # 输出结果
    if missing_dates:
        print("\n以下日期缺少对应的CSV文件:")
        for date in missing_dates:
            print(f"{date}")
        print(f"\n总计缺少 {len(missing_dates)} 个文件")
    else:
        print("\n所有日期都有对应的CSV文件,文件完整!")
    
    return missing_dates

if __name__ == "__main__":
    summary_file = "cal.csv"  # 替换为你的汇总文件路径
    folder_path = "/mnt/workspace/qlib/qlib/qlib_data/akshare_temp1"   # 替换为你的CSV文件夹路径
    missing_files = check_missing_csv_files(summary_file, folder_path)
#获取最新的数据,第五步按照代码维度重新整理并生成csv文件
import pandas as pd
import os
import glob

def process_ts_code_files(folder_path):
    """
    处理文件夹中的日期CSV文件,按照ts_code列重新组织数据
    
    参数:
    folder_path: 包含日期CSV文件的文件夹路径
    """
    # 获取所有日期格式的CSV文件
    csv_files = glob.glob(os.path.join(folder_path, "*.csv"))
    date_files = [f for f in csv_files if os.path.basename(f).replace('.csv', '').isdigit() and len(os.path.basename(f).replace('.csv', '')) == 8]
    
    print(f"找到 {len(date_files)} 个日期CSV文件")
    
    # 用于存储每个ts_code对应的数据
    ts_code_data = {}
    
    # 处理每个日期文件
    for file_path in date_files:
        filename = os.path.basename(file_path)
        date_str = filename.replace('.csv', '')
        
        print(f"正在处理文件: {filename}")
        
        try:
            # 读取CSV文件
            df = pd.read_csv(file_path)
            
            # 检查是否存在ts_code列
            if 'ts_code' not in df.columns:
                print(f"警告: 文件 {filename} 中不存在 'ts_code' 列,跳过此文件")
                continue
            
            # 按ts_code分组处理
            for ts_code, group in df.groupby('ts_code'):
                # 解析ts_code格式:YYYYYY.ZZ -> 转换为 ZZYYYY
                if '.' in str(ts_code):
                    parts = str(ts_code).split('.')
                    if len(parts) == 2:
                        new_filename = f"{parts[1]}{parts[0]}.csv"
                    else:
                        print(f"警告: ts_code格式异常: {ts_code},使用原值作为文件名")
                        new_filename = f"{ts_code}.csv"
                else:
                    new_filename = f"{ts_code}.csv"
                
                # 如果这个ts_code是第一次遇到,初始化DataFrame
                if new_filename not in ts_code_data:
                    ts_code_data[new_filename] = group.copy()
                else:
                    # 追加到已存在的数据
                    ts_code_data[new_filename] = pd.concat([ts_code_data[new_filename], group], ignore_index=True)
                    
        except Exception as e:
            print(f"处理文件 {filename} 时出错: {e}")
            continue
    
    return ts_code_data

def save_ts_code_files(ts_code_data, output_folder):
    """
    将处理后的ts_code数据保存到CSV文件
    
    参数:
    ts_code_data: 包含ts_code数据的字典
    output_folder: 输出文件夹路径
    """
    if not os.path.exists(output_folder):
        os.makedirs(output_folder)
    
    saved_files = []
    for filename, data in ts_code_data.items():
        output_path = os.path.join(output_folder, filename)
        data.to_csv(output_path, index=False)
        saved_files.append(filename)
        print(f"已创建文件: {filename} (包含 {len(data)} 行数据)")
    
    return saved_files

def main():
    # 配置参数
    input_folder = "/mnt/workspace/qlib/qlib/qlib_data/akshare_temp1"  # 替换为您的CSV文件夹路径
    output_folder = "/mnt/workspace/qlib/qlib/qlib_data/akshare_temp2"    # 替换为您希望保存结果的文件夹路径
    
    # 处理文件
    print("开始处理CSV文件...")
    ts_code_data = process_ts_code_files(input_folder)
    
    if not ts_code_data:
        print("未找到有效数据,程序结束")
        return
    
    print(f"\n成功处理数据,共涉及 {len(ts_code_data)} 个不同的ts_code")
    
    # 保存结果
    print("\n开始保存结果文件...")
    saved_files = save_ts_code_files(ts_code_data, output_folder)
    
    print(f"\n处理完成!共生成 {len(saved_files)} 个按ts_code分类的CSV文件")
    print("生成的文件列表:")
    for file in sorted(saved_files):
        print(f"  - {file}")

# 支持大文件处理和进度显示
def process_ts_code_files_enhanced(folder_path, chunksize=10000):
    """
    增强版本:支持大文件分块读取,节省内存
    """
    csv_files = glob.glob(os.path.join(folder_path, "*.csv"))
    date_files = [f for f in csv_files if os.path.basename(f).replace('.csv', '').isdigit() and len(os.path.basename(f).replace('.csv', '')) == 8]
    
    print(f"找到 {len(date_files)} 个日期CSV文件")
    
    ts_code_data = {}
    
    for file_path in date_files:
        filename = os.path.basename(file_path)
        print(f"正在处理文件: {filename}")
        
        try:
            # 分块读取大文件
            chunk_number = 0
            for chunk in pd.read_csv(file_path, chunksize=chunksize):
                chunk_number += 1
                print(f"  处理第 {chunk_number} 块数据...")
                
                if 'ts_code' not in chunk.columns:
                    continue
                
                for ts_code, group in chunk.groupby('ts_code'):
                    if '.' in str(ts_code):
                        parts = str(ts_code).split('.')
                        new_filename = f"{parts[1]}{parts[0]}.csv" if len(parts) == 2 else f"{ts_code}.csv"
                    else:
                        new_filename = f"{ts_code}.csv"
                    
                    if new_filename not in ts_code_data:
                        ts_code_data[new_filename] = group.copy()
                    else:
                        ts_code_data[new_filename] = pd.concat([ts_code_data[new_filename], group], ignore_index=True)
                        
        except Exception as e:
            print(f"处理文件 {filename} 时出错: {e}")
            continue
    
    return ts_code_data

# 使用示例
if __name__ == "__main__":
    main()
    
    # 如果需要处理大文件,使用增强版本
    # ts_code_data = process_ts_code_files_enhanced("./csv_files")
    # saved_files = save_ts_code_files(ts_code_data, "./output")
#获取最新的数据,第六步通过qlib自动工具把csv文件格式化成qlib的二进制数据,注意会覆盖所有的原始数据
!python /mnt/workspace/qlib/scripts/dump_bin.py dump_all --data_path /mnt/workspace/qlib/qlib/qlib_data/akshare_temp2 --qlib_dir /mnt/workspace/qlib/qlib/qlib_data/cn_data --include_fields open,close,high,low,volume,factor
#获取最新的数据,第七步如果不放心可以check一下
!python /mnt/workspace/qlib/scripts/check_data_health.py check_data --qlib_dir /mnt/workspace/qlib/qlib/qlib_data/cn_data

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值