feat(model): add StockMixer (AAAI 2024) to model zoo and benchmarks - #2363
Open
Sourish-07 wants to merge 3 commits into
Open
Sourish-07 wants to merge 3 commits into
Sourish-07 wants to merge 3 commits into
Conversation
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Description
Adds a Qlib-native implementation of StockMixer (AAAI 2024), an
MLP-based architecture for stock price forecasting, to the contrib model
zoo.
qlib/contrib/model/pytorch_stockmixer.pyfollows the same skeleton aspytorch_tcn.py(Modelinterface,fit/predict, early stopping,seed/GPU handling,
count_parameterslogging). The network implementsthe paper's three mixers: indicator mixing, multi-scale time mixing
(upper-triangular
TriUmixing of the raw series and a conv-down-sampledview), and market-aware stock mixing (
NoGraphMixer, no prior stockgraph required) - plus the paper's MSE + α·pairwise-ranking loss.
Key design decisions:
DatasetH(notTSDatasetH), like TCN. Since the stock-mixingblock needs the full daily cross-section, the wrapper batches
day by day (grouped by the
datetimeindex level); each day is fedas one
(n_stock, time_steps, d_feat)tensor.n_stock. A booleanmask is threaded through the model so padded rows are excluded from the
NoGraphMixerLayerNorm statistics and zeroed before/after its denselayers - the only layer in the network that mixes across stocks; every
other layer operates per-stock, where padding is harmless by
construction. Padded rows are also masked out of the loss and dropped
from predictions. A day with more real instruments than
n_stockraises a clear error.
d_feat=6, time_steps=60(paper-faithful). Alpha158:d_feat=157, time_steps=1; the multi-scale time-mixing branchdegrades to a single-step mapping in that case (documented in the
model's docstring).
New files:
qlib/contrib/model/pytorch_stockmixer.pyexamples/benchmarks/StockMixer/workflow_config_stockmixer_Alpha360.yamlexamples/benchmarks/StockMixer/workflow_config_stockmixer_Alpha158.yamlexamples/benchmarks/StockMixer/requirements.txttests/model/test_stockmixer.pyREADME / model-zoo documentation and the official 20-seed benchmark-table
numbers are intentionally left out of this PR and can be added in a
follow-up once more thorough runs and hyperparameter tuning are done.
Motivation and Context
StockMixer is a recent MLP-based architecture that avoids relying on a
pre-defined stock graph while still modeling cross-sectional relationships.
Adding it expands Qlib's Quant Model Zoo with an architecture distinct
from the existing RNN/GNN/Transformer baselines.
Paper: https://ojs.aaai.org/index.php/AAAI/article/view/28681
Official code: https://github.com/SJTU-DMTai/StockMixer
How Has This Been Tested?
pytest qlib/tests/test_all_pipeline.pyunder upper directory ofqlib.Automated tests (
tests/model/test_stockmixer.py, 5 tests, all passing):instantiation from both real workflow configs; forward/backward for both
time_steps=60andtime_steps=1; a correctness test proving paddedstock rows cannot influence real stocks' outputs through the masked
NoGraphMixer(torch.allcloseunder extreme injected noise in thepadding); the
n_stockoverflow error path; a fullfit/predictcycleon synthetic data.
Full regression:
pytest tests/test_all_pipeline.py- 3 passed, nofailures (3 pre-existing unrelated warnings).
Real-data runs on CSI300 (via
qrun, full unmodified configs,n_epochs: 200,early_stop: 20, no crashes, no NaN/Inf anywhere):IC 0.0264 ± 0.0028, ICIR 0.207 ± 0.044, Rank IC 0.0326 ± 0.0041,
Rank ICIR 0.246 ± 0.020, annualized return (with cost) 5.45% ± 1.17%,
information ratio (with cost) 0.686 ± 0.164. All 5 runs early-stopped
within epochs 20–22, best epoch between 0–2.
epoch 21 (best @ epoch 1), with-cost annualized return −8.8%,
information ratio −0.946.
Known limitation: both configs' validation score
peaks very early (epoch 0–2) and then degrades, triggering early stop
well before 200 epochs. Alpha158 ends up modestly positive; Alpha360's
backtest is currently negative after cost. This looks like a
hyperparameter-tuning issue (learning rate / rank-loss weight / possibly
batch composition) rather than a correctness bug - the architecture
itself is verified correct via the masking test above, and training
losses decrease smoothly and finitely throughout. I'd value input from
maintainers familiar with similar architectures on reasonable defaults,
and plan to tune further before/alongside the full 20-seed benchmark run.
Screenshots of Test Results (if appropriate):
tests/test_all_pipeline.py- 3 passed, 3 pre-existing warnings, exit 0.tests/model/test_stockmixer.py- 5 passed, exit 0.Types of changes