diff --git a/quantitative_data/importer.py b/quantitative_data/importer.py index 32aeba2..1fa5c3d 100644 --- a/quantitative_data/importer.py +++ b/quantitative_data/importer.py @@ -307,7 +307,7 @@ def import_trade_cal(start_date: Optional[str] = None, end_date: Optional[str] = def import_daily_for_stock(ts_code: str, start_date: str, end_date: str, conn) -> int: """ - 导入单只股票的日线行情 + 导入单只股票的日线行情(保留用于单只股票补充/重试) 返回导入的记录数 """ pro = get_ts_pro() @@ -325,11 +325,9 @@ def import_daily_for_stock(ts_code: str, start_date: str, end_date: str, conn) - df = normalize_columns(df) - # 转换日期 if "trade_date" in df.columns: df["trade_date"] = pd.to_datetime(df["trade_date"], format="%Y%m%d", errors="coerce") - # 数值列处理 NaN numeric_cols = [ "open", "high", "low", "close", "pre_close", "change", "pct_chg", "vol", "amount", @@ -338,7 +336,6 @@ def import_daily_for_stock(ts_code: str, start_date: str, end_date: str, conn) - if col in df.columns: df[col] = pd.to_numeric(df[col], errors="coerce") - # 额外列 (Tushare Pro 不同版本返回字段可能不同) for col in ["turnover_rate", "volume_ratio", "ma5", "ma10", "ma20", "ma_v_5", "ma_v_10", "ma_v_20"]: if col not in df.columns: df[col] = None @@ -347,85 +344,145 @@ def import_daily_for_stock(ts_code: str, start_date: str, end_date: str, conn) - return batch_insert("daily", df, conn, conflict_cols) -def import_daily_batch( - stock_list: List[str], +def _normalize_daily_df(df: pd.DataFrame) -> pd.DataFrame: + """标准化 daily DataFrame 的列和类型 (供 import_daily_by_date 复用)""" + df = normalize_columns(df) + if "trade_date" in df.columns: + df["trade_date"] = pd.to_datetime(df["trade_date"], format="%Y%m%d", errors="coerce") + numeric_cols = [ + "open", "high", "low", "close", "pre_close", "change", "pct_chg", + "vol", "amount", + ] + for col in numeric_cols: + if col in df.columns: + df[col] = pd.to_numeric(df[col], errors="coerce") + for col in ["turnover_rate", "volume_ratio", "ma5", "ma10", "ma20", "ma_v_5", "ma_v_10", "ma_v_20"]: + if col not in df.columns: + df[col] = None + return df + + +def import_daily_by_date( start_date: Optional[str] = None, end_date: Optional[str] = None, + conn=None, sleep_interval: float = 0.3, ): """ - 批量导入多只股票的日线行情 - - stock_list: 股票代码列表 - - sleep_interval: API 调用间隔 (避免频率限制) + 按交易日批量导入日线行情 (高效模式) + 使用 pro.daily(trade_date='YYYYMMDD') 一次性拉取全市场当日数据 + 大幅减少 API 调用次数: 约250交易日/年 × 16年 ≈ 4000次 (原来需要 5000股票 × 16年 = 80000次) + + 返回: 失败的交易日列表 """ if start_date is None: start_date = START_DATE if end_date is None: end_date = END_DATE - total = len(stock_list) - logger.info("=" * 60) - logger.info( - f"[3/7] 导入日线行情 (daily): {start_date} ~ {end_date}, " - f"共 {total} 只股票" - ) + own_conn = conn is None + if own_conn: + conn = get_pg_connection() - conn = get_pg_connection() + # 从 trade_cal 获取交易日列表 + try: + cursor = conn.cursor() + cursor.execute( + """ + SELECT DISTINCT cal_date FROM trade_cal + WHERE is_open = 1 + AND cal_date >= %s AND cal_date <= %s + ORDER BY cal_date + """, + (start_date, end_date), + ) + trade_dates = [row[0] for row in cursor.fetchall()] + cursor.close() + finally: + if own_conn: + conn.close() + conn = get_pg_connection() + + total = len(trade_dates) + if total == 0: + logger.warning(f" 日期范围 {start_date} ~ {end_date} 内无交易日") + if own_conn: + conn.close() + return [] + + logger.info(f" 日期范围 {start_date} ~ {end_date}: 共 {total} 个交易日") + + pro = get_ts_pro() success_count = 0 fail_list = [] - for i, ts_code in enumerate(stock_list, 1): + for i, td in enumerate(trade_dates, 1): try: - n = import_daily_for_stock(ts_code, start_date, end_date, conn) - if n > 0: - success_count += 1 - if i % 50 == 0 or i == total: - logger.info(f" 进度: {i}/{total} 成功={success_count} 失败={len(fail_list)}") + td_str = td.strftime("%Y%m%d") if hasattr(td, "strftime") else str(td).replace("-", "") + + def fetch(): + return pro.daily(trade_date=td_str) + + df = fetch_with_retry(fetch, max_retries=3) + if df is None or df.empty: + logger.warning(f" [{td_str}] 返回空数据 (可能非交易日或API限制)") + continue + + df = _normalize_daily_df(df) + conflict_cols = ["ts_code", "trade_date"] + batch_insert("daily", df, conn, conflict_cols) + success_count += 1 + except Exception as e: - logger.error(f" [{ts_code}] 导入失败: {e}") - fail_list.append(ts_code) - conn.rollback() + logger.error(f" [{td}] 导入失败: {e}") + fail_list.append(str(td)) + try: + conn.rollback() + except Exception: + pass - time.sleep(sleep_interval) # API 频率控制 + if i % 50 == 0 or i == total: + logger.info(f" 进度: {i}/{total} 成功={success_count} 失败={len(fail_list)}") - conn.close() - logger.info(f" 日线行情导入完成: 成功 {success_count}/{total}") + time.sleep(sleep_interval) + if own_conn: + conn.close() + + logger.info( + f" 日线行情按日期导入完成: 成功 {success_count}/{total} 个交易日" + ) if fail_list: - logger.warning(f" 失败列表({len(fail_list)}): {fail_list[:20]}...") + logger.warning(f" 失败日期({len(fail_list)}): {fail_list[:20]}...") return fail_list def import_daily_by_year( - stock_list: List[str], start_year: int = 2010, end_year: int = 2025, + sleep_interval: float = 0.3, ): """ - 按年份逐批导入日线行情 (断点续传友好) - 适合大数据量导入,每年每只股票可单独重试 + 按年份逐批导入日线行情 (按交易日循环拉取全市场数据) + 不再需要 stock_list 参数 — 每次 API 调用拉取当日全市场数据 """ logger.info("=" * 60) logger.info( - f"[3/7] 按年导入日线行情: {start_year} ~ {end_year}, " - f"共 {len(stock_list)} 只股票" + f"[3/7] 按年导入日线行情 (按交易日): {start_year} ~ {end_year}" ) - total_imported = 0 for year in range(start_year, end_year + 1): year_start = f"{year}-01-01" year_end = f"{year}-12-31" logger.info(f"--- 导入 {year} 年日线行情 ---") - fail_list = import_daily_batch( - stock_list, + import_daily_by_date( start_date=year_start, end_date=year_end, - sleep_interval=0.2, + sleep_interval=sleep_interval, ) - total_imported += 1 logger.info(f" {year} 年完成\n") - logger.info(f" 所有年份日线行情导入完成!") + logger.info(" 所有年份日线行情导入完成!") # ============================================================ @@ -956,9 +1013,8 @@ def full_import( logger.error("无法获取股票列表,请先导入 stock_basic") return - # Step 3: 日线行情 (按年导入) + # Step 3: 日线行情 (按年导入,按交易日循环拉取全市场数据) import_daily_by_year( - stock_codes, start_year=int(start_date[:4]), end_year=int(end_date[:4]), ) @@ -1071,58 +1127,136 @@ def check_table_summary(conn=None): conn.close() -def resume_daily_by_year(stock_list, start_year=2010, end_year=2025): +def get_missing_daily_dates( + start_date: Optional[str] = None, + end_date: Optional[str] = None, + conn=None, +) -> List[str]: """ - 从中断点恢复按年导入日线行情 - 自动跳过数据库已有的年份,只导入缺失年份的数据 + 获取 daily 表中缺失的交易日列表 + 对比 trade_cal 中 is_open=1 的日期和 daily 表已有的 trade_date, + 返回未导入的交易日列表。 + + 返回: 缺失交易日字符串列表 (YYYY-MM-DD 格式) """ - conn = get_pg_connection() + if start_date is None: + start_date = START_DATE + if end_date is None: + end_date = END_DATE + + own_conn = conn is None + if own_conn: + conn = get_pg_connection() + try: cursor = conn.cursor() - cursor.execute(""" - SELECT DISTINCT EXTRACT(YEAR FROM trade_date)::int AS year - FROM daily - ORDER BY year - """) - completed_years = set(row[0] for row in cursor.fetchall()) - cursor.close() - finally: - conn.close() - - logger.info(f"已完成年份: {sorted(completed_years)}") - logger.info(f"待导入年份: {[y for y in range(start_year, end_year+1) if y not in completed_years]}") - - for year in range(start_year, end_year + 1): - if year in completed_years: - # 检查该年的股票覆盖是否完整 - conn = get_pg_connection() - try: - cursor = conn.cursor() - cursor.execute( - sql.SQL(""" - SELECT COUNT(DISTINCT ts_code) - FROM {} - WHERE EXTRACT(YEAR FROM trade_date) = %s - """).format(sql.Identifier("daily")), - (year,), - ) - stock_count = cursor.fetchone()[0] - cursor.close() - finally: - conn.close() - logger.info(f" {year} 年: 已有 {stock_count} 只股票, 跳过") - continue - - year_start = f"{year}-01-01" - year_end = f"{year}-12-31" - logger.info(f"--- 导入 {year} 年日线行情 ---") - import_daily_batch( - stock_list, - start_date=year_start, - end_date=year_end, - sleep_interval=0.2, + cursor.execute( + """ + SELECT tc.cal_date + FROM trade_cal tc + WHERE tc.is_open = 1 + AND tc.cal_date >= %s + AND tc.cal_date <= %s + AND NOT EXISTS ( + SELECT 1 FROM daily d + WHERE d.trade_date = tc.cal_date + ) + ORDER BY tc.cal_date + """, + (start_date, end_date), ) - logger.info(f" {year} 年完成\n") + missing_dates = [row[0].strftime("%Y-%m-%d") if hasattr(row[0], "strftime") else str(row[0])[:10] + for row in cursor.fetchall()] + cursor.close() + return missing_dates + finally: + if own_conn: + conn.close() + + +def resume_daily_by_date( + start_date: Optional[str] = None, + end_date: Optional[str] = None, + sleep_interval: float = 0.3, +): + """ + 按缺失日期断点续传日线行情 + 自动查询 daily 表已有的 trade_date 与 trade_cal 对比, + 只导入缺失日期的全市场数据。 + + 用法: + resume_daily_by_date(start_date="2010-01-01", end_date="2025-12-31") + """ + if start_date is None: + start_date = START_DATE + if end_date is None: + end_date = END_DATE + + logger.info("=" * 60) + logger.info(f"[断点续传] 检测缺失日期: {start_date} ~ {end_date}") + + missing_dates = get_missing_daily_dates(start_date, end_date) + + if not missing_dates: + logger.info(" 所有交易日数据已完整,无需续传!") + return + + total = len(missing_dates) + logger.info(f" 发现 {total} 个缺失交易日待导入") + if total <= 20: + logger.info(f" 缺失日期: {missing_dates}") + else: + logger.info(f" 缺失日期 (前20): {missing_dates[:20]}") + + # 按年份分组统计 + years_map: Dict[int, List[str]] = {} + for d in missing_dates: + y = int(d[:4]) + years_map.setdefault(y, []).append(d) + for y in sorted(years_map): + logger.info(f" {y} 年: {len(years_map[y])} 个缺失交易日") + + # 逐日期导入 + conn = get_pg_connection() + pro = get_ts_pro() + success_count = 0 + fail_list = [] + + for i, td_str in enumerate(missing_dates, 1): + try: + td_compact = td_str.replace("-", "") + + def fetch(): + return pro.daily(trade_date=td_compact) + + df = fetch_with_retry(fetch, max_retries=3) + if df is None or df.empty: + logger.warning(f" [{td_str}] 返回空数据,跳过") + continue + + df = _normalize_daily_df(df) + conflict_cols = ["ts_code", "trade_date"] + batch_insert("daily", df, conn, conflict_cols) + success_count += 1 + + except Exception as e: + logger.error(f" [{td_str}] 导入失败: {e}") + fail_list.append(td_str) + try: + conn.rollback() + except Exception: + pass + + if i % 50 == 0 or i == total: + logger.info(f" 续传进度: {i}/{total} 成功={success_count} 失败={len(fail_list)}") + + time.sleep(sleep_interval) + + conn.close() + logger.info(f" 断点续传完成: 成功 {success_count}/{total}") + if fail_list: + logger.warning(f" 失败日期({len(fail_list)}): {fail_list[:20]}...") + return fail_list if __name__ == "__main__": diff --git a/quantitative_data/数据批量导入.ipynb b/quantitative_data/数据批量导入.ipynb index ba684be..1986e65 100644 --- a/quantitative_data/数据批量导入.ipynb +++ b/quantitative_data/数据批量导入.ipynb @@ -42,16 +42,7 @@ "cell_type": "code", "execution_count": 8, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "工作目录: t:\\jupyter\\notebook\\quantitative_data\n", - "Python 版本: 3.10.2 (heads/master:d9999f5, Dec 16 2022, 16:20:32) [MSC v.1929 64 bit (AMD64)]\n" - ] - } - ], + "outputs": [], "source": [ "import sys\n", "import os\n", @@ -64,28 +55,7 @@ "cell_type": "code", "execution_count": 9, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "pandas 1.5.0\n", - "python-dotenv 1.2.2\n", - "SQLAlchemy 1.3.24\n", - "tushare 1.2.89\n", - "vnpy-tushare 1.2.85.1\n" - ] - }, - { - "name": "stderr", - "output_type": "stream", - "text": [ - "\n", - "[notice] A new release of pip available: 22.2.2 -> 26.2\n", - "[notice] To update, run: python.exe -m pip install --upgrade pip\n" - ] - } - ], + "outputs": [], "source": [ "# 检查依赖包\n", "!pip list | findstr -i \"tushare pandas psycopg2-binary sqlalchemy python-dotenv\"" @@ -105,15 +75,7 @@ "cell_type": "code", "execution_count": 10, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "模块导入成功!\n" - ] - } - ], + "outputs": [], "source": [ "# 导入核心模块\n", "from importer import (\n", @@ -122,7 +84,8 @@ " init_database,\n", " import_stock_basic,\n", " import_trade_cal,\n", - " import_daily_batch,\n", + " import_daily_for_stock,\n", + " import_daily_by_date,\n", " import_daily_by_year,\n", " import_daily_basic,\n", " import_daily_basic_by_date,\n", @@ -136,7 +99,8 @@ " batch_insert,\n", " check_daily_progress,\n", " check_table_summary,\n", - " resume_daily_by_year,\n", + " get_missing_daily_dates,\n", + " resume_daily_by_date,\n", " logger,\n", ")\n", "from config import DB_CONFIG, TUSHARE_TOKEN, START_DATE, END_DATE\n", @@ -148,20 +112,7 @@ "cell_type": "code", "execution_count": 11, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "✗ 连接失败: connection to server at \"192.168.27.11\", port 5438 failed: fe_sendauth: no password supplied\n", - "\n", - "请检查:\n", - " 1. Docker 容器是否已启动: docker ps | findstr postgres\n", - " 2. 环境变量 (.env) 中的连接参数是否正确\n", - " 3. 防火墙是否开放 5438 端口\n" - ] - } - ], + "outputs": [], "source": [ "# 测试数据库连接\n", "try:\n", @@ -309,28 +260,14 @@ "---\n", "## Step 4: 导入日线行情 (核心表,最耗时)\n", "\n", - "### 4.1 获取股票列表" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "# 获取所有需要导入的股票代码\n", - "stock_list = get_stock_codes_from_db()\n", - "print(f\"共 {len(stock_list)} 只股票需要导入日线行情\")\n", - "print(f\"前 10 只: {stock_list[:10]}\")" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "### 4.2 按年批量导入 (推荐 - 断点续传友好)\n", + "> **新版改进:** 按交易日循环拉取全市场数据 `pro.daily(trade_date='20180810')`,\n", + "> API 调用从 ~80,000次 (5000只×16年) 降至 ~4,000次 (250交易日×16年),速度提升约 **20倍**。\n", "\n", - "数据量估算: 约5000只股票 × 250交易日/年 × 16年 ≈ 2000万条记录" + "### 4.1 按年批量导入 (推荐)\n", + "\n", + "数据量估算: ~4000个交易日,每个交易日约5000条记录 ≈ 2000万条记录\n", + "\n", + "> **注意:** 新版 `import_daily_by_year` 不再需要 `stock_list` 参数,内部自动从 `trade_cal` 获取交易日列表后逐日拉取全市场数据。" ] }, { @@ -339,12 +276,13 @@ "metadata": {}, "outputs": [], "source": [ - "# 按年份逐批导入日线行情\n", - "# 如果中断,可以修改年份范围从断点继续\n", + "# 按年份逐批导入日线行情 (按交易日循环拉取全市场数据)\n", + "# 不再需要 stock_list 参数 — 自动查询 trade_cal 获取交易日\n", + "# 如果中断,修改年份范围从断点继续即可 (UPSERT 幂等,不会重复)\n", "import_daily_by_year(\n", - " stock_list,\n", " start_year=2010,\n", " end_year=2025,\n", + " sleep_interval=0.3, # API 频率控制,免费版建议 0.3~0.5\n", ")" ] }, @@ -352,9 +290,9 @@ "cell_type": "markdown", "metadata": {}, "source": [ - "### 4.2.1 查看日线导入进度\n", + "### 4.1.1 查看日线导入进度\n", "\n", - "按年份统计 daily 表中已导入的记录数和独立股票数,用于确认导入到哪个年份了。" + "按年份统计 daily 表中已导入的记录数和独立股票数,了解导入到哪个年份了。" ] }, { @@ -373,7 +311,7 @@ "metadata": {}, "outputs": [], "source": [ - "# 或者查看所有表的整体概览\n", + "# 查看所有表的整体概览\n", "check_table_summary()" ] }, @@ -381,11 +319,14 @@ "cell_type": "markdown", "metadata": {}, "source": [ - "### 4.2.2 断点续传 — 从中断处继续导入\n", + "### 4.1.2 按缺失日期断点续传 (推荐)\n", "\n", - "**方法一(推荐):** 使用 `resume_daily_by_year` 自动跳过已完成的年份,只导入缺失年份。\n", + "新版 `resume_daily_by_date` 自动对比 `trade_cal` 和 `daily` 表,**只导入缺失日期的全市场数据**。\n", + "粒度精确到交易日级别,比旧的按年份续传更精细。\n", "\n", - "**方法二:** 手动修改 `start_year` 参数重新调用 `import_daily_by_year`(因为 UPSERT 幂等,重复导入不会造成数据问题)。" + "**工作流程:**\n", + "1. 先调用 `get_missing_daily_dates()` 查看缺失的交易日列表\n", + "2. 调用 `resume_daily_by_date()` 仅补缺缺失的交易日" ] }, { @@ -394,9 +335,21 @@ "metadata": {}, "outputs": [], "source": [ - "# 方法一:自动检测已有年份,只导入缺失年份(推荐)\n", - "stock_list = get_stock_codes_from_db()\n", - "resume_daily_by_year(stock_list, start_year=2010, end_year=2025)" + "# 第1步:查看缺失的交易日 (仅查询,不导入)\n", + "missing = get_missing_daily_dates(\n", + " start_date=\"2010-01-01\",\n", + " end_date=\"2025-12-31\",\n", + ")\n", + "print(f\"缺失交易日总数: {len(missing)}\")\n", + "if len(missing) <= 30:\n", + " print(f\"缺失日期: {missing}\")\n", + "else:\n", + " # 按年份汇总显示\n", + " from collections import Counter\n", + " year_counts = Counter(d[:4] for d in missing)\n", + " for y in sorted(year_counts):\n", + " print(f\" {y} 年: {year_counts[y]} 个缺失交易日\")\n", + " print(f\" (前10个缺失日期): {missing[:10]}\")" ] }, { @@ -405,24 +358,36 @@ "metadata": {}, "outputs": [], "source": [ - "# 方法二:手动指定断点年份重新调用 import_daily_by_year\n", - "# 例如假设 2010~2020 已完成,从 2021 年开始继续\n", - "# stock_list = get_stock_codes_from_db()\n", - "# import_daily_by_year(stock_list, start_year=2021, end_year=2025)" + "# 第2步:按缺失日期断点续传,自动补充缺失的交易日数据\n", + "# 例如之前中断了,这里只会导入尚未导入的交易日数据\n", + "resume_daily_by_date(\n", + " start_date=\"2010-01-01\",\n", + " end_date=\"2025-12-31\",\n", + " sleep_interval=0.3,\n", + ")" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ - "> **提示:** 以上两个方法都可以安全使用。因为 daily 表使用 `ON CONFLICT (ts_code, trade_date) DO UPDATE`,重复导入已存在的数据不会产生重复记录。另外也可查看 `import_data.log` 文件获取最后一次成功的日志输出。" + "> **备用方案:** 也可以手动指定断点年份重新调用 `import_daily_by_year`,因为 UPSERT 幂等,重复导入已存在的数据不会产生重复记录。\n", + ">\n", + "> ```python\n", + "> # 例如假设 2010~2020 已完成,从 2021 年继续\n", + "> import_daily_by_year(start_year=2021, end_year=2025)\n", + "> ```\n", + ">\n", + "> 也可查看 `import_data.log` 文件获取最后一次成功的日志输出。" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ - "### 4.3 单只股票导入 (补充/重试)" + "### 4.2 单只股票补充导入\n", + "\n", + "如果某只股票数据缺失,可以单独补导(保留原接口兼容)。" ] }, { @@ -433,8 +398,6 @@ "source": [ "# 导入单只股票日线行情 (用于补充导入或测试)\n", "conn = get_pg_connection()\n", - "from importer import import_daily_for_stock\n", - "\n", "n = import_daily_for_stock(\"000001.SZ\", \"2020-01-01\", \"2020-12-31\", conn)\n", "print(f\"导入 000001.SZ 2020年数据: {n} 条\")\n", "conn.close()" @@ -517,6 +480,10 @@ "outputs": [], "source": [ "# 导入复权因子 (用于前复权/后复权价格计算)\n", + "# 需要先获取股票列表\n", + "stock_list = get_stock_codes_from_db()\n", + "print(f\"共 {len(stock_list)} 只股票\")\n", + "\n", "import_adj_factor_batch(\n", " stock_list,\n", " start_date=\"2010-01-01\",\n", @@ -541,6 +508,10 @@ "outputs": [], "source": [ "# 导入利润表、资产负债表、现金流量表、财务指标\n", + "# stock_list 从上一步已获取,或重新获取\n", + "if 'stock_list' not in dir():\n", + " stock_list = get_stock_codes_from_db()\n", + "\n", "import_financial_statements(\n", " stock_list,\n", " start_date=\"2010-01-01\",\n", @@ -613,7 +584,7 @@ "metadata": {}, "outputs": [], "source": [ - "# 一键全量导入 (需数小时~数十小时,请谨慎)\n", + "# 一键全量导入 (需数小时,请谨慎)\n", "# full_import(\n", "# start_date=\"2010-01-01\",\n", "# end_date=\"2025-12-31\",\n", @@ -768,4 +739,4 @@ }, "nbformat": 4, "nbformat_minor": 4 -} +} \ No newline at end of file