start experiment 57 (exp/57-signal-quality-gate-gate-trades-based-on)

This commit is contained in:
zhaoli
2026-08-20 23:06:35 +00:00
parent fd5382caa4
commit ceb1e196e2
27 changed files with 1948 additions and 0 deletions
@@ -0,0 +1,11 @@
from .ic_gate import ICGateTopkDropoutStrategy # noqa: F401
from .optimal_stop import OptimalStopControl # noqa: F401
from .regime_gate import RegimeGateTopkDropoutStrategy # noqa: F401
from .weekly_rebalance import WeeklyRebalanceDropoutStrategy # noqa: F401
__all__ = [
"ICGateTopkDropoutStrategy",
"OptimalStopControl",
"RegimeGateTopkDropoutStrategy",
"WeeklyRebalanceDropoutStrategy",
]
@@ -0,0 +1,117 @@
"""Realized-IC circuit breaker TopkDropout strategy.
Subclass of ``qlib.contrib.strategy.signal_strategy.TopkDropoutStrategy`` that
holds the book (issues NO orders) while the streaming realized RankIC of the
deployed signal is below threshold — i.e. the model's cross-sectional
predictions are no longer earning against realized forward returns. When the
gate is open it behaves exactly like the reference TopkDropoutStrategy.
The gate is evaluated per trade step on the trailing mean realized RankIC of
the signal over the last ``ic_window`` trading days whose label is fully
realized as of the decision date (no lookahead — a 5d fwd label ``close[t+6]/
close[t+1]-1`` is only known at ``t+6``).
Two wiring modes:
* ``ic_gate``: a precomputed ``pd.Series`` indexed by datetime of booleans
(True = gate open / trade allowed). Computed once by the caller (e.g.
``rd_backtest``) and looked up per step. Missing dates default to open.
* realized-IC self-computation: when ``ic_min_rankic`` is given but no
``ic_gate``, the strategy computes the per-date realized RankIC itself from
``self.signal`` (the pred scores) and the lake 1d bars via
``tac_qlib.risk_limits.realized_rankic_series``, then applies the same
trailing-window comparison. Works when instantiated from a workflow YAML
PortAnaRecord config (``lake_root`` / ``market`` must be provided).
"""
from __future__ import annotations
import pandas as pd
from qlib.backtest.decision import TradeDecisionWO
from qlib.contrib.strategy.signal_strategy import TopkDropoutStrategy
from tac_qlib.risk_limits import ic_circuit_breaker, realized_rankic_series
__all__ = ["ICGateTopkDropoutStrategy"]
class ICGateTopkDropoutStrategy(TopkDropoutStrategy):
"""TopkDropout with a streaming realized-IC circuit breaker.
Parameters
----------
topk, n_drop, method_sell, method_buy, hold_thresh, only_tradable,
forbid_all_trade_at_limit : same as ``TopkDropoutStrategy``.
ic_min_rankic : float — pause new trading while trailing realized RankIC is
below this threshold (0 disables the gate).
ic_window : int — trailing window for the realized RankIC mean (default 22).
ic_label_horizon : int — label horizon in trading days (default 6).
ic_min_obs : int — min realized labels before the gate arms (default 10).
ic_gate : pd.Series, optional — precomputed per-date gate (bool indexed by
datetime). When provided, it overrides self-computation.
lake_root, market : str — lake location for self-computed realized IC.
"""
def __init__(
self,
*,
topk,
n_drop,
ic_min_rankic: float = 0.0,
ic_window: int = 22,
ic_label_horizon: int = 6,
ic_min_obs: int = 10,
ic_gate=None,
lake_root: str = "",
market: str = "US",
**kwargs,
):
super().__init__(topk=topk, n_drop=n_drop, **kwargs)
self.ic_min_rankic = float(ic_min_rankic or 0.0)
self.ic_window = int(ic_window or 22)
self.ic_label_horizon = int(ic_label_horizon or 6)
self.ic_min_obs = int(ic_min_obs or 10)
self._ic_gate = ic_gate
self._realized_ic = None
self.lake_root = lake_root or ""
self.market = market or "US"
def _load_realized_ic(self):
if self._realized_ic is None:
pred_start_time, pred_end_time = self.trade_calendar.get_step_time(
self.trade_calendar.get_trade_step(), shift=-self.ic_label_horizon
)
pred = self.signal.get_signal(start_time=pred_start_time, end_time=pred_end_time)
if isinstance(pred, pd.DataFrame):
pred = pred.iloc[:, 0]
self._realized_ic = realized_rankic_series(
pred, self.lake_root, self.market, label_horizon=self.ic_label_horizon
)
return self._realized_ic
def _gate_open(self, trade_start_time) -> bool:
ts = pd.Timestamp(trade_start_time)
if self._ic_gate is not None:
# precomputed gate series: look up the latest known decision date <= ts
known = self._ic_gate[self._ic_gate.index <= ts]
if len(known):
return bool(known.iloc[-1])
return True
if self.ic_min_rankic <= 0:
return True
realized = self._load_realized_ic()
limits = {
"ic_min_rankic": self.ic_min_rankic,
"ic_window": self.ic_window,
"ic_min_obs": self.ic_min_obs,
}
tripped, _reason, _trail = ic_circuit_breaker(realized, ts, limits)
return not tripped
def generate_trade_decision(self, execute_result=None):
trade_step = self.trade_calendar.get_trade_step()
trade_start_time, _ = self.trade_calendar.get_step_time(trade_step)
if not self._gate_open(trade_start_time):
return TradeDecisionWO([], self)
return super().generate_trade_decision(execute_result)
@@ -0,0 +1,217 @@
"""Optimal-stopping / stochastic-control strategy for cross-sectional signals.
Entry is a control policy: a symbol opens a position only when its cross-sectional
signal percentile is at or above ``entry_pct`` (i.e. it is one of the top-ranked
names) and the portfolio has fewer than ``topk`` open positions.
Exit is an optimal-stopping rule: a held position is stopped (closed) when its
signal percentile falls below ``exit_pct`` (the continuation value of holding is
no longer worth the risk), OR after ``max_hold_days`` (time stop / finite
horizon), OR when the position P&L breaches ``sl`` (loss control) and the
position has been held at least ``min_hold_days``.
Sizing is fixed ``notional`` per position (equal-weight control), unlike the
TopkDropout cash-allocation heuristic.
Wired into qrun workflows like any ``BaseStrategy`` (see ``PortAnaRecord``
config). Mirrors the API usage of qlib's ``TopkDropoutStrategy``: ``Order``/
``OrderDir`` from ``qlib.backtest.decision``, ``trade_calendar`` /
``trade_exchange`` / ``trade_position`` injected by the backtest executor.
"""
from __future__ import annotations
from typing import List
import pandas as pd
from qlib.backtest import Order
from qlib.backtest.decision import OrderDir, TradeDecisionWO
from qlib.contrib.strategy.signal_strategy import BaseSignalStrategy
__all__ = ["OptimalStopControl"]
DEFAULT_NOTIONAL = 20_000.0
DEFAULT_ENTRY_PCT = 0.80
DEFAULT_EXIT_PCT = 0.50
DEFAULT_MAX_HOLD_DAYS = 10
DEFAULT_MIN_HOLD_DAYS = 2
DEFAULT_SL = -0.06
class OptimalStopControl(BaseSignalStrategy):
"""Optimal-stopping long-only strategy over a cross-sectional signal.
Parameters
----------
topk : max number of concurrent positions.
entry_pct : min cross-sectional score percentile required to OPEN (0..1).
exit_pct : held positions are stopped when score percentile < exit_pct.
max_hold_days : hard time stop (finite-horizon close).
min_hold_days : minimum holding days before stop-loss is evaluated.
notional : $ per position (equal-weight control).
sl : stop-loss threshold as fraction of entry price (<= 0), disabled if 0.
"""
def __init__(
self,
*,
signal=None,
topk: int = 10,
entry_pct: float = DEFAULT_ENTRY_PCT,
exit_pct: float = DEFAULT_EXIT_PCT,
max_hold_days: int = DEFAULT_MAX_HOLD_DAYS,
min_hold_days: int = DEFAULT_MIN_HOLD_DAYS,
notional: float = DEFAULT_NOTIONAL,
sl: float = DEFAULT_SL,
risk_degree: float = 0.95,
trade_exchange=None,
level_infra=None,
common_infra=None,
**kwargs,
):
super().__init__(
signal=signal,
trade_exchange=trade_exchange,
level_infra=level_infra,
common_infra=common_infra,
**kwargs,
)
self.topk = topk
self.entry_pct = entry_pct
self.exit_pct = exit_pct
self.max_hold_days = max_hold_days
self.min_hold_days = min_hold_days
self.notional = notional
self.sl = sl
# ------------------------------------------------------------------ utils
@staticmethod
def _pct_rank(score: pd.Series) -> pd.Series:
return score.rank(pct=True)
def _entry_price(self, pos) -> float:
# Position stores avg entry price under key "price" (see Position.position)
price = pos.position.get("price")
if price is None:
price = pos.get_stock_amount("price")
return float(price)
def _pnl_pct(self, pos, mark: float) -> float:
entry = self._entry_price(pos)
if not entry or entry != entry:
return 0.0
return mark / entry - 1.0
def _is_tradable(self, code, start, end, direction) -> bool:
try:
return self.trade_exchange.is_stock_tradable(
stock_id=code, start_time=start, end_time=end, direction=direction
)
except TypeError: # some exchanges take no direction kwarg
return self.trade_exchange.is_stock_tradable(stock_id=code, start_time=start, end_time=end)
# ------------------------------------------------------------ decision
def generate_trade_decision(self, execute_result=None):
trade_step = self.trade_calendar.get_trade_step()
trade_start, trade_end = self.trade_calendar.get_step_time(trade_step)
pred_start, pred_end = self.trade_calendar.get_step_time(trade_step, shift=1)
pred_score = self.signal.get_signal(start_time=pred_start, end_time=pred_end)
if isinstance(pred_score, pd.DataFrame):
pred_score = pred_score.iloc[:, 0]
if pred_score is None or len(pred_score) == 0:
return TradeDecisionWO([], self)
pct = self._pct_rank(pred_score)
time_per_step = self.trade_calendar.get_freq()
current_temp = __import__("copy").deepcopy(self.trade_position)
holdings = {}
for code in current_temp.get_stock_list():
if abs(current_temp.get_stock_amount(code)) > 1e-6:
holdings[code] = current_temp
# ---- optimal stopping: close held positions -----------------------
sell_orders: List[Order] = []
closed_today = set()
kept = {}
for code, pos in holdings.items():
held = current_temp.get_stock_count(code, bar=time_per_step)
mark = self.trade_exchange.get_deal_price(
stock_id=code, start_time=trade_start, end_time=trade_end, direction=Order.SELL
)
if mark is None or mark != mark:
continue
rank = pct.get(code, 0.0)
stop_pnl = held >= self.min_hold_days and self.sl < 0 and self._pnl_pct(pos, mark) <= self.sl
if held >= self.max_hold_days or rank < self.exit_pct or stop_pnl:
amt = abs(current_temp.get_stock_amount(code))
o = Order(stock_id=code, amount=amt, start_time=trade_start,
end_time=trade_end, direction=Order.SELL)
if self.trade_exchange.check_order(o):
sell_orders.append(o)
self.trade_exchange.deal_order(o, position=current_temp)
closed_today.add(code)
else:
kept[code] = mark
# ---- equal-weight control: target notional per name -----------------
# candidate opens: top-ranked names whose signal pct >= entry_pct
rank_desc = pred_score.sort_values(ascending=False)
held_codes = set(kept)
opens = []
for sym in rank_desc.index:
if len(opens) >= self.topk:
break
if sym in held_codes:
continue
if pct.get(sym, 0.0) < self.entry_pct:
continue
if not self._is_tradable(sym, trade_start, trade_end, OrderDir.BUY):
continue
opens.append(sym)
targets = held_codes | set(opens)
if not targets:
return TradeDecisionWO(sell_orders, self)
# total value (cash + marked positions) -> per-target notional
total_value = current_temp.get_cash()
for code, mark in kept.items():
total_value += abs(current_temp.get_stock_amount(code)) * mark
target_notional = total_value * self.risk_degree / max(1, len(targets))
# ---- rebalance kept positions toward target weight ------------------
buy_orders: List[Order] = []
for code, mark in kept.items():
cur = abs(current_temp.get_stock_amount(code)) * mark
diff_notional = target_notional - cur
if abs(diff_notional) / target_notional < 0.02:
continue # skip tiny rebalances
amount_delta = diff_notional / mark
direction = Order.BUY if amount_delta > 0 else Order.SELL
o = Order(stock_id=code, amount=abs(amount_delta), start_time=trade_start,
end_time=trade_end, direction=direction)
if self.trade_exchange.check_order(o):
(buy_orders if direction == Order.BUY else sell_orders).append(o)
self.trade_exchange.deal_order(o, position=current_temp)
# ---- open new positions at target weight ----------------------------
for sym in opens:
px = self.trade_exchange.get_deal_price(
stock_id=sym, start_time=trade_start, end_time=trade_end, direction=OrderDir.BUY
)
if px is None or px != px or px <= 0:
continue
amount = target_notional / px
factor = self.trade_exchange.get_factor(
stock_id=sym, start_time=trade_start, end_time=trade_end
)
amount = self.trade_exchange.round_amount_by_trade_unit(amount, factor)
o = Order(stock_id=sym, amount=amount, start_time=trade_start,
end_time=trade_end, direction=Order.BUY)
if self.trade_exchange.check_order(o):
buy_orders.append(o)
return TradeDecisionWO(sell_orders + buy_orders, self)
@@ -0,0 +1,215 @@
"""Regime-gate TopkDropout strategy.
Subclass of ``qlib.contrib.strategy.signal_strategy.TopkDropoutStrategy`` that
holds the book (issues NO orders) while a regime detector says the market is in
an unfavorable state. When the gate is open it behaves exactly like the
reference TopkDropoutStrategy.
Three detector types are supported (all causal — no lookahead):
* ``dispersion``: cross-sectional standard deviation of 22-day rolling returns
across the universe. Gate closes when CS dispersion < threshold (low
dispersion means the spread between winners and losers is too narrow for
TopkDropout to exploit).
* ``vol``: cross-sectional mean of 22-day rolling realized volatility. Gate
closes when avg vol is outside a band ``[vol_low, vol_high]`` (strategy
needs moderate vol — too calm or too turbulent both hurt).
* ``hmm``: pre-computed HMM posterior for regime 1 (``sp_hmm_p_regime1``).
Gate closes when posterior < threshold (model is not confident the calm
regime is active).
The gate is provided as a precomputed ``pd.Series`` of booleans indexed by
datetime (True = trade allowed). The companion ``compute_regime_gate``
function builds this series from lake bars; call it once before backtesting
and pass the result as the ``regime_gate`` parameter.
"""
from __future__ import annotations
import pandas as pd
from qlib.backtest.decision import TradeDecisionWO
from qlib.contrib.strategy.signal_strategy import TopkDropoutStrategy
__all__ = ["RegimeGateTopkDropoutStrategy", "compute_regime_gate"]
class RegimeGateTopkDropoutStrategy(TopkDropoutStrategy):
"""TopkDropout with a regime-gate circuit breaker.
Parameters
----------
topk, n_drop, method_sell, method_buy, hold_thresh, only_tradable,
forbid_all_trade_at_limit : same as ``TopkDropoutStrategy``.
regime_gate : pd.Series — precomputed per-date gate (bool indexed by
datetime). True = trade allowed, False = no orders. Missing dates
default to open (trade allowed).
"""
def __init__(self, *, regime_gate=None, **kwargs):
super().__init__(**kwargs)
self._regime_gate = regime_gate
def _gate_open(self, trade_start_time) -> bool:
if self._regime_gate is None:
return True
ts = pd.Timestamp(trade_start_time)
known = self._regime_gate[self._regime_gate.index <= ts]
if len(known):
return bool(known.iloc[-1])
return True # default open if no history yet
def generate_trade_decision(self, execute_result=None):
trade_step = self.trade_calendar.get_trade_step()
trade_start_time, _ = self.trade_calendar.get_step_time(trade_step)
if not self._gate_open(trade_start_time):
return TradeDecisionWO([], self)
return super().generate_trade_decision(execute_result)
# ---------------------------------------------------------------------------
# Precomputation helper
# ---------------------------------------------------------------------------
def compute_regime_gate(
detector: str,
threshold: float = 0.0,
*,
lake_root: str = "",
market: str = "US",
start: str = "2015-01-03",
end: str = "2026-08-19",
vol_low: float = 0.0,
vol_high: float = 999.0,
hmm_field: str = "sp_hmm_p_regime1",
) -> pd.Series:
"""Build a per-date regime gate series from lake bars.
Parameters
----------
detector : str — ``"dispersion"``, ``"vol"``, or ``"hmm"``.
threshold : float — for ``dispersion``: min CS dispersion to allow trading.
For ``hmm``: min HMM posterior to allow trading.
Ignored for ``vol`` (uses ``vol_low``/``vol_high`` band instead).
lake_root, market : str — lake location.
start, end : str — date window.
vol_low, vol_high : float — annualized vol band for the ``vol`` detector.
hmm_field : str — HMM feature column name for the ``hmm`` detector.
Returns
-------
pd.Series — bool, indexed by datetime. True = trade allowed.
"""
from tac_qlib.data.config import LakeConfig, resolve_lake_root
cfg = LakeConfig(resolve_lake_root(lake_root or None), market)
symbols = _universe_symbols(cfg)
close_df, vol_df = _load_daily_bars(symbols, cfg, start, end)
if close_df.empty:
return pd.Series(dtype=bool)
if detector == "dispersion":
return _dispersion_gate(close_df, threshold)
elif detector == "vol":
return _vol_gate(close_df, vol_low, vol_high)
elif detector == "hmm":
return _hmm_gate(cfg, symbols, threshold, start, end, hmm_field)
else:
raise ValueError(f"Unknown detector: {detector!r}")
def _universe_symbols(cfg) -> list:
"""Read symbols from the lake symbols.parquet."""
import pathlib
sp = cfg.lake_root / "symbols.parquet"
if sp.exists():
df = pd.read_parquet(sp)
col = "symbol" if "symbol" in df.columns else df.columns[0]
return sorted(df[col].astype(str).str.upper().tolist())
return []
def _load_daily_bars(symbols, cfg, start, end):
"""Load daily close prices for all symbols into a wide DataFrame."""
closes = {}
vols = {}
for sym in symbols:
p = cfg.bar_path("1d", sym)
if not p.exists():
continue
try:
df = pd.read_parquet(p)
except Exception:
continue
if not len(df):
continue
tcol = df["t"] if "t" in df.columns else df["date"]
ts = pd.to_datetime(tcol)
df = df.assign(_t=ts).set_index("_t").sort_index()
df = df.loc[start:end]
if len(df) < 22:
continue
closes[sym] = df["c"]
if "v" in df.columns:
vols[sym] = df["v"]
close_df = pd.DataFrame(closes)
vol_df = pd.DataFrame(vols) if vols else None
return close_df, vol_df
def _dispersion_gate(close_df, threshold):
"""Cross-sectional dispersion of 22-day rolling returns."""
if close_df.empty or close_df.shape[1] < 2:
return pd.Series(dtype=bool)
ret = close_df.pct_change(22)
cs_disp = ret.std(axis=1)
gate = cs_disp >= threshold
gate.iloc[:22] = True # warmup: allow trading
return gate
def _vol_gate(close_df, vol_low, vol_high):
"""Cross-sectional mean of 22-day rolling realized vol."""
if close_df.empty or close_df.shape[1] < 2:
return pd.Series(dtype=bool)
import numpy as np
log_ret = np.log(close_df / close_df.shift(1))
rv22 = log_ret.rolling(22).std() * (252 ** 0.5)
cs_mean_vol = rv22.mean(axis=1)
gate = (cs_mean_vol >= vol_low) & (cs_mean_vol <= vol_high)
gate.iloc[:22] = True # warmup
return gate
def _hmm_gate(cfg, symbols, threshold, start, end, hmm_field):
"""HMM regime posterior gate from persisted SP features."""
feat_root = cfg.lake_root / "features"
all_posteriors = {}
for sym in symbols:
# check both ta and sp family paths
for family in ("sp", "ta"):
p = feat_root / f"market=US" / f"timeframe=1d" / f"family={family}" / f"symbol={sym}.parquet"
if not p.exists():
continue
try:
df = pd.read_parquet(p)
except Exception:
continue
if hmm_field not in df.columns:
continue
tcol = df["t"] if "t" in df.columns else df["date"]
ts = pd.to_datetime(tcol)
s = pd.Series(df[hmm_field].values, index=ts, name=sym)
s = s.loc[start:end].dropna()
if len(s) > 0:
all_posteriors[sym] = s
break
if not all_posteriors:
# no HMM features found — default open
idx = pd.date_range(start, end, freq="B")
return pd.Series(True, index=idx)
post_df = pd.DataFrame(all_posteriors)
cs_mean = post_df.mean(axis=1)
gate = cs_mean >= threshold
return gate
@@ -0,0 +1,202 @@
"""Weekly-rebalance TopkDropout strategy.
Turnover-reduction variant of ``qlib.contrib.strategy.signal_strategy.TopkDropoutStrategy``:
the topk/n_drop selection and sizing are identical to the reference, but the
target book is recomputed only on the first trading day of each ISO week; on the
other days the strategy issues NO orders (holds the book untouched).
The weekly cadence is derived from the qlib trade calendar: a rebalance happens
when the current trade step's date belongs to a different ISO ``(year, week)``
than the previous trade step. ``hold_band_pct`` (default 0) optionally skips
tiny rebalances: when a name's existing position differs from the new target by
less than this fraction, no order is generated for it.
"""
from __future__ import annotations
from typing import List
import numpy as np
import pandas as pd
from qlib.backtest import Order
from qlib.backtest.decision import OrderDir, TradeDecisionWO
from qlib.contrib.strategy.signal_strategy import TopkDropoutStrategy
__all__ = ["WeeklyRebalanceDropoutStrategy"]
DEFAULT_HOLD_BAND_PCT = 0.0
class WeeklyRebalanceDropoutStrategy(TopkDropoutStrategy):
"""TopkDropout rebalanced once per ISO week; holds otherwise.
Parameters
----------
topk, n_drop, method_sell, method_buy, hold_thresh, only_tradable,
forbid_all_trade_at_limit : same as ``TopkDropoutStrategy``.
hold_band_pct : skip order for a name whose deviation from target weight is
below this fraction of the target (no-trade buffer band).
"""
def __init__(self, *, topk, n_drop, hold_band_pct: float = DEFAULT_HOLD_BAND_PCT, **kwargs):
super().__init__(topk=topk, n_drop=n_drop, **kwargs)
self.hold_band_pct = hold_band_pct
@staticmethod
def _iso_week(ts) -> tuple:
return (ts.year, ts.week)
def generate_trade_decision(self, execute_result=None):
import copy
trade_step = self.trade_calendar.get_trade_step()
trade_start_time, trade_end_time = self.trade_calendar.get_step_time(trade_step)
cur_week = self._iso_week(trade_start_time)
prev_week = getattr(self, "_last_week", None)
self._last_week = cur_week
if prev_week is not None and prev_week == cur_week:
# not the first trading day of this ISO week -> hold
return TradeDecisionWO([], self)
pred_start_time, pred_end_time = self.trade_calendar.get_step_time(trade_step, shift=1)
pred_score = self.signal.get_signal(start_time=pred_start_time, end_time=pred_end_time)
if isinstance(pred_score, pd.DataFrame):
pred_score = pred_score.iloc[:, 0]
if pred_score is None:
return TradeDecisionWO([], self)
if self.only_tradable:
def get_first_n(li, n, reverse=False):
cur_n = 0
res = []
for si in reversed(li) if reverse else li:
if self.trade_exchange.is_stock_tradable(
stock_id=si, start_time=trade_start_time, end_time=trade_end_time
):
res.append(si)
cur_n += 1
if cur_n >= n:
break
return res[::-1] if reverse else res
def get_last_n(li, n):
return get_first_n(li, n, reverse=True)
def filter_stock(li):
return [
si
for si in li
if self.trade_exchange.is_stock_tradable(
stock_id=si, start_time=trade_start_time, end_time=trade_end_time
)
]
else:
def get_first_n(li, n):
return list(li)[:n]
def get_last_n(li, n):
return list(li)[-n:]
def filter_stock(li):
return li
current_temp: "object" = copy.deepcopy(self.trade_position)
sell_order_list: List[Order] = []
buy_order_list: List[Order] = []
cash = current_temp.get_cash()
current_stock_list = current_temp.get_stock_list()
last = pred_score.reindex(current_stock_list).sort_values(ascending=False).index
if self.method_buy == "top":
today = get_first_n(
pred_score[~pred_score.index.isin(last)].sort_values(ascending=False).index,
self.n_drop + self.topk - len(last),
)
elif self.method_buy == "random":
topk_candi = get_first_n(pred_score.sort_values(ascending=False).index, self.topk)
candi = list(filter(lambda x: x not in last, topk_candi))
n = self.n_drop + self.topk - len(last)
try:
today = np.random.choice(candi, n, replace=False)
except ValueError:
today = candi
else:
raise NotImplementedError(f"This type of input is not supported")
comb = pred_score.reindex(last.union(pd.Index(today))).sort_values(ascending=False).index
if self.method_sell == "bottom":
sell = last[last.isin(get_last_n(comb, self.n_drop))]
elif self.method_sell == "random":
candi = filter_stock(last)
try:
sell = pd.Index(np.random.choice(candi, self.n_drop, replace=False) if len(last) else [])
except ValueError:
sell = candi
else:
raise NotImplementedError(f"This type of input is not supported")
buy = today[: len(sell) + self.topk - len(last)]
for code in current_stock_list:
if not self.trade_exchange.is_stock_tradable(
stock_id=code,
start_time=trade_start_time,
end_time=trade_end_time,
direction=None if self.forbid_all_trade_at_limit else OrderDir.SELL,
):
continue
if code in sell:
time_per_step = self.trade_calendar.get_freq()
if current_temp.get_stock_count(code, bar=time_per_step) < self.hold_thresh:
continue
sell_amount = current_temp.get_stock_amount(code=code)
sell_order = Order(
stock_id=code,
amount=sell_amount,
start_time=trade_start_time,
end_time=trade_end_time,
direction=Order.SELL,
)
if self.trade_exchange.check_order(sell_order):
sell_order_list.append(sell_order)
trade_val, trade_cost, trade_price = self.trade_exchange.deal_order(
sell_order, position=current_temp
)
cash += trade_val - trade_cost
if len(buy) == 0:
return TradeDecisionWO(sell_order_list, self)
value = cash * self.risk_degree / len(buy)
for code in buy:
if not self.trade_exchange.is_stock_tradable(
stock_id=code,
start_time=trade_start_time,
end_time=trade_end_time,
direction=None if self.forbid_all_trade_at_limit else OrderDir.BUY,
):
continue
buy_price = self.trade_exchange.get_deal_price(
stock_id=code, start_time=trade_start_time, end_time=trade_end_time, direction=OrderDir.BUY
)
buy_amount = value / buy_price
factor = self.trade_exchange.get_factor(
stock_id=code, start_time=trade_start_time, end_time=trade_end_time
)
buy_amount = self.trade_exchange.round_amount_by_trade_unit(buy_amount, factor)
buy_order = Order(
stock_id=code,
amount=buy_amount,
start_time=trade_start_time,
end_time=trade_end_time,
direction=Order.BUY,
)
buy_order_list.append(buy_order)
return TradeDecisionWO(sell_order_list + buy_order_list, self)