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 . import data # noqa: F401 (registers tac_qlib.contrib.data)
from . import model, strategy # noqa: F401
from .data import TACHandler # noqa: F401
from .model import RankICLGBModel # noqa: F401
from .strategy import OptimalStopControl # noqa: F401
__all__ = [
"TACHandler",
"RankICLGBModel",
"OptimalStopControl",
]
@@ -0,0 +1,3 @@
from .handler import TACHandler
__all__ = ["TACHandler"]
@@ -0,0 +1,254 @@
"""TACHandler: a qlib DataHandlerLP that builds datasets from the TradeAC lake.
This is the "custom DataHandler" entry point (Option B): the handler is referenced from the
workflow yaml's ``dataset.handler`` and reads OHLCV + pre-computed ta-lib features straight
from the lake parquet files through ``QLibDataLoader`` + the tac_qlib feature provider.
The standard qlib processor pipeline (``infer_processors`` / ``learn_processors``) still runs
on top, so existing recipes such as ``DropnaLabel``, ``CSZScoreNorm`` or ``RobustZScoreNorm``
keep working unchanged.
"""
from __future__ import annotations
import os
from inspect import getfullargspec
from typing import List, Optional, Tuple, Union
from qlib.data.dataset import processor as processor_module
from qlib.data.dataset.handler import DataHandlerLP
from qlib.utils import get_callable_kwargs
from ...data.config import (
LakeConfig,
timeframe_for_freq,
NON_FEATURE_COLUMNS,
)
DEFAULT_INFER_PROCESSORS = [
{"class": "DropAllNaN", "kwargs": {}},
{"class": "ProcessInf", "kwargs": {}},
{"class": "ZScoreNorm", "kwargs": {}},
{"class": "Fillna", "kwargs": {}},
]
DEFAULT_LEARN_PROCESSORS = [
{"class": "DropnaLabel"},
{"class": "CSZScoreNorm", "kwargs": {"fields_group": "label"}},
]
#: always include raw OHLCV; ta-lib columns are discovered from the lake and appended.
RAW_FEATURE_FIELDS = ("$open", "$high", "$low", "$close", "$vwap", "$volume")
DEFAULT_LABEL = "Ref($close,-2)/Ref($close,-1)-1"
def check_transform_proc(proc_l, fit_start_time, fit_end_time):
"""Port of ``qlib.contrib.data.handler.check_transform_proc`` (inject fit window into procs)."""
new_l = []
for p in proc_l:
if not isinstance(p, processor_module.Processor):
klass, pkwargs = get_callable_kwargs(p, processor_module)
args = getfullargspec(klass).args
if "fit_start_time" in args and "fit_end_time" in args:
assert fit_start_time is not None and fit_end_time is not None, (
"Make sure `fit_start_time` and `fit_end_time` are not None."
)
pkwargs.update({"fit_start_time": fit_start_time, "fit_end_time": fit_end_time})
proc_config = {"class": klass.__name__, "kwargs": pkwargs}
if isinstance(p, dict) and "module_path" in p:
proc_config["module_path"] = p["module_path"]
new_l.append(proc_config)
else:
new_l.append(p)
return new_l
def get_common_feature_fields(lake_root=None, market="US", timeframe="1d") -> List[str]:
"""Discover feature columns present in *every* feature file of the lake.
Walks the `family=ta|sp` partition layout (plus any legacy flat files).
TA and SP columns are disjoint by construction, so the common set is
computed per family (columns shared by all symbol files of that family),
then the per-family results are unioned. Returns sorted field names
(without the ``$`` prefix). Empty if no features are persisted.
"""
cfg = LakeConfig(lake_root, market)
feat_dir = cfg.features_dir(timeframe)
if not feat_dir.exists():
return []
import pyarrow.parquet as pq
def _family_common(fam_dir: Path) -> set:
common = None
for p in sorted(fam_dir.glob("symbol=*.parquet")):
try:
cols = set(pq.read_schema(p).names) - set(NON_FEATURE_COLUMNS)
except Exception: # pragma: no cover - skip unreadable files
continue
common = cols if common is None else (common & cols)
if not common:
break
return common or set()
common: set = set()
# family tier: features/market=*/timeframe=*/family=*/symbol=*.parquet
for fam in ("ta", "sp"):
fam_dir = feat_dir / f"family={fam}"
if fam_dir.is_dir():
common |= _family_common(fam_dir)
# legacy flat: features/market=*/timeframe=*/symbol=*.parquet
if (feat_dir / "family=ta").exists() or (feat_dir / "family=sp").exists():
pass # family layout already covered
else:
common |= _family_common(feat_dir)
return sorted(common)
class DropAllNaN(processor_module.Processor):
"""Drop feature columns that are all-NaN over the fit window.
The lake can hold fully-empty indicator columns (e.g. a ta-lib output that was NaN
from the start). Such columns carry no learnable signal and make ``ZScoreNorm.fit``
warn on empty slices, so we drop them before any other processor runs. The drop set
is fixed on the fit window once (during ``fit``), then applied consistently to every
segment so train/valid/test keep identical feature columns.
"""
def __init__(self, fit_start_time=None, fit_end_time=None):
self.fit_start_time = fit_start_time
self.fit_end_time = fit_end_time
self.cols_to_drop = []
def fit(self, df=None):
if df is None or len(df) == 0:
return self
window = df
if self.fit_start_time is not None and self.fit_end_time is not None:
try:
from qlib.data.dataset.utils import fetch_df_by_index
window = fetch_df_by_index(
df, slice(self.fit_start_time, self.fit_end_time), level="datetime"
)
except Exception: # pragma: no cover - defensive
window = df
if len(window) == 0:
return self
self.cols_to_drop = [c for c in window.columns if window[c].isna().all()]
return self
def __call__(self, df):
if self.cols_to_drop:
return df.drop(columns=self.cols_to_drop, errors="ignore")
return df
class TACHandler(DataHandlerLP):
"""DataHandlerLP backed by the TradeAC parquet lake.
Parameters mirror ``Alpha158``: ``instruments``/``start_time``/``end_time``/``freq`` define
the queried window; ``feature_fields`` selects the features (default: raw OHLCV + all common
ta-lib columns found in the lake); ``label`` is a qlib expression for the target.
"""
def __init__(
self,
instruments="all",
start_time=None,
end_time=None,
freq="day",
infer_processors=DEFAULT_INFER_PROCESSORS,
learn_processors=DEFAULT_LEARN_PROCESSORS,
fit_start_time=None,
fit_end_time=None,
process_type=DataHandlerLP.PTYPE_A,
filter_pipe=None,
feature_fields=None,
label=DEFAULT_LABEL,
lake_root=None,
market="US",
**kwargs,
):
# default the processor fit window to the queried window (like Alpha158 without a split)
if fit_start_time is None:
fit_start_time = start_time
if fit_end_time is None:
fit_end_time = end_time
infer_processors = check_transform_proc(infer_processors, fit_start_time, fit_end_time)
learn_processors = check_transform_proc(learn_processors, fit_start_time, fit_end_time)
feature_fields = self._normalize_feature_fields(feature_fields, freq, lake_root, market)
if not feature_fields:
raise ValueError(
"no feature fields available for the lake; set `feature_fields` explicitly "
"(e.g. ['$close', '$rsi_14', '$sma_20'])"
)
label_expr, label_names = self._normalize_label(label)
data_loader = {
"class": "QlibDataLoader",
"kwargs": {
"config": {
"feature": (feature_fields, feature_fields),
"label": (label_expr, label_names),
},
"filter_pipe": filter_pipe,
"freq": freq,
},
}
super().__init__(
instruments=instruments,
start_time=start_time,
end_time=end_time,
data_loader=data_loader,
infer_processors=infer_processors,
learn_processors=learn_processors,
process_type=process_type,
**kwargs,
)
# ------------------------------------------------------------------ config
@staticmethod
def _normalize_feature_fields(feature_fields, freq, lake_root, market) -> List[str]:
if feature_fields is None:
common = get_common_feature_fields(lake_root, market, timeframe_for_freq(freq))
feature_fields = list(RAW_FEATURE_FIELDS) + ["$" + f for f in common if "$" + f not in RAW_FEATURE_FIELDS]
elif isinstance(feature_fields, str):
feature_fields = [f.strip() for f in feature_fields.split(",") if f.strip()]
fields = [f if f.startswith("$") else "$" + f for f in feature_fields]
# de-dup while preserving order
seen, out = set(), []
for f in fields:
if f not in seen:
seen.add(f)
out.append(f)
return out
@staticmethod
def _normalize_label(label) -> Tuple[List[str], List[str]]:
if isinstance(label, str):
return [label], ["LABEL0"]
if isinstance(label, (list, tuple)):
if len(label) == 2 and isinstance(label[0], str):
return [label[0]], list(label[1]) if isinstance(label[1], (list, tuple)) else [label[1]]
return list(label), ["LABEL%d" % i for i in range(len(label))]
raise TypeError(f"unsupported label config: {label!r}")
# ------------------------------------------------------------------ utils
def get_label_config(self):
return DEFAULT_LABEL
@staticmethod
def discover_feature_fields(lake_root=None, market="US", freq="day") -> List[str]:
return get_common_feature_fields(lake_root, market, timeframe_for_freq(freq))
__all__ = ["TACHandler", "DropAllNaN", "get_common_feature_fields"]
# Make `DropAllNaN` resolvable by bare name from processor configs (e.g. the default
# ``infer_processors`` and workflow yamls that reference it without a ``module_path``),
# mirroring how qlib registers its own processors in ``qlib.data.dataset.processor``.
processor_module.DropAllNaN = DropAllNaN
@@ -0,0 +1,4 @@
from .rank_ensemble import RankICEnsembleLGBModel # noqa: F401
from .rank_gbdt import RankICLGBModel, rankic_feval # noqa: F401
__all__ = ["RankICLGBModel", "rankic_feval", "RankICEnsembleLGBModel"]
@@ -0,0 +1,189 @@
"""Seed-ensembled LightGBM that early-stops on cross-sectional RankIC.
``RankICEnsembleLGBModel`` wraps ``RankICLGBModel`` (per-day RankIC feval +
``metric='None'`` + ``first_metric_only`` early stopping) over a seed ensemble:
one sub-model is trained per seed with identical hyper-parameters, and
predictions are averaged across seeds. This is the model class the
``tac-rd-rank-ensemble-isolated`` reference run wires into its workflow
(``module_path: tac_qlib.contrib.model.rank_ensemble``).
The ensemble inherits the RankIC early-stopping behaviour of the single-seed
model (valid RankIC drives the stopping iteration) while the seed averaging
stabilizes the prediction against any single seed's early-stopping path.
Training is parallelized: the seed sub-models train in a thread pool —
``lgb.train`` is C++ and releases the GIL, so concurrent seeds do not block on
the GIL (5 seeds ~40min/5 on this box). Measured on a 6-physical-core / 12 SMT
host: the seeds scale ~2x, not linearly — the runs are memory-bandwidth bound
and each Booster caps its threads at ``cores // workers`` so 5 concurrent
boosters don't oversubscribe; larger-core hosts scale better. The qlib data
pipeline is warmed once on the calling thread (fills the handler cache), and
each worker then prepares its **own** ``lgb.Dataset`` (independent handle, so
no concurrent ``construct()`` on a shared handle — LightGBM's ``Dataset`` is
not thread-safe to build). qlib's ``R`` recorder is also not thread-safe, so
the per-seed evaluation curves are logged on the calling thread after the pool
finishes.
Wired into a workflow yaml like:
model:
class: RankICEnsembleLGBModel
module_path: tac_qlib.contrib.model.rank_ensemble
kwargs:
loss: mse
learning_rate: 0.02
num_leaves: 31
n_estimators: 3000
num_boost_round: 3000
early_stopping_rounds: 200
min_data_in_leaf: 20
lambda_l2: 0.5
colsample_bytree: 0.8
subsample: 0.8
subsample_freq: 1
reg_alpha: 0.1
reg_lambda: 1.0
seeds: "42,7,2026,99,123"
parallel: 5
Any ``**kwargs`` other than ``seeds``/``parallel`` are forwarded unchanged to
every ``RankICLGBModel`` sub-model (same params, different ``seed``).
"""
from __future__ import annotations
import os
from concurrent.futures import ThreadPoolExecutor
from typing import List, Optional
import pandas as pd
from qlib.data.dataset import DatasetH
from qlib.data.dataset.handler import DataHandlerLP
from tac_qlib.contrib.model.rank_gbdt import RankICLGBModel
__all__ = ["RankICEnsembleLGBModel"]
class RankICEnsembleLGBModel(RankICLGBModel):
"""Seed ensemble of RankIC-early-stopping LightGBM models.
Parameters
----------
seeds : comma-separated integers, one sub-model per seed.
parallel : number of seeds to train concurrently. ``0`` (default) = auto
(all seeds, bounded by the available cores); ``1`` = sequential.
**kwargs : forwarded to every ``RankICLGBModel`` sub-model (model
hyper-parameters). ``seeds``/``parallel`` are consumed here and not
forwarded.
"""
def __init__(self, seeds: str = "42", parallel: int = 0, **kwargs):
self.seeds = [int(s.strip()) for s in str(seeds).split(",") if s.strip()]
if not self.seeds:
raise ValueError("seeds must contain at least one integer")
self.parallel = int(parallel)
# drop seed/parallel handling from the base kwargs, keep everything else
self._model_kwargs = dict(kwargs)
super().__init__(**self._model_kwargs)
self._models: List[RankICLGBModel] = []
# --------------------------------------------------------------- helpers
@staticmethod
def _cores() -> int:
try:
return max(1, len(os.sched_getaffinity(0)))
except AttributeError:
return max(1, os.cpu_count() or 1)
def _worker_count(self) -> int:
if self.parallel > 0:
return min(len(self.seeds), self.parallel)
return min(len(self.seeds), self._cores())
# ------------------------------------------------------------------ fit
def fit(
self,
dataset: DatasetH,
num_boost_round: Optional[int] = None,
early_stopping_rounds: Optional[int] = None,
verbose_eval: int = 20,
evals_result=None,
reweighter=None,
**kwargs,
):
"""Train one RankICLGBModel per seed and keep them for prediction.
The qlib data pipeline is warmed once on this thread (handler cache),
then each seed sub-model trains in a parallel worker thread on its own
``lgb.Dataset`` (LightGBM releases the GIL in ``lgb.train``). Evals
are logged on this thread after the pool (qlib's ``R`` is not
thread-safe).
"""
n_round = num_boost_round or self.num_boost_round
n_es = early_stopping_rounds or self.early_stopping_rounds
if len(self.seeds) == 1:
m = RankICLGBModel(seed=self.seeds[0], **self._model_kwargs)
m.fit(
dataset,
num_boost_round=n_round,
early_stopping_rounds=n_es,
verbose_eval=verbose_eval,
evals_result=evals_result,
reweighter=reweighter,
**kwargs,
)
self._models = [m]
return
# Warm the qlib handler cache once on this thread so the workers'
# concurrent prepare() calls only hit cached frames (no first-write race).
proto = RankICLGBModel(seed=self.seeds[0], **self._model_kwargs)
proto._prepare_data(dataset, reweighter)
workers = self._worker_count()
# Cap per-Booster threads so concurrent seeds don't oversubscribe
# (LightGBM's num_threads=0 uses ALL cores per Booster).
per_booster = max(1, self._cores() // workers)
def fit_seed(seed):
m = RankICLGBModel(seed=seed, **self._model_kwargs)
if workers > 1 and "num_threads" not in m.params:
m.params["num_threads"] = per_booster
ds_l = m._prepare_data(dataset, reweighter)
booster, evals, names = m._train_from_datasets(
ds_l,
num_boost_round=n_round,
early_stopping_rounds=n_es,
verbose_eval=verbose_eval,
**kwargs,
)
m.model = booster
return m, evals, names
with ThreadPoolExecutor(max_workers=workers) as ex:
results = list(ex.map(fit_seed, self.seeds))
self._models = [m for m, _, _ in results]
# Merge + log evals on the main thread (qlib's R is not thread-safe).
if evals_result is not None:
for m, evals, names in results:
for k in names:
for key, val in evals.get(k, {}).items():
evals_result.setdefault(f"{k}.seed{m.params['seed']}", {})[key] = val
for m, evals, names in results:
self._log_evals(evals, names, prefix=f"seed{m.params['seed']}.")
# -------------------------------------------------------------- predict
def predict(self, dataset: DatasetH, segment="test") -> pd.Series:
"""Average the per-seed predictions over the given segment."""
if not self._models:
raise ValueError("model is not fitted yet!")
preds = [m.predict(dataset, segment=segment) for m in self._models]
if len(preds) == 1:
return preds[0]
frame = pd.concat(preds, axis=1)
return frame.mean(axis=1)
@@ -0,0 +1,238 @@
"""LGBModel variant that early-stops on cross-sectional RankIC instead of l2.
Standard qlib ``LGBModel`` early-stops on the regression loss (mse). For
cross-sectional alpha signals the quantity we actually care about is the per-day
rank correlation (Rank IC), which mse early-stopping does not optimize for.
Experiments on the 50-ETF lake (SP-5d 55-feature panel) show that early-stopping
on a custom RankIC feval lifts RankIC 0.047 -> 0.075 vs. the mse-stopped model.
This class reuses ``LGBModel``'s data preparation but:
- tags each ``lgb.Dataset`` with per-day query ``group`` sizes so a ranking
metric can be computed per trading day;
- injects a custom ``feval`` (mean per-day Spearman of pred vs label) into
``lgb.train``; early stopping then selects the iteration that maximizes
RankIC on the valid set;
- forces ``metric='None'`` + ``first_metric_only=True`` so early-stopping
tracks RankIC only (not the regression loss).
Wired into a workflow yaml like:
model:
class: RankICLGBModel
module_path: tac_qlib.contrib.model.rank_gbdt
kwargs:
loss: mse
learning_rate: 0.03
num_leaves: 31
n_estimators: 500
...
The rank feval is used for early-stopping selection only; the objective stays
the configured loss (default mse). Set ``rank_eval=False`` to fall back to the
plain LGBModel behaviour (early-stop on the loss).
Generic: works for any cross-sectional panel whose qlib dataset index has a
``datetime`` level (each level value = one query group). The per-day groups are
derived automatically, so no universe-specific configuration is needed.
"""
from __future__ import annotations
from typing import List, Optional, Tuple
import numpy as np
import pandas as pd
import lightgbm as lgb
from qlib.data.dataset import DatasetH
from qlib.data.dataset.handler import DataHandlerLP
from qlib.contrib.model.gbdt import LGBModel
from qlib.workflow import R
__all__ = ["RankICLGBModel", "rankic_feval"]
def _group_averaged_rank(values: np.ndarray, gid: np.ndarray, offs: np.ndarray) -> np.ndarray:
"""Averaged (tie-corrected) rank of ``values`` within each group, vectorized.
``gid`` maps each row to its group id; ``offs`` holds the cumulative row
offsets so that group ``i`` occupies rows ``[offs[i], offs[i+1])``. Returns
the same result as ``pandas.Series.rank(method='average')`` applied per
group, but in one pass (``np.lexsort`` is the only non-linear step).
"""
n = len(values)
order = np.lexsort((values, gid))
ord_rank = np.empty(n, dtype=np.float64)
ord_rank[order] = np.arange(n, dtype=np.float64) - offs[gid[order]] + 1.0
sg = gid[order]
sv = values[order]
newblock = np.empty(n, dtype=bool)
newblock[0] = True
newblock[1:] = (sg[1:] != sg[:-1]) | (sv[1:] != sv[:-1])
blockid = np.cumsum(newblock) - 1
block_mean = np.bincount(blockid, weights=ord_rank[order]) / np.bincount(blockid)
out = np.empty(n)
out[order] = block_mean[blockid]
return out
def _per_day_spearman(preds: np.ndarray, labels: np.ndarray, group: np.ndarray) -> float:
"""Mean per-day Spearman rank correlation of preds vs labels.
``group`` holds the number of rows of each trading day (query group), in
order. Days with <3 valid rows or a constant pred/label are skipped.
Vectorized: per-day Spearman == Pearson of the per-day rank transforms,
and the Pearson moments (``sum``, ``sum`` of products/squares) aggregate
over each day with ``np.bincount``. Runs ~10x faster than the per-day
``pd.Series.rank()`` loop that preceded it — this feval is invoked on the
train and valid panels every boosting round, per seed.
"""
if group is None or len(group) == 0:
return 0.0
offs = np.concatenate([[0], np.cumsum(group.astype(int))])
gid = np.repeat(np.arange(len(group)), group.astype(int))
rp = _group_averaged_rank(preds, gid, offs)
rl = _group_averaged_rank(labels, gid, offs)
n_g = group.astype(float)
s_p = np.bincount(gid, weights=rp)
s_l = np.bincount(gid, weights=rl)
s_pl = np.bincount(gid, weights=rp * rl)
s_pp = np.bincount(gid, weights=rp * rp)
s_ll = np.bincount(gid, weights=rl * rl)
cov = n_g * s_pl - s_p * s_l
var_p = n_g * s_pp - s_p ** 2
var_l = n_g * s_ll - s_l ** 2
denom = np.sqrt(var_p * var_l)
valid = (n_g >= 3) & (denom > 0)
corr = np.where(valid, cov / np.where(denom == 0, 1, denom), 0.0)
return float(corr[valid].mean()) if valid.any() else 0.0
def rankic_feval(preds, dataset):
"""LightGBM feval: mean RankIC (higher is better in lgb convention)."""
labels = dataset.get_label()
group = dataset.get_group()
ric = _per_day_spearman(preds, labels, group)
return "rankic", ric, True # (name, value, higher_is_better)
class RankICLGBModel(LGBModel):
"""LGBModel that early-stops on per-day RankIC via a custom feval."""
def __init__(self, rank_eval: bool = True, **kwargs):
super().__init__(**kwargs)
self.rank_eval = rank_eval
def _prepare_data(self, dataset: DatasetH, reweighter=None) -> List[Tuple[lgb.Dataset, str]]:
ds_l = []
assert "train" in dataset.segments
for key in ["train", "valid"]:
if key in dataset.segments:
df = dataset.prepare(key, col_set=["feature", "label"], data_key=DataHandlerLP.DK_L)
if df.empty:
raise ValueError("Empty data from dataset, please check your dataset config.")
x, y = df["feature"], df["label"]
if y.values.ndim == 2 and y.values.shape[1] == 1:
y = np.squeeze(y.values)
else:
raise ValueError("LightGBM doesn't support multi-label training")
if reweighter is None:
w = None
elif hasattr(reweighter, "reweight"):
w = reweighter.reweight(df)
else:
raise ValueError("Unsupported reweighter type.")
# per-day query groups: each trading day is one group
if self.rank_eval and isinstance(df.index, pd.MultiIndex) and "datetime" in df.index.names:
group = df.groupby(level="datetime").size().to_numpy(dtype=np.int32)
else:
group = None
d = lgb.Dataset(x.values, label=y, weight=w, group=group, free_raw_data=False)
ds_l.append((d, key))
return ds_l
def _train_from_datasets(
self,
ds_l: List[Tuple[lgb.Dataset, str]],
num_boost_round: Optional[int] = None,
early_stopping_rounds: Optional[int] = None,
verbose_eval: int = 20,
evals_result=None,
**kwargs,
) -> Tuple[lgb.Booster, dict, List[str]]:
"""Train a Booster from already-prepared ``lgb.Dataset`` objects.
Pure training — no ``R.log_metrics`` — so it can be called from worker
threads (qlib's ``R`` recorder is not thread-safe; the caller decides
when/where to log). Returns ``(booster, evals_result, segment_names)``.
"""
if evals_result is None:
evals_result = {}
ds, names = list(zip(*ds_l))
callbacks = [
lgb.early_stopping(
self.early_stopping_rounds if early_stopping_rounds is None else early_stopping_rounds
),
lgb.log_evaluation(period=verbose_eval),
lgb.record_evaluation(evals_result),
]
if self.rank_eval:
# early-stopping must be driven ONLY by the RankIC feval, not l2.
# metric='None' suppresses the default l2 metric; first_metric_only
# makes early_stopping track the single remaining (rankic) metric.
self.params["metric"] = "None"
self.params["first_metric_only"] = True
feval = rankic_feval
else:
self.params.pop("metric", None)
self.params.pop("first_metric_only", None)
feval = None
booster = lgb.train(
self.params,
ds[0],
num_boost_round=self.num_boost_round if num_boost_round is None else num_boost_round,
valid_sets=ds,
valid_names=names,
feval=feval,
callbacks=callbacks,
**kwargs,
)
return booster, evals_result, list(names)
def _log_evals(self, evals_result, names: List[str], prefix: str = "") -> None:
"""Log recorded evaluation curves to qlib's active recorder."""
for k in names:
for key, val in evals_result.get(k, {}).items():
name = f"{prefix}{key}.{k}"
for epoch, m in enumerate(val):
R.log_metrics(**{name.replace("@", "_"): m}, step=epoch)
def fit(
self,
dataset: DatasetH,
num_boost_round: Optional[int] = None,
early_stopping_rounds: Optional[int] = None,
verbose_eval: int = 20,
evals_result=None,
reweighter=None,
**kwargs,
):
if evals_result is None:
evals_result = {}
ds_l = self._prepare_data(dataset, reweighter)
self.model, evals_result, names = self._train_from_datasets(
ds_l,
num_boost_round=num_boost_round,
early_stopping_rounds=early_stopping_rounds,
verbose_eval=verbose_eval,
evals_result=evals_result,
**kwargs,
)
self._log_evals(evals_result, names)
@@ -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)