CoolFace
Apppublic

Qionk/a-share-quant

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
data_input.py382 linesDownload Raw Back to predict
1"""2价格预测 - 数据输入3支持 AKShare API 获取 和 Excel 手动上传两种方式4akshare 失败时自动切换到腾讯财经备用源,并尝试通过网易财经补充换手率5"""6 7import io8import time9import json10import subprocess11import pandas as pd12import numpy as np13import akshare as ak14 15REQUIRED_COLUMNS = ["日期", "开盘", "最高", "最低", "收盘", "成交量", "成交额", "涨跌幅"]16 17COLUMN_MAP = {18    "日期": "date", "开盘": "open", "最高": "high", "最低": "low",19    "收盘": "close", "成交量": "volume", "成交额": "amount", "涨跌幅": "pct_change",20}21 22MAX_RETRIES = 223RETRY_DELAY = 124 25 26def _market_prefix(stock_code: str) -> str:27    """根据股票代码判断市场前缀"""28    if stock_code.startswith(("6", "9")):29        return "sh"30    return "sz"31 32 33def _fetch_via_tencent(stock_code: str, start_date: str, end_date: str) -> pd.DataFrame:34    """35    备用数据源:通过腾讯财经 newfqkline 接口获取前复权日K数据(含换手率)36    """37    import requests as _req38 39    prefix = _market_prefix(stock_code)40    symbol = f"{prefix}{stock_code}"41    sd = f"{start_date[:4]}-{start_date[4:6]}-{start_date[6:8]}"42    ed = f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}"43 44    all_klines = []45    current_end = ed46    for _ in range(20):47        # newfqkline 接口:返回含换手率的完整日K48        url = (f"https://proxy.finance.qq.com/ifzqgtimg/appstock/app/newfqkline/get"49               f"?param={symbol},day,{sd},{current_end},800,qfq")50        try:51            resp = _req.get(url, timeout=10)52            if resp.status_code != 200:53                break54            data = resp.json()55        except Exception:56            # fallback 旧接口57            try:58                url_old = (f"https://web.ifzq.gtimg.cn/appstock/app/fqkline/get"59                           f"?param={symbol},day,{sd},{current_end},800,qfq")60                resp = _req.get(url_old, timeout=10)61                data = resp.json()62            except Exception:63                break64 65        stock_data = data.get("data", {}).get(symbol, {})66        klines = stock_data.get("qfqday") or stock_data.get("day", [])67        if not klines:68            break69 70        all_klines = klines + all_klines71        earliest = klines[0][0]72        if earliest <= sd:73            break74        from datetime import datetime as dt, timedelta75        prev = (dt.strptime(earliest, "%Y-%m-%d") - timedelta(days=1)).strftime("%Y-%m-%d")76        current_end = prev77        time.sleep(0.3)78 79    if not all_klines:80        return pd.DataFrame()81 82    # 补齐尾部83    last_fetched = all_klines[-1][0] if all_klines else None84    if last_fetched and last_fetched < ed:85        url = (f"https://proxy.finance.qq.com/ifzqgtimg/appstock/app/newfqkline/get"86               f"?param={symbol},day,{last_fetched},{ed},30,qfq")87        try:88            resp = _req.get(url, timeout=10)89            if resp.status_code == 200:90                data2 = resp.json()91                klines2 = data2.get("data", {}).get(symbol, {})92                klines2 = klines2.get("qfqday") or klines2.get("day", [])93                if klines2:94                    extra = [k for k in klines2 if k[0] > last_fetched]95                    all_klines.extend(extra)96        except Exception:97            pass98 99    # 去重(按日期)100    seen = set()101    unique = []102    for k in all_klines:103        if k[0] not in seen:104            seen.add(k[0])105            unique.append(k)106    unique.sort(key=lambda x: x[0])107 108    # newfqkline 字段: [date, open, close, high, low, volume, {ma}, turnover%, amount, ?]109    # 旧接口字段:      [date, open, close, high, low, volume]110    rows = []111    for row in unique:112        entry = {113            "date": row[0],114            "open": row[1],115            "close": row[2],116            "high": row[3],117            "low": row[4],118            "volume": row[5],119        }120        # newfqkline 有更多字段(第7个是dict/ma,第8个是换手率)121        if len(row) >= 8 and not isinstance(row[7], dict):122            try:123                turnover_val = float(row[7])124                if turnover_val > 0:125                    entry["turnover"] = turnover_val126            except (ValueError, TypeError):127                pass128        if len(row) >= 9:129            try:130                amount_val = float(row[8]) if not isinstance(row[8], dict) else None131                if amount_val and amount_val > 0:132                    entry["amount"] = amount_val * 10000  # 腾讯单位是万元133            except (ValueError, TypeError):134                pass135        rows.append(entry)136 137    df = pd.DataFrame(rows)138    df["date"] = pd.to_datetime(df["date"])139    df = df.sort_values("date").set_index("date")140 141    for col in ["open", "close", "high", "low", "volume"]:142        if col in df.columns:143            df[col] = pd.to_numeric(df[col], errors="coerce")144 145    # 计算缺失的列146    df["pct_change"] = df["close"].pct_change() * 100147    if "amount" not in df.columns:148        df["amount"] = df["close"] * df["volume"] * 100149 150    return df151 152 153def _fetch_turnover_netease(stock_code: str, start_date: str, end_date: str) -> pd.Series:154    """155    通过网易财经接口获取换手率数据,返回 Series(index=DatetimeIndex, values=turnover%)。156    失败返回 None。157    """158    try:159        prefix = "0" + stock_code if stock_code.startswith(("6", "9")) else "1" + stock_code160        sd = f"{start_date[:4]}{start_date[4:6]}{start_date[6:8]}"161        ed = f"{end_date[:4]}{end_date[4:6]}{end_date[6:8]}"162        url = (f"https://quotes.money.163.com/service/chddata.html"163               f"?code={prefix}&start={sd}&end={ed}&fields=TURNOVER")164 165        content = None166        # 优先用 requests(Streamlit Cloud 无 curl)167        try:168            import requests169            resp = requests.get(url, timeout=15)170            if resp.status_code == 200 and resp.content:171                content = resp.content.decode("gb2312", errors="ignore")172        except Exception:173            pass174        # fallback: curl175        if not content:176            result = subprocess.run(177                ["curl", "-s", "-m", "15", "-L", url],178                capture_output=True, text=True, timeout=20,179            )180            if result.returncode == 0 and result.stdout.strip():181                content = result.stdout182 183        if not content:184            return None185        from io import StringIO186        ndf = pd.read_csv(StringIO(content), engine="python")187        if ndf.empty:188            return None189        # 网易列名:日期, 股票代码, 名称, 换手率190        date_col = ndf.columns[0]191        turnover_col = [c for c in ndf.columns if "换手" in c]192        if not turnover_col:193            return None194        ndf[date_col] = pd.to_datetime(ndf[date_col])195        ndf = ndf.sort_values(date_col).set_index(date_col)196        s = pd.to_numeric(ndf[turnover_col[0]], errors="coerce")197        s.index.name = None198        return s199    except Exception:200        return None201 202 203def load_from_akshare(stock_code: str, start_date: str, end_date: str) -> pd.DataFrame:204    """205    获取个股日线数据206    优先用腾讯财经 newfqkline(稳定、含换手率),失败后切换 AKShare 备用207    """208    last_err = None209 210    # 主源:腾讯财经 newfqkline(前复权,含换手率)211    try:212        df = _fetch_via_tencent(stock_code, start_date, end_date)213        if df is not None and not df.empty:214            if "pct_change" not in df.columns or df["pct_change"].isna().all():215                df["pct_change"] = df["close"].pct_change() * 100216            return df217    except Exception as e:218        last_err = e219 220    # 备用源:AKShare(东方财富,可能被限流)221    for attempt in range(MAX_RETRIES):222        try:223            df = ak.stock_zh_a_hist(224                symbol=stock_code, period="daily",225                start_date=start_date, end_date=end_date,226                adjust="qfq",227            )228            if df is not None and not df.empty:229                df = df.rename(columns={230                    "日期": "date", "开盘": "open", "收盘": "close",231                    "最高": "high", "最低": "low", "成交量": "volume",232                    "成交额": "amount", "涨跌幅": "pct_change", "换手率": "turnover",233                })234                df["date"] = pd.to_datetime(df["date"])235                df = df.sort_values("date").set_index("date")236                needed = ["open", "high", "low", "close", "volume", "amount", "pct_change"]237                for col in needed:238                    if col not in df.columns:239                        if col == "pct_change":240                            df["pct_change"] = df["close"].pct_change() * 100241                return df242        except Exception as e:243            last_err = e244            if attempt < MAX_RETRIES - 1:245                time.sleep(RETRY_DELAY * (attempt + 1))246 247    msg = f"未获取到 {stock_code} 的数据(主备数据源均失败)"248    if last_err:249        msg += f"\n原始错误: {type(last_err).__name__}: {last_err}"250    raise ValueError(msg)251 252 253STOCK_NAME_CACHE = {254    "601869": "长飞光纤",255    "603601": "再升科技",256    "601138": "工业富联",257}258 259 260def get_stock_name(stock_code: str) -> str:261    """查询股票名称(优先用缓存,失败时回退到代码)"""262    if stock_code in STOCK_NAME_CACHE:263        return STOCK_NAME_CACHE[stock_code]264    try:265        info = ak.stock_info_a_code_name()266        info.columns = ["code", "name"]267        match = info[info["code"] == stock_code]268        if not match.empty:269            name = match.iloc[0]["name"]270            STOCK_NAME_CACHE[stock_code] = name271            return name272    except Exception:273        pass274    return stock_code275 276 277def validate_dataframe(df: pd.DataFrame) -> tuple:278    """279    验证上传的 DataFrame 格式280    返回: (是否通过, 错误信息列表)281    """282    errors = []283 284    missing = [c for c in REQUIRED_COLUMNS if c not in df.columns]285    if missing:286        errors.append(f"缺少必需列: {', '.join(missing)}")287        return False, errors288 289    try:290        pd.to_datetime(df["日期"])291    except Exception:292        errors.append("'日期' 列格式无法解析,请使用 YYYY-MM-DD 或 YYYYMMDD 格式")293 294    numeric_cols = [c for c in REQUIRED_COLUMNS if c != "日期"]295    for col in numeric_cols:296        if not pd.api.types.is_numeric_dtype(df[col]):297            try:298                pd.to_numeric(df[col])299            except Exception:300                errors.append(f"'{col}' 列包含非数值数据")301 302    if len(df) < 30:303        errors.append(f"数据行数不足: {len(df)} 行(至少需要 30 行)")304 305    na_counts = df[REQUIRED_COLUMNS].isna().sum()306    cols_with_na = na_counts[na_counts > 0]307    if not cols_with_na.empty:308        for col, cnt in cols_with_na.items():309            errors.append(f"'{col}' 列有 {cnt} 个缺失值")310 311    return len(errors) == 0, errors312 313 314def load_from_excel(uploaded_file) -> pd.DataFrame:315    """316    读取上传的 Excel 文件并标准化317    返回标准化 DataFrame(date 为 DatetimeIndex)318    """319    for encoding in ["utf-8", "gbk", "gb2312"]:320        try:321            df = pd.read_excel(uploaded_file, engine=None)322            break323        except Exception:324            uploaded_file.seek(0)325            continue326    else:327        raise ValueError("无法读取 Excel 文件,请确认文件格式正确")328 329    ok, errors = validate_dataframe(df)330    if not ok:331        raise ValueError("Excel 数据验证失败:\n" + "\n".join(f"  - {e}" for e in errors))332 333    result = pd.DataFrame()334    for cn_col, en_col in COLUMN_MAP.items():335        if cn_col in df.columns:336            result[en_col] = df[cn_col]337 338    result["date"] = pd.to_datetime(result["date"])339    result = result.sort_values("date").set_index("date")340 341    for col in ["open", "high", "low", "close", "volume", "amount", "pct_change"]:342        if col in result.columns:343            result[col] = pd.to_numeric(result[col], errors="coerce")344 345    result = result.dropna(subset=["close"])346    return result347 348 349def generate_template() -> bytes:350    """生成可下载的 Excel 模板(含示例数据和说明)"""351    dates = pd.bdate_range("2024-01-02", periods=5)352    sample = pd.DataFrame({353        "日期": dates.strftime("%Y-%m-%d"),354        "开盘": [10.50, 10.80, 10.60, 10.90, 11.00],355        "最高": [10.90, 10.95, 10.85, 11.10, 11.20],356        "最低": [10.40, 10.55, 10.50, 10.75, 10.90],357        "收盘": [10.80, 10.60, 10.80, 11.05, 11.10],358        "成交量": [500000, 450000, 520000, 600000, 550000],359        "成交额": [5300000, 4800000, 5500000, 6500000, 6100000],360        "涨跌幅": [1.50, -1.85, 1.89, 2.31, 0.45],361    })362 363    buf = io.BytesIO()364    with pd.ExcelWriter(buf, engine="xlsxwriter") as writer:365        sample.to_excel(writer, sheet_name="日线数据", index=False)366 367        instructions = pd.DataFrame({368            "说明": [369                "请按照'日线数据'工作表的格式填入数据",370                "日期格式: YYYY-MM-DD",371                "涨跌幅单位: %(如 1.50 表示涨 1.50%)",372                "成交量单位: 股",373                "成交额单位: 元",374                "所有列均为必填项",375                "至少需要 30 个交易日的数据",376            ]377        })378        instructions.to_excel(writer, sheet_name="填写说明", index=False)379 380    buf.seek(0)381    return buf.getvalue()382