claraleeee/hftpredict
0
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 