Qionk/a-share-quant
0
1#!/usr/bin/env python32"""3清理数据库并重新拉取股票数据(只保留最近 500 个交易日)4 5步骤:6 1. 列出当前数据库中所有股票代码7 2. 清空 stock_daily_data / training_sessions / model_results 三张表8 3. 逐只股票从 AKShare 拉取全量数据,截取最后 500 天写入 MySQL9 10用法:11 set MYSQL_HOST=mysql3.sqlpub.com12 set MYSQL_PORT=330813 set MYSQL_USER=root_quant14 set MYSQL_PASSWORD=BLnVlQ8qASfhA9xZ15 set MYSQL_DATABASE=a_share_quant16 17 python scripts/clean_and_refetch.py18"""19import sys, os, time20from datetime import datetime21 22sys.stdout.reconfigure(line_buffering=True) if hasattr(sys.stdout, "reconfigure") else None23 24ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))25sys.path.insert(0, ROOT)26 27os.environ.setdefault("MYSQL_HOST", os.environ.get("MYSQL_HOST", ""))28os.environ.setdefault("MYSQL_PORT", os.environ.get("MYSQL_PORT", "3306"))29os.environ.setdefault("MYSQL_USER", os.environ.get("MYSQL_USER", ""))30os.environ.setdefault("MYSQL_PASSWORD", os.environ.get("MYSQL_PASSWORD", ""))31os.environ.setdefault("MYSQL_DATABASE", os.environ.get("MYSQL_DATABASE", ""))32 33import pymysql34import pandas as pd35import numpy as np36 37MAX_DAYS = 500 # 每只股票最多保留的交易天数38 39print(f"\n{'='*60}")40print(f" 数据库清理 & 重新拉取 (保留最近 {MAX_DAYS} 天)")41print(f" 时间: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")42print(f"{'='*60}")43 44 45# ── 1. 连接数据库 ──────────────────────────────────────46def get_conn():47 host = os.environ.get("MYSQL_HOST", "")48 port = int(os.environ.get("MYSQL_PORT", "3306"))49 user = os.environ.get("MYSQL_USER", "")50 password = os.environ.get("MYSQL_PASSWORD", "")51 database = os.environ.get("MYSQL_DATABASE", "")52 53 if not host or not user or not database:54 print("[错误] MySQL 环境变量未完整设置!")55 sys.exit(1)56 57 return pymysql.connect(58 host=host, port=port, user=user, password=password,59 database=database, charset="utf8mb4",60 connect_timeout=10, read_timeout=30, write_timeout=30,61 autocommit=True,62 )63 64 65conn = get_conn()66cur = conn.cursor()67print("\n[OK] MySQL 连接成功")68 69 70# ── 2. 列出当前所有股票 ────────────────────────────────71print("\n── 当前数据库中的股票 ──")72cur.execute("""73 SELECT stock_code, stock_name, COUNT(*) AS cnt,74 MIN(trade_date) AS start_date, MAX(trade_date) AS end_date75 FROM stock_daily_data76 GROUP BY stock_code, stock_name77 ORDER BY stock_code78""")79existing = [(r[0], r[1], r[2], r[3], r[4]) for r in cur.fetchall()]80 81if not existing:82 print(" 数据库中没有股票数据,无需清理")83 print(" 请先在 Streamlit 网页中添加股票,或手动指定股票代码列表")84 conn.close()85 sys.exit(0)86 87stock_codes = []88for code, name, cnt, start, end in existing:89 s_str = str(start) if hasattr(start, "strftime") else str(start)90 e_str = str(end) if hasattr(end, "strftime") else str(end)91 print(f" {code} {name:<12} {cnt}天 ({s_str} ~ {e_str})")92 stock_codes.append(code)93 94print(f"\n共 {len(stock_codes)} 只股票")95 96# ── 2b. 查看训练数据 ────────────────────────────────────97cur.execute("SELECT COUNT(*) FROM training_sessions")98ts_count = cur.fetchone()[0]99cur.execute("SELECT COUNT(*) FROM model_results")100mr_count = cur.fetchone()[0]101print(f"训练记录: {ts_count} 条 sessions, {mr_count} 条 model_results")102 103# ── 3. 确认清库 ─────────────────────────────────────────104print(f"\n{'!'*60}")105print(f" 即将执行以下操作:")106print(f" 1. 清空 stock_daily_data 表 ({sum(c for _,_,c,_,_ in existing)} 条)")107print(f" 2. 清空 training_sessions 表 ({ts_count} 条)")108print(f" 3. 清空 model_results 表 ({mr_count} 条)")109print(f" 4. 重新拉取 {len(stock_codes)} 只股票,各保留最近 {MAX_DAYS} 天")110print(f"{'!'*60}")111 112confirm = input("\n确认执行? (输入 yes 继续): ").strip().lower()113if confirm != "yes":114 print("已取消")115 conn.close()116 sys.exit(0)117 118 119# ── 4. 清空表 ───────────────────────────────────────────120print("\n── 清理数据库 ──")121 122tables = ["model_results", "training_sessions", "stock_daily_data"]123for table in tables:124 t0 = time.time()125 cur.execute(f"DELETE FROM {table}")126 deleted = cur.rowcount127 print(f" [OK] 清空 {table}: {deleted} 条 ({time.time()-t0:.1f}s)")128 129print(" 清理完成")130 131 132# ── 5. 重新拉取数据 ────────────────────────────────────133print(f"\n── 重新拉取数据 (每只股票保留最近 {MAX_DAYS} 天) ──")134 135from src.predict.data_input import load_from_akshare, get_stock_name136 137INSERT_SQL = """138 INSERT INTO stock_daily_data139 (stock_code, stock_name, trade_date, open, high, low, close, volume, amount, pct_change, turnover)140 VALUES (%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s)141"""142 143def safe_float(v):144 if v is None:145 return None146 try:147 f = float(v)148 return None if (np.isnan(f) or np.isinf(f)) else f149 except (TypeError, ValueError):150 return None151 152success_list = []153fail_list = []154 155for idx, code in enumerate(stock_codes):156 name = ""157 try:158 print(f"\n[{idx+1}/{len(stock_codes)}] {code} ...")159 160 # 获取股票名称161 try:162 name = get_stock_name(code)163 except Exception:164 name = code165 166 # 拉取全量数据(AKShare 默认返回全部历史)167 df = load_from_akshare(code, "20200101", datetime.now().strftime("%Y%m%d"))168 if df is None or df.empty:169 fail_list.append((code, name, "无数据返回"))170 print(f" [失败] 无数据")171 continue172 173 total_days = len(df)174 print(f" 拉取到 {total_days} 天数据")175 176 # 截取最后 MAX_DAYS 天177 if len(df) > MAX_DAYS:178 df = df.tail(MAX_DAYS)179 print(f" 截取最近 {MAX_DAYS} 天: "180 f"{df.index[0].strftime('%Y-%m-%d')} ~ {df.index[-1].strftime('%Y-%m-%d')}")181 else:182 print(f" 数据不足 {MAX_DAYS} 天,全部保留")183 184 # 写入 MySQL185 rows = []186 for idx_row, row in df.iterrows():187 trade_date = idx_row.date() if hasattr(idx_row, "date") else pd.Timestamp(idx_row).date()188 rows.append((189 code, name, trade_date,190 safe_float(row.get("open")),191 safe_float(row.get("high")),192 safe_float(row.get("low")),193 safe_float(row.get("close")),194 safe_float(row.get("volume")),195 safe_float(row.get("amount")),196 safe_float(row.get("pct_change")),197 safe_float(row.get("turnover")),198 ))199 200 batch_size = 500201 for i in range(0, len(rows), batch_size):202 batch = rows[i:i + batch_size]203 cur.executemany(INSERT_SQL, batch)204 205 last_close = df["close"].iloc[-1]206 success_list.append((code, name, len(rows)))207 print(f" [OK] {name} ({code}) 写入 {len(rows)} 天 "208 f"最新价 {last_close:.2f} "209 f"{df.index[0].strftime('%Y-%m-%d')} ~ {df.index[-1].strftime('%Y-%m-%d')}")210 211 # 请求间隔,避免触发反爬212 time.sleep(1)213 214 except Exception as e:215 fail_list.append((code, name, str(e)))216 print(f" [失败] {e}")217 import traceback218 traceback.print_exc()219 220 221# ── 6. 汇总 ─────────────────────────────────────────────222print(f"\n{'='*60}")223print(f" 完成!")224print(f"{'='*60}")225print(f" 成功: {len(success_list)}/{len(stock_codes)} 只")226print(f" 失败: {len(fail_list)}/{len(stock_codes)} 只")227print(f"")228 229if success_list:230 total_rows = sum(r for _,_,r in success_list)231 print(f" 成功列表 ({total_rows} 条总记录):")232 for code, name, cnt in success_list:233 print(f" [OK] {code} {name} ({cnt} 天)")234 235if fail_list:236 print(f" 失败列表:")237 for code, name, err in fail_list:238 print(f" [FAIL] {code} {name} ({err})")239 240cur.close()241conn.close()242print(f"\nDone.\n")