CoolFace
Apppublic

Qionk/a-share-quant

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
clean_and_refetch.py242 linesDownload Raw Back to scripts
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")