CoolFace
Apppublic

claraleeee/hftpredict

sourceHugging Faceupdated 10mo agoView on Hugging Face
0likes
predict.py267 linesDownload Raw Back to root
1# predict.py2# Universal prediction module for A-share, ETF, and US stocks3# Returns DataFrame with columns: datetime, predicted_price, model4 5import os6import numpy as np7import pandas as pd8from datetime import datetime9from sklearn.preprocessing import MinMaxScaler10 11import akshare as ak12import yfinance as yf13 14import tensorflow as tf15from tensorflow.keras import layers, models, callbacks16import torch17from torch import nn18 19# ---------- CONFIG ----------20TIME_STEP = 6021FUTURE_MINUTES_A_SH = 24222FUTURE_MINUTES_US = 39023HISTORY_DAYS = 3024EPOCHS_META = 425EPOCHS_TFT = 626SEED = 4227np.random.seed(SEED)28tf.random.set_seed(SEED)29torch.manual_seed(SEED)30 31 32# ---------- UTILITIES ----------33def build_a_stock_minutes(code, days=HISTORY_DAYS):34    """Fetch A-share minute-level data robustly and normalize column names."""35    print(f"[DBG] fetching A股分钟数据: {code}, days={days}")36    try:37        df = ak.stock_zh_a_hist_min_em(symbol=code, period="1", adjust="qfq")38    except Exception as e:39        print(f"[WARN] hist_min_em failed ({e}), fallback to stock_zh_a_minute")40        df = ak.stock_zh_a_minute(symbol=code, period="1", adjust="qfq")41 42    if df is None or df.shape[0] == 0:43        raise RuntimeError(f"No data fetched for {code}")44 45    # try to rename common columns to standard english names46    rename_map = {47        "时间": "datetime", "日期时间": "datetime", "day": "datetime",48        "开盘": "open", "收盘": "close", "最高": "high", "最低": "low", "成交量": "volume",49        "开盘价": "open", "收盘价": "close", "最高价": "high", "最低价": "low", "成交额": "volume"50    }51    df = df.rename(columns={k: v for k, v in rename_map.items() if k in df.columns})52 53    if "datetime" not in df.columns:54        # maybe already english names55        if "time" in df.columns:56            df = df.rename(columns={"time": "datetime"})57        elif "datetime" in df.columns:58            pass59        else:60            raise RuntimeError(f"Cannot find time column in data: {df.columns.tolist()}")61 62    df["datetime"] = pd.to_datetime(df["datetime"], errors="coerce")63    for c in ["open", "high", "low", "close", "volume"]:64        if c in df.columns:65            df[c] = pd.to_numeric(df[c], errors="coerce")66 67    df = df.dropna(subset=["close"]).sort_values("datetime").reset_index(drop=True)68    latest = df["datetime"].max()69    start = latest - pd.Timedelta(days=days)70    df = df[df["datetime"] >= start].reset_index(drop=True)71    print(f"[DBG] rows={len(df)}, columns={list(df.columns)}")72    return df73 74 75def build_us_minutes(ticker, days=HISTORY_DAYS):76    """Fetch US minute data using yfinance and standardize columns."""77    print(f"[DBG] fetching US minute data for {ticker}")78    df = yf.download(ticker, period=f"{days}d", interval="1m", progress=False)79    if df is None or df.shape[0] == 0:80        raise RuntimeError(f"yfinance returned no data for {ticker}")81    df = df.reset_index().rename(columns={82        "Datetime": "datetime", "Open": "open", "High": "high", "Low": "low",83        "Close": "close", "Volume": "volume"84    })85    df["datetime"] = pd.to_datetime(df["datetime"])86    df = df.dropna(subset=["close"]).reset_index(drop=True)87    return df88 89 90def trading_minutes_for_date(predict_date, market="cn"):91    """Return exact trading minute timeline for the given predict_date (strict)."""92    date = pd.to_datetime(predict_date).strftime("%Y-%m-%d")93    if market == "cn":94        morning = pd.date_range(f"{date} 09:30", f"{date} 11:30", freq="1min")95        afternoon = pd.date_range(f"{date} 13:00", f"{date} 15:00", freq="1min")96        timeline = morning.append(afternoon)97        # enforce FUTURE_MINUTES_A_SH length98        return timeline[:FUTURE_MINUTES_A_SH]99    else:100        timeline = pd.date_range(f"{date} 09:30", f"{date} 16:00", freq="1min")101        return timeline[:FUTURE_MINUTES_US]102 103 104# ---------- META (A股) ----------105def predict_meta_a(stock_code, predict_date, open_price=None):106    print(f"[DBG] [META] start {stock_code} for {predict_date}")107    df = build_a_stock_minutes(stock_code)108    df["return_1"] = df["close"].pct_change().fillna(0)109    df["ma5"] = df["close"].rolling(5).mean().bfill()110    df["ma20"] = df["close"].rolling(20).mean().bfill()111    df["vma5"] = df["volume"].rolling(5).mean().bfill()112    df["vma20"] = df["volume"].rolling(20).mean().bfill()113 114    features = ["close", "volume", "return_1", "ma5", "ma20", "vma5", "vma20"]115    df = df.dropna(subset=features)116    scaler = MinMaxScaler().fit(df[features])117    Xs = scaler.transform(df[features])118 119    FUT = FUTURE_MINUTES_A_SH120    X, Y = [], []121    for i in range(TIME_STEP, len(Xs) - FUT):122        X.append(Xs[i - TIME_STEP:i])123        Y.append(Xs[i:i + FUT, 0])124    X, Y = np.array(X), np.array(Y)125 126    if len(X) < 2:127        raise RuntimeError("META: insufficient data")128 129    tf.keras.backend.clear_session()130    def build_model():131        inp = layers.Input(shape=(TIME_STEP, X.shape[2]))132        x = layers.LSTM(64, return_sequences=True)(inp)133        x = layers.LSTM(32)(x)134        out = layers.Dense(FUT)(x)135        return models.Model(inp, out)136 137    model = build_model()138    model.compile(optimizer="adam", loss="mse")139    model.fit(X[:-1], Y[:-1], epochs=EPOCHS_META, batch_size=32, verbose=0)140 141    pred_scaled = model.predict(X[-1:])[0]142    tmp = np.zeros((len(pred_scaled), len(features)))143    tmp[:, 0] = pred_scaled144    # inverse-transform entire feature-space placeholder then take close column145    pred_prices = scaler.inverse_transform(tmp)[:, 0]146 147    if open_price is not None and len(pred_prices) > 0:148        pred_prices = open_price * (pred_prices / (pred_prices[0] if pred_prices[0] != 0 else 1.0))149 150    timeline = trading_minutes_for_date(predict_date, market="cn")151    pred_prices = pred_prices[:len(timeline)]152    return pd.DataFrame({"datetime": timeline, "predicted_price": pred_prices, "model": "META"})153 154 155# ---------- TFT (A股) ----------156class SimpleTFT(nn.Module):157    def __init__(self, n_features, hidden=64, heads=4, out_steps=FUTURE_MINUTES_A_SH):158        super().__init__()159        self.proj = nn.Linear(n_features, hidden)160        layer = nn.TransformerEncoderLayer(d_model=hidden, nhead=heads)161        self.enc = nn.TransformerEncoder(layer, num_layers=1)162        self.out = nn.Linear(hidden, out_steps)163    def forward(self, x):164        x = self.proj(x)165        x = x.permute(1, 0, 2)166        x = self.enc(x)167        x = x[-1]168        return self.out(x)169 170def predict_tft_a(stock_code, predict_date, open_price=None):171    print(f"[DBG] [TFT] start {stock_code} for {predict_date}")172    df = build_a_stock_minutes(stock_code)173    df["return_1"] = df["close"].pct_change().fillna(0)174    df["ma5"] = df["close"].rolling(5).mean().bfill()175    df["ma20"] = df["close"].rolling(20).mean().bfill()176    df["vma5"] = df["volume"].rolling(5).mean().bfill()177    df["vma20"] = df["volume"].rolling(20).mean().bfill()178 179    features = ["close", "volume", "return_1", "ma5", "ma20", "vma5", "vma20"]180    df = df.dropna(subset=features)181    scaler = MinMaxScaler().fit(df[features])182    scaled = scaler.transform(df[features])183 184    X = np.array([scaled[i - TIME_STEP:i] for i in range(TIME_STEP, len(scaled))])185    if X.shape[0] < 1:186        raise RuntimeError("TFT: insufficient windows")187    model = SimpleTFT(n_features=X.shape[2])188    opt = torch.optim.Adam(model.parameters(), lr=1e-3)189    loss_fn = nn.MSELoss()190    y = scaled[TIME_STEP:, 0]191    X_train = torch.tensor(X[:-1], dtype=torch.float32) if X.shape[0] > 1 else torch.tensor(X, dtype=torch.float32)192    y_train = torch.tensor(y[:-1], dtype=torch.float32) if X.shape[0] > 1 else torch.tensor(y, dtype=torch.float32)193 194    model.train()195    for ep in range(EPOCHS_TFT):196        perm = np.random.permutation(len(X_train))197        for i in perm:198            xb = X_train[i:i+1]; yb = y_train[i:i+1]199            opt.zero_grad(); pr = model(xb)200            loss = loss_fn(pr[:, 0], yb); loss.backward(); opt.step()201 202    model.eval()203    with torch.no_grad():204        last_seq = torch.tensor(scaled[-TIME_STEP:], dtype=torch.float32).unsqueeze(0)205        pred_scaled = model(last_seq).cpu().numpy().flatten()206 207    tmp = np.zeros((len(pred_scaled), len(features)))208    tmp[:, 0] = pred_scaled209    pred_prices = scaler.inverse_transform(tmp)[:, 0]210 211    if open_price is not None and len(pred_prices) > 0:212        pred_prices = open_price * (pred_prices / (pred_prices[0] if pred_prices[0] != 0 else 1.0))213 214    timeline = trading_minutes_for_date(predict_date, market="cn")215    pred_prices = pred_prices[:len(timeline)]216    return pd.DataFrame({"datetime": timeline, "predicted_price": pred_prices, "model": "TFT"})217 218 219# ---------- META (US) ----------220def predict_meta_us(ticker, predict_date, open_price=None):221    print(f"[DBG] [META_US] start {ticker} for {predict_date}")222    df = build_us_minutes(ticker)223    df["return_1"] = df["close"].pct_change().fillna(0)224    df["ma5"] = df["close"].rolling(5).mean().bfill()225    df["ma20"] = df["close"].rolling(20).mean().bfill()226    features = ["close", "return_1", "ma5", "ma20"]227    df = df.dropna(subset=features)228    Xs = MinMaxScaler().fit_transform(df[features])229 230    FUT = FUTURE_MINUTES_US231    X = [Xs[i - TIME_STEP:i] for i in range(TIME_STEP, len(Xs) - FUT)]232    if len(X) < 2:233        raise RuntimeError("META_US: insufficient data")234 235    X = np.array(X)236    model = models.Sequential([237        layers.Input(shape=(TIME_STEP, X.shape[2])),238        layers.LSTM(64, return_sequences=True),239        layers.LSTM(32),240        layers.Dense(FUT)241    ])242    model.compile(optimizer="adam", loss="mse")243    model.fit(X[:-1], np.zeros((len(X[:-1]), FUT)), epochs=EPOCHS_META, batch_size=32, verbose=0)  # lightweight placeholder244 245    pred_scaled = model.predict(X[-1:])[0]246    tmp = np.zeros((len(pred_scaled), X.shape[2]))247    tmp[:, 0] = pred_scaled248    pred_prices = MinMaxScaler().fit(df[features]).inverse_transform(tmp)[:, 0]249 250    if open_price is not None and len(pred_prices) > 0:251        pred_prices = open_price * (pred_prices / (pred_prices[0] if pred_prices[0] != 0 else 1.0))252 253    timeline = trading_minutes_for_date(predict_date, market="us")254    pred_prices = pred_prices[:len(timeline)]255    return pd.DataFrame({"datetime": timeline, "predicted_price": pred_prices, "model": "META_US"})256 257 258# ---------- DISPATCH ----------259def predict_for(market, code, predict_date, open_price=None):260    if market in ("a", "etf"):261        return [predict_meta_a(code, predict_date, open_price),262                predict_tft_a(code, predict_date, open_price)]263    elif market == "us":264        return [predict_meta_us(code, predict_date, open_price)]265    else:266        raise ValueError("Unknown market")267