mirror of
https://github.com/microsoft/qlib.git
synced 2026-07-21 19:27:36 +08:00
Compare commits
119 Commits
v0.8.2
...
mini_proje
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a02ac95538 | ||
|
|
cc94c32db6 | ||
|
|
9a40fd3cdc | ||
|
|
c4281121e3 | ||
|
|
2de9903200 | ||
|
|
2cf842bcfe | ||
|
|
9e381493c2 | ||
|
|
a73b60d05a | ||
|
|
64979ad769 | ||
|
|
c5cf8fb9cc | ||
|
|
5d579d1a20 | ||
|
|
3c9c76b384 | ||
|
|
9d0a8f61d1 | ||
|
|
701b18af1b | ||
|
|
84ff662a26 | ||
|
|
00e40e775b | ||
|
|
45fe5e6974 | ||
|
|
366a9c33f3 | ||
|
|
982e0da715 | ||
|
|
cd5e5d5235 | ||
|
|
caea495f40 | ||
|
|
d934c8caba | ||
|
|
a139986f4e | ||
|
|
12c3de42d0 | ||
|
|
fe0f9427f2 | ||
|
|
a973e4fb66 | ||
|
|
c60366addd | ||
|
|
41447f320b | ||
|
|
e1271a83f7 | ||
|
|
30b531086c | ||
|
|
87926513cb | ||
|
|
7bfc7e1797 | ||
|
|
85e7cdcac3 | ||
|
|
08fd1d3f42 | ||
|
|
defd6758f6 | ||
|
|
61cc1a3867 | ||
|
|
655ed982cf | ||
|
|
2952c443ca | ||
|
|
7f1293ec34 | ||
|
|
73438807f9 | ||
|
|
962751c72d | ||
|
|
56cfa480dc | ||
|
|
6edd0bf298 | ||
|
|
fe155703b0 | ||
|
|
3c4f4bfd44 | ||
|
|
5200ff520a | ||
|
|
30e457119c | ||
|
|
243e516cf1 | ||
|
|
e229b567ad | ||
|
|
f129bfef5d | ||
|
|
9dd5e07819 | ||
|
|
00ed35fc1b | ||
|
|
3f53a097b0 | ||
|
|
fb230a8097 | ||
|
|
ff4724e248 | ||
|
|
73d90f7f44 | ||
|
|
b7988e6428 | ||
|
|
8efc8b92ef | ||
|
|
f2a5ecd98a | ||
|
|
705354cc28 | ||
|
|
1b5d0d4d6d | ||
|
|
f4a481945b | ||
|
|
5f18ba7970 | ||
|
|
2681c61c60 | ||
|
|
776b0c5bb4 | ||
|
|
829ad9f5e9 | ||
|
|
921c13cc90 | ||
|
|
0f519f6053 | ||
|
|
2ed806c846 | ||
|
|
d2f0bebf60 | ||
|
|
615a381038 | ||
|
|
568a88fddb | ||
|
|
058f976727 | ||
|
|
faa99f30fa | ||
|
|
837067b9e1 | ||
|
|
3a911bc09b | ||
|
|
90be21bb40 | ||
|
|
7540b1257b | ||
|
|
57f7ed9914 | ||
|
|
9e3d0249f7 | ||
|
|
2ac964c470 | ||
|
|
07f0d4f599 | ||
|
|
ea4fb33ff2 | ||
|
|
ed0c238787 | ||
|
|
80af395b3c | ||
|
|
4dc66932d5 | ||
|
|
40dd84857c | ||
|
|
74cc21fc2c | ||
|
|
ec8969a3ae | ||
|
|
528f74af09 | ||
|
|
d482726f28 | ||
|
|
cfc3e886ed | ||
|
|
60d45ad770 | ||
|
|
0e8b94a552 | ||
|
|
4bf127eba5 | ||
|
|
c149c8616c | ||
|
|
3274e16c95 | ||
|
|
d496cf7476 | ||
|
|
357ee74b6f | ||
|
|
5da5cf5175 | ||
|
|
6a946761cf | ||
|
|
76b7b5f24b | ||
|
|
d7d19feb4e | ||
|
|
bba6972a55 | ||
|
|
18af288692 | ||
|
|
ba056850cb | ||
|
|
aed5b8ebc0 | ||
|
|
79355666a9 | ||
|
|
144e1e2459 | ||
|
|
635632e4ed | ||
|
|
c5834476e2 | ||
|
|
01afd06e18 | ||
|
|
d533219738 | ||
|
|
5b5c99fe75 | ||
|
|
da48f42f3f | ||
|
|
f979dcf5e8 | ||
|
|
97aa16a078 | ||
|
|
094be9be86 | ||
|
|
d9b9386032 |
2
.github/PULL_REQUEST_TEMPLATE.md
vendored
2
.github/PULL_REQUEST_TEMPLATE.md
vendored
@@ -8,7 +8,7 @@
|
|||||||
<!--- Why is this change required? What problem does it solve? -->
|
<!--- Why is this change required? What problem does it solve? -->
|
||||||
|
|
||||||
## How Has This Been Tested?
|
## How Has This Been Tested?
|
||||||
<! --- Put an `x` in all the boxes that apply: --->
|
<!--- Put an `x` in all the boxes that apply: --->
|
||||||
- [ ] Pass the test by running: `pytest qlib/tests/test_all_pipeline.py` under upper directory of `qlib`.
|
- [ ] Pass the test by running: `pytest qlib/tests/test_all_pipeline.py` under upper directory of `qlib`.
|
||||||
- [ ] If you are adding a new feature, test on your own test scripts.
|
- [ ] If you are adding a new feature, test on your own test scripts.
|
||||||
|
|
||||||
|
|||||||
3
.github/workflows/python-publish.yml
vendored
3
.github/workflows/python-publish.yml
vendored
@@ -12,7 +12,8 @@ jobs:
|
|||||||
runs-on: ${{ matrix.os }}
|
runs-on: ${{ matrix.os }}
|
||||||
strategy:
|
strategy:
|
||||||
matrix:
|
matrix:
|
||||||
os: [windows-latest, macos-latest, macos-11]
|
os: [windows-latest, macos-11]
|
||||||
|
# FIXME: macos-latest will raise error now.
|
||||||
# not supporting 3.6 due to annotations is not supported https://stackoverflow.com/a/52890129
|
# not supporting 3.6 due to annotations is not supported https://stackoverflow.com/a/52890129
|
||||||
python-version: [3.7, 3.8]
|
python-version: [3.7, 3.8]
|
||||||
|
|
||||||
|
|||||||
81
.github/workflows/test.yml
vendored
81
.github/workflows/test.yml
vendored
@@ -33,11 +33,85 @@ jobs:
|
|||||||
- name: Install Qlib with pip
|
- name: Install Qlib with pip
|
||||||
run: |
|
run: |
|
||||||
pip install numpy==1.19.5 ruamel.yaml
|
pip install numpy==1.19.5 ruamel.yaml
|
||||||
pip install pyqlib --ignore-installed
|
pip install pyqlib --ignore-installed
|
||||||
|
|
||||||
|
- name: Make html with sphinx
|
||||||
|
run: |
|
||||||
|
pip install -U sphinx
|
||||||
|
pip install sphinx_rtd_theme readthedocs_sphinx_ext
|
||||||
|
pip install --exists-action=w --no-cache-dir -r docs/requirements.txt
|
||||||
|
cd docs
|
||||||
|
sphinx-build -b html . build
|
||||||
|
cd ..
|
||||||
|
|
||||||
|
# Check Qlib with pylint
|
||||||
|
# TODO: These problems we will solve in the future. Important among them are: W0221, W0223, W0237, E1102
|
||||||
|
# C0103: invalid-name
|
||||||
|
# C0209: consider-using-f-string
|
||||||
|
# R0402: consider-using-from-import
|
||||||
|
# R1705: no-else-return
|
||||||
|
# R1710: inconsistent-return-statements
|
||||||
|
# R1725: super-with-arguments
|
||||||
|
# R1735: use-dict-literal
|
||||||
|
# W0102: dangerous-default-value
|
||||||
|
# W0212: protected-access
|
||||||
|
# W0221: arguments-differ
|
||||||
|
# W0223: abstract-method
|
||||||
|
# W0231: super-init-not-called
|
||||||
|
# W0237: arguments-renamed
|
||||||
|
# W0612: unused-variable
|
||||||
|
# W0621: redefined-outer-name
|
||||||
|
# W0622: redefined-builtin
|
||||||
|
# FIXME: specify exception type
|
||||||
|
# W0703: broad-except
|
||||||
|
# W1309: f-string-without-interpolation
|
||||||
|
# E1102: not-callable
|
||||||
|
# E1136: unsubscriptable-object
|
||||||
|
# References for parameters: https://github.com/PyCQA/pylint/issues/4577#issuecomment-1000245962
|
||||||
|
- name: Check Qlib with pylint
|
||||||
|
run: |
|
||||||
|
pip install --upgrade pip
|
||||||
|
pip install pylint
|
||||||
|
pylint --disable=C0104,C0114,C0115,C0116,C0301,C0302,C0411,C0413,C1802,R0201,R0401,R0801,R0902,R0903,R0911,R0912,R0913,R0914,R0915,R1720,W0105,W0123,W0201,W0511,W0613,W1113,W1514,E0401,E1121,C0103,C0209,R0402,R1705,R1710,R1725,R1735,W0102,W0212,W0221,W0223,W0231,W0237,W0612,W0621,W0622,W0703,W1309,E1102,E1136 --const-rgx='[a-z_][a-z0-9_]{2,30}$' qlib --init-hook "import astroid; astroid.context.InferenceContext.max_inferred = 500"
|
||||||
|
|
||||||
|
# The following flake8 error codes were ignored:
|
||||||
|
# E501 line too long
|
||||||
|
# Description: We have used black to limit the length of each line to 120.
|
||||||
|
# F541 f-string is missing placeholders
|
||||||
|
# Description: The same thing is done when using pylint for detection.
|
||||||
|
# E266 too many leading '#' for block comment
|
||||||
|
# Description: To make the code more readable, a lot of "#" is used.
|
||||||
|
# This error code appears centrally in:
|
||||||
|
# qlib/backtest/executor.py
|
||||||
|
# qlib/data/ops.py
|
||||||
|
# qlib/utils/__init__.py
|
||||||
|
# E402 module level import not at top of file
|
||||||
|
# Description: There are times when module level import is not available at the top of the file.
|
||||||
|
# W503 line break before binary operator
|
||||||
|
# Description: Since black formats the length of each line of code, it has to perform a line break when a line of arithmetic is too long.
|
||||||
|
# E731 do not assign a lambda expression, use a def
|
||||||
|
# Description: Restricts the use of lambda expressions, but at some point lambda expressions are required.
|
||||||
|
# E203 whitespace before ':'
|
||||||
|
# Description: If there is whitespace before ":", it cannot pass the black check.
|
||||||
|
- name: Check Qlib with flake8
|
||||||
|
run: |
|
||||||
|
pip install --upgrade pip
|
||||||
|
pip install flake8
|
||||||
|
flake8 --ignore=E501,F541,E266,E402,W503,E731,E203 --per-file-ignores="__init__.py:F401,F403" qlib
|
||||||
|
|
||||||
|
# https://github.com/python/mypy/issues/10600
|
||||||
|
- name: Check Qlib with mypy
|
||||||
|
run: |
|
||||||
|
pip install mypy
|
||||||
|
mypy qlib --install-types --non-interactive || true
|
||||||
|
mypy qlib
|
||||||
|
|
||||||
- name: Test data downloads
|
- name: Test data downloads
|
||||||
run: |
|
run: |
|
||||||
python scripts/get_data.py qlib_data --target_dir ~/.qlib/qlib_data/cn_data --interval 1d --region cn
|
python scripts/get_data.py qlib_data --name qlib_data_simple --target_dir ~/.qlib/qlib_data/cn_data_simple --interval 1d --region cn
|
||||||
|
python -c "import os; userpath=os.path.expanduser('~'); os.rename(userpath + '/.qlib/qlib_data/cn_data_simple', userpath + '/.qlib/qlib_data/cn_data')"
|
||||||
|
azcopy copy https://qlibpublic.blob.core.windows.net/data /tmp/qlibpublic --recursive
|
||||||
|
mv /tmp/qlibpublic/data tests/.data
|
||||||
|
|
||||||
- name: Test workflow by config (install from pip)
|
- name: Test workflow by config (install from pip)
|
||||||
run: |
|
run: |
|
||||||
@@ -48,6 +122,7 @@ jobs:
|
|||||||
- name: Install Qlib from source
|
- name: Install Qlib from source
|
||||||
run: |
|
run: |
|
||||||
pip install --upgrade cython jupyter jupyter_contrib_nbextensions numpy scipy scikit-learn # installing without this line will cause errors on GitHub Actions, while instsalling locally won't
|
pip install --upgrade cython jupyter jupyter_contrib_nbextensions numpy scipy scikit-learn # installing without this line will cause errors on GitHub Actions, while instsalling locally won't
|
||||||
|
pip install gym tianshou torch
|
||||||
pip install -e .
|
pip install -e .
|
||||||
|
|
||||||
- name: Install test dependencies
|
- name: Install test dependencies
|
||||||
@@ -57,10 +132,10 @@ jobs:
|
|||||||
|
|
||||||
- name: Unit tests with Pytest
|
- name: Unit tests with Pytest
|
||||||
run: |
|
run: |
|
||||||
|
pip install -r scripts/data_collector/pit/requirements.txt
|
||||||
cd tests
|
cd tests
|
||||||
python -m pytest . --durations=10
|
python -m pytest . --durations=10
|
||||||
|
|
||||||
- name: Test workflow by config (install from source)
|
- name: Test workflow by config (install from source)
|
||||||
run: |
|
run: |
|
||||||
python qlib/workflow/cli.py examples/benchmarks/LightGBM/workflow_config_lightgbm_Alpha158.yaml
|
python qlib/workflow/cli.py examples/benchmarks/LightGBM/workflow_config_lightgbm_Alpha158.yaml
|
||||||
|
|
||||||
|
|||||||
21
.github/workflows/test_macos.yml
vendored
21
.github/workflows/test_macos.yml
vendored
@@ -34,10 +34,24 @@ jobs:
|
|||||||
python -m black qlib -l 120 --check --diff
|
python -m black qlib -l 120 --check --diff
|
||||||
# Test Qlib installed with pip
|
# Test Qlib installed with pip
|
||||||
|
|
||||||
|
- name: Check Qlib with flake8
|
||||||
|
run: |
|
||||||
|
pip install --upgrade pip
|
||||||
|
pip install flake8
|
||||||
|
flake8 --ignore=E501,F541,E266,E402,W503,E731,E203 --per-file-ignores="__init__.py:F401,F403" qlib
|
||||||
|
|
||||||
- name: Install Qlib with pip
|
- name: Install Qlib with pip
|
||||||
run: |
|
run: |
|
||||||
python -m pip install numpy==1.19.5
|
python -m pip install numpy==1.19.5
|
||||||
python -m pip install pyqlib --ignore-installed ruamel.yaml numpy
|
python -m pip install pyqlib --ignore-installed ruamel.yaml numpy
|
||||||
|
- name: Make html with sphnix
|
||||||
|
run: |
|
||||||
|
pip install -U sphinx
|
||||||
|
pip install sphinx_rtd_theme readthedocs_sphinx_ext
|
||||||
|
pip install --exists-action=w --no-cache-dir -r docs/requirements.txt
|
||||||
|
cd docs
|
||||||
|
sphinx-build -b html . build
|
||||||
|
cd ..
|
||||||
- name: Install Lightgbm for MacOS
|
- name: Install Lightgbm for MacOS
|
||||||
run: |
|
run: |
|
||||||
/bin/bash -c "$(curl -fsSL https://raw.githubusercontent.com/Microsoft/qlib/main/.github/brew_install.sh)"
|
/bin/bash -c "$(curl -fsSL https://raw.githubusercontent.com/Microsoft/qlib/main/.github/brew_install.sh)"
|
||||||
@@ -49,7 +63,10 @@ jobs:
|
|||||||
brew install libomp.rb
|
brew install libomp.rb
|
||||||
- name: Test data downloads
|
- name: Test data downloads
|
||||||
run: |
|
run: |
|
||||||
python scripts/get_data.py qlib_data --target_dir ~/.qlib/qlib_data/cn_data --interval 1d --region cn
|
python scripts/get_data.py qlib_data --name qlib_data_simple --target_dir ~/.qlib/qlib_data/cn_data_simple --interval 1d --region cn
|
||||||
|
python -c "import os; userpath=os.path.expanduser('~'); os.rename(userpath + '/.qlib/qlib_data/cn_data_simple', userpath + '/.qlib/qlib_data/cn_data')"
|
||||||
|
azcopy copy https://qlibpublic.blob.core.windows.net/data /tmp/qlibpublic --recursive
|
||||||
|
mv /tmp/qlibpublic/data tests/.data
|
||||||
- name: Test workflow by config (install from pip)
|
- name: Test workflow by config (install from pip)
|
||||||
run: |
|
run: |
|
||||||
python qlib/workflow/cli.py examples/benchmarks/LightGBM/workflow_config_lightgbm_Alpha158.yaml
|
python qlib/workflow/cli.py examples/benchmarks/LightGBM/workflow_config_lightgbm_Alpha158.yaml
|
||||||
@@ -60,6 +77,7 @@ jobs:
|
|||||||
python -m pip install --upgrade cython
|
python -m pip install --upgrade cython
|
||||||
python -m pip install numpy jupyter jupyter_contrib_nbextensions
|
python -m pip install numpy jupyter jupyter_contrib_nbextensions
|
||||||
python -m pip install -U scipy scikit-learn # installing without this line will cause errors on GitHub Actions, while instsalling locally won't
|
python -m pip install -U scipy scikit-learn # installing without this line will cause errors on GitHub Actions, while instsalling locally won't
|
||||||
|
python -m pip install gym tianshou torch
|
||||||
pip install -e .
|
pip install -e .
|
||||||
- name: Install test dependencies
|
- name: Install test dependencies
|
||||||
run: |
|
run: |
|
||||||
@@ -68,6 +86,7 @@ jobs:
|
|||||||
python -m pip install black pytest
|
python -m pip install black pytest
|
||||||
- name: Unit tests with Pytest
|
- name: Unit tests with Pytest
|
||||||
run: |
|
run: |
|
||||||
|
pip install -r scripts/data_collector/pit/requirements.txt
|
||||||
cd tests
|
cd tests
|
||||||
python -m pytest . --durations=0
|
python -m pytest . --durations=0
|
||||||
- name: Test workflow by config (install from source)
|
- name: Test workflow by config (install from source)
|
||||||
|
|||||||
5
.gitignore
vendored
5
.gitignore
vendored
@@ -27,6 +27,10 @@ examples/estimator/estimator_example/
|
|||||||
|
|
||||||
*.egg-info/
|
*.egg-info/
|
||||||
|
|
||||||
|
# test related
|
||||||
|
test-output.xml
|
||||||
|
.output
|
||||||
|
.data
|
||||||
|
|
||||||
# special software
|
# special software
|
||||||
mlruns/
|
mlruns/
|
||||||
@@ -34,6 +38,7 @@ mlruns/
|
|||||||
tags
|
tags
|
||||||
|
|
||||||
.pytest_cache/
|
.pytest_cache/
|
||||||
|
.mypy_cache/
|
||||||
.vscode/
|
.vscode/
|
||||||
|
|
||||||
*.swp
|
*.swp
|
||||||
|
|||||||
17
.mypy.ini
Normal file
17
.mypy.ini
Normal file
@@ -0,0 +1,17 @@
|
|||||||
|
[mypy]
|
||||||
|
exclude = (?x)(
|
||||||
|
^qlib/backtest
|
||||||
|
| ^qlib/contrib
|
||||||
|
| ^qlib/data
|
||||||
|
| ^qlib/model
|
||||||
|
| ^qlib/strategy
|
||||||
|
| ^qlib/tests
|
||||||
|
| ^qlib/utils
|
||||||
|
| ^qlib/workflow
|
||||||
|
| ^qlib/config\.py$
|
||||||
|
| ^qlib/log\.py$
|
||||||
|
| ^qlib/__init__\.py$
|
||||||
|
)
|
||||||
|
ignore_missing_imports = true
|
||||||
|
disallow_incomplete_defs = true
|
||||||
|
follow_imports = skip
|
||||||
12
.pre-commit-config.yaml
Normal file
12
.pre-commit-config.yaml
Normal file
@@ -0,0 +1,12 @@
|
|||||||
|
repos:
|
||||||
|
- repo: https://github.com/psf/black
|
||||||
|
rev: 22.1.0
|
||||||
|
hooks:
|
||||||
|
- id: black
|
||||||
|
args: ["qlib", "-l 120"]
|
||||||
|
|
||||||
|
- repo: https://github.com/PyCQA/flake8
|
||||||
|
rev: 4.0.1
|
||||||
|
hooks:
|
||||||
|
- id: flake8
|
||||||
|
args: ["--ignore=E501,F541,E266,E402,W503,E731,E203"]
|
||||||
5
.pylintrc
Normal file
5
.pylintrc
Normal file
@@ -0,0 +1,5 @@
|
|||||||
|
[TYPECHECK]
|
||||||
|
# https://stackoverflow.com/a/53572939
|
||||||
|
# List of members which are set dynamically and missed by Pylint inference
|
||||||
|
# system, and so shouldn't trigger E1101 when accessed.
|
||||||
|
generated-members=numpy.*, torch.*
|
||||||
93
README.md
93
README.md
@@ -11,7 +11,11 @@
|
|||||||
Recent released features
|
Recent released features
|
||||||
| Feature | Status |
|
| Feature | Status |
|
||||||
| -- | ------ |
|
| -- | ------ |
|
||||||
| Arctic Provider Backend & Orderbook data example | :hammer: [Rleased](https://github.com/microsoft/qlib/pull/744) on Jan 17, 2022 |
|
| HIST and IGMTF models | :chart_with_upwards_trend: [Released](https://github.com/microsoft/qlib/pull/1040) on Apr 10, 2022 |
|
||||||
|
| Qlib [notebook tutorial](https://github.com/microsoft/qlib/tree/main/examples/tutorial) | 📖 [Released](https://github.com/microsoft/qlib/pull/1037) on Apr 7, 2022 |
|
||||||
|
| Ibovespa index data | :rice: [Released](https://github.com/microsoft/qlib/pull/990) on Apr 6, 2022 |
|
||||||
|
| Point-in-Time database | :hammer: [Released](https://github.com/microsoft/qlib/pull/343) on Mar 10, 2022 |
|
||||||
|
| Arctic Provider Backend & Orderbook data example | :hammer: [Released](https://github.com/microsoft/qlib/pull/744) on Jan 17, 2022 |
|
||||||
| Meta-Learning-based framework & DDG-DA | :chart_with_upwards_trend: :hammer: [Released](https://github.com/microsoft/qlib/pull/743) on Jan 10, 2022 |
|
| Meta-Learning-based framework & DDG-DA | :chart_with_upwards_trend: :hammer: [Released](https://github.com/microsoft/qlib/pull/743) on Jan 10, 2022 |
|
||||||
| Planning-based portfolio optimization | :hammer: [Released](https://github.com/microsoft/qlib/pull/754) on Dec 28, 2021 |
|
| Planning-based portfolio optimization | :hammer: [Released](https://github.com/microsoft/qlib/pull/754) on Dec 28, 2021 |
|
||||||
| Release Qlib v0.8.0 | :octocat: [Released](https://github.com/microsoft/qlib/releases/tag/v0.8.0) on Dec 8, 2021 |
|
| Release Qlib v0.8.0 | :octocat: [Released](https://github.com/microsoft/qlib/releases/tag/v0.8.0) on Dec 8, 2021 |
|
||||||
@@ -28,7 +32,7 @@ Recent released features
|
|||||||
| High-frequency data processing example | :hammer: [Released](https://github.com/microsoft/qlib/pull/257) on Feb 5, 2021 |
|
| High-frequency data processing example | :hammer: [Released](https://github.com/microsoft/qlib/pull/257) on Feb 5, 2021 |
|
||||||
| High-frequency trading example | :chart_with_upwards_trend: [Part of code released](https://github.com/microsoft/qlib/pull/227) on Jan 28, 2021 |
|
| High-frequency trading example | :chart_with_upwards_trend: [Part of code released](https://github.com/microsoft/qlib/pull/227) on Jan 28, 2021 |
|
||||||
| High-frequency data(1min) | :rice: [Released](https://github.com/microsoft/qlib/pull/221) on Jan 27, 2021 |
|
| High-frequency data(1min) | :rice: [Released](https://github.com/microsoft/qlib/pull/221) on Jan 27, 2021 |
|
||||||
| Tabnet Model | :chart_with_upwards_trend: [Released](https://github.com/microsoft/qlib/pull/205) on Jan 22, 2021 |
|
| Tabnet Model | :chart_with_upwards_trend: [Released](https://github.com/microsoft/qlib/pull/205) on Jan 22, 2021 |
|
||||||
|
|
||||||
Features released before 2021 are not listed here.
|
Features released before 2021 are not listed here.
|
||||||
|
|
||||||
@@ -45,34 +49,58 @@ With Qlib, users can easily try ideas to create better Quant investment strategi
|
|||||||
|
|
||||||
For more details, please refer to our paper ["Qlib: An AI-oriented Quantitative Investment Platform"](https://arxiv.org/abs/2009.11189).
|
For more details, please refer to our paper ["Qlib: An AI-oriented Quantitative Investment Platform"](https://arxiv.org/abs/2009.11189).
|
||||||
|
|
||||||
- [**Plans**](#plans)
|
|
||||||
- [Framework of Qlib](#framework-of-qlib)
|
|
||||||
- [Quick Start](#quick-start)
|
|
||||||
- [Installation](#installation)
|
|
||||||
- [Data Preparation](#data-preparation)
|
|
||||||
- [Auto Quant Research Workflow](#auto-quant-research-workflow)
|
|
||||||
- [Building Customized Quant Research Workflow by Code](#building-customized-quant-research-workflow-by-code)
|
|
||||||
- [Main Challenges & Solutions in Quant Research](#main-challenges--solutions-in-quant-research)
|
|
||||||
- [Forecasting: Finding Valuable Signals/Patterns](#forecasting-finding-valuable-signalspatterns)
|
|
||||||
- [**Quant Model (Paper) Zoo**](#quant-model-paper-zoo)
|
|
||||||
- [Run a Single Model](#run-a-single-model)
|
|
||||||
- [Run Multiple Models](#run-multiple-models)
|
|
||||||
- [Adapting to Market Dynamics](#adapting-to-market-dynamics)
|
|
||||||
- [**Quant Dataset Zoo**](#quant-dataset-zoo)
|
|
||||||
- [More About Qlib](#more-about-qlib)
|
|
||||||
- [Offline Mode and Online Mode](#offline-mode-and-online-mode)
|
|
||||||
- [Performance of Qlib Data Server](#performance-of-qlib-data-server)
|
|
||||||
- [Related Reports](#related-reports)
|
|
||||||
- [Contact Us](#contact-us)
|
|
||||||
- [Contributing](#contributing)
|
|
||||||
|
|
||||||
|
<table>
|
||||||
|
<tbody>
|
||||||
|
<tr>
|
||||||
|
<th>Frameworks, Tutorial, Data & DevOps</th>
|
||||||
|
<th>Main Challenges & Solutions in Quant Research</th>
|
||||||
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td>
|
||||||
|
<li><a href="#plans"><strong>Plans</strong></a></li>
|
||||||
|
<li><a href="#framework-of-qlib">Framework of Qlib</a></li>
|
||||||
|
<li><a href="#quick-start">Quick Start</a></li>
|
||||||
|
<ul dir="auto">
|
||||||
|
<li type="circle"><a href="#installation">Installation</a> </li>
|
||||||
|
<li type="circle"><a href="#data-preparation">Data Preparation</a></li>
|
||||||
|
<li type="circle"><a href="#auto-quant-research-workflow">Auto Quant Research Workflow</a></li>
|
||||||
|
<li type="circle"><a href="#building-customized-quant-research-workflow-by-code">Building Customized Quant Research Workflow by Code</a></li></ul>
|
||||||
|
<li><a href="#quant-dataset-zoo"><strong>Quant Dataset Zoo</strong></a></li>
|
||||||
|
<li><a href="#more-about-qlib">More About Qlib</a></li>
|
||||||
|
<li><a href="#offline-mode-and-online-mode">Offline Mode and Online Mode</a>
|
||||||
|
<ul>
|
||||||
|
<li type="circle"><a href="#performance-of-qlib-data-server">Performance of Qlib Data Server</a></li></ul>
|
||||||
|
<li><a href="#related-reports">Related Reports</a></li>
|
||||||
|
<li><a href="#contact-us">Contact Us</a></li>
|
||||||
|
<li><a href="#contributing">Contributing</a></li>
|
||||||
|
</td>
|
||||||
|
<td valign="baseline">
|
||||||
|
<li><a href="#main-challenges--solutions-in-quant-research">Main Challenges & Solutions in Quant Research</a>
|
||||||
|
<ul>
|
||||||
|
<li type="circle"><a href="#forecasting-finding-valuable-signalspatterns">Forecasting: Finding Valuable Signals/Patterns</a>
|
||||||
|
<ul>
|
||||||
|
<li type="disc"><a href="#quant-model-paper-zoo"><strong>Quant Model (Paper) Zoo</strong></a>
|
||||||
|
<ul>
|
||||||
|
<li type="circle"><a href="#run-a-single-model">Run a Single Model</a></li>
|
||||||
|
<li type="circle"><a href="#run-multiple-models">Run Multiple Models</a></li>
|
||||||
|
</ul>
|
||||||
|
</li>
|
||||||
|
</ul>
|
||||||
|
</li>
|
||||||
|
<li type="circle"><a href="#adapting-to-market-dynamics">Adapting to Market Dynamics</a></li>
|
||||||
|
</ul>
|
||||||
|
</li>
|
||||||
|
</td>
|
||||||
|
</tr>
|
||||||
|
</tbody>
|
||||||
|
</table>
|
||||||
|
|
||||||
# Plans
|
# Plans
|
||||||
New features under development(order by estimated release time).
|
New features under development(order by estimated release time).
|
||||||
Your feedbacks about the features are very important.
|
Your feedbacks about the features are very important.
|
||||||
| Feature | Status |
|
<!-- | Feature | Status | -->
|
||||||
| -- | ------ |
|
<!-- | -- | ------ | -->
|
||||||
| Point-in-Time database | Under review: https://github.com/microsoft/qlib/pull/343 |
|
|
||||||
|
|
||||||
# Framework of Qlib
|
# Framework of Qlib
|
||||||
|
|
||||||
@@ -80,7 +108,6 @@ Your feedbacks about the features are very important.
|
|||||||
<img src="docs/_static/img/framework.svg" />
|
<img src="docs/_static/img/framework.svg" />
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
|
|
||||||
At the module level, Qlib is a platform that consists of the above components. The components are designed as loose-coupled modules, and each component could be used stand-alone.
|
At the module level, Qlib is a platform that consists of the above components. The components are designed as loose-coupled modules, and each component could be used stand-alone.
|
||||||
|
|
||||||
| Name | Description |
|
| Name | Description |
|
||||||
@@ -92,6 +119,8 @@ At the module level, Qlib is a platform that consists of the above components. T
|
|||||||
* The modules with hand-drawn style are under development and will be released in the future.
|
* The modules with hand-drawn style are under development and will be released in the future.
|
||||||
* The modules with dashed borders are highly user-customizable and extendible.
|
* The modules with dashed borders are highly user-customizable and extendible.
|
||||||
|
|
||||||
|
(p.s. framework image is created with https://draw.io/)
|
||||||
|
|
||||||
|
|
||||||
# Quick Start
|
# Quick Start
|
||||||
|
|
||||||
@@ -311,6 +340,8 @@ Here is a list of models built on `Qlib`.
|
|||||||
- [TCN based on pytorch (Shaojie Bai, et al. 2018)](examples/benchmarks/TCN/)
|
- [TCN based on pytorch (Shaojie Bai, et al. 2018)](examples/benchmarks/TCN/)
|
||||||
- [ADARNN based on pytorch (YunTao Du, et al. 2021)](examples/benchmarks/ADARNN/)
|
- [ADARNN based on pytorch (YunTao Du, et al. 2021)](examples/benchmarks/ADARNN/)
|
||||||
- [ADD based on pytorch (Hongshun Tang, et al.2020)](examples/benchmarks/ADD/)
|
- [ADD based on pytorch (Hongshun Tang, et al.2020)](examples/benchmarks/ADD/)
|
||||||
|
- [IGMTF based on pytorch (Wentao Xu, et al.2021)](examples/benchmarks/IGMTF/)
|
||||||
|
- [HIST based on pytorch (Wentao Xu, et al.2021)](examples/benchmarks/HIST/)
|
||||||
|
|
||||||
Your PR of new Quant models is highly welcomed.
|
Your PR of new Quant models is highly welcomed.
|
||||||
|
|
||||||
@@ -359,6 +390,8 @@ Dataset plays a very important role in Quant. Here is a list of the datasets bui
|
|||||||
Your PR to build new Quant dataset is highly welcomed.
|
Your PR to build new Quant dataset is highly welcomed.
|
||||||
|
|
||||||
# More About Qlib
|
# More About Qlib
|
||||||
|
If you want to have a quick glance at the most frequently used components of qlib, you can try notebooks [here](examples/tutorial/).
|
||||||
|
|
||||||
The detailed documents are organized in [docs](docs/).
|
The detailed documents are organized in [docs](docs/).
|
||||||
[Sphinx](http://www.sphinx-doc.org) and the readthedocs theme is required to build the documentation in html formats.
|
[Sphinx](http://www.sphinx-doc.org) and the readthedocs theme is required to build the documentation in html formats.
|
||||||
```bash
|
```bash
|
||||||
@@ -441,9 +474,13 @@ If you don't know how to start to contribute, you can refer to the following exa
|
|||||||
| Docs | [Improve docs quality](https://github.com/microsoft/qlib/pull/797/files) ; [Fix a typo](https://github.com/microsoft/qlib/pull/774) |
|
| Docs | [Improve docs quality](https://github.com/microsoft/qlib/pull/797/files) ; [Fix a typo](https://github.com/microsoft/qlib/pull/774) |
|
||||||
| Feature | Implement a [requested feature](https://github.com/microsoft/qlib/projects) like [this](https://github.com/microsoft/qlib/pull/754); [Refactor interfaces](https://github.com/microsoft/qlib/pull/539/files) |
|
| Feature | Implement a [requested feature](https://github.com/microsoft/qlib/projects) like [this](https://github.com/microsoft/qlib/pull/754); [Refactor interfaces](https://github.com/microsoft/qlib/pull/539/files) |
|
||||||
| Dataset | [Add a dataset](https://github.com/microsoft/qlib/pull/733) |
|
| Dataset | [Add a dataset](https://github.com/microsoft/qlib/pull/733) |
|
||||||
| Models | [Implement a new model](https://github.com/microsoft/qlib/pull/689) |
|
| Models | [Implement a new model](https://github.com/microsoft/qlib/pull/689), [some instructions to contribute models](https://github.com/microsoft/qlib/tree/main/examples/benchmarks#contributing) |
|
||||||
|
|
||||||
If you would like to become one of Qlib's maintainers to contribute more (e.g. help merge PR, triage issues), please contact us by email([qlib@microsoft.com](mailto:qlib@microsoft.com)). We are glad to help you to set the right permission.
|
[Good first issues](https://github.com/microsoft/qlib/labels/good%20first%20issue) are labelled to indicate that they are easy to start your contributions.
|
||||||
|
|
||||||
|
You can find some impefect implementation in Qlib by `rg 'TODO|FIXME' qlib`
|
||||||
|
|
||||||
|
If you would like to become one of Qlib's maintainers to contribute more (e.g. help merge PR, triage issues), please contact us by email([qlib@microsoft.com](mailto:qlib@microsoft.com)). We are glad to help to upgrade your permission.
|
||||||
|
|
||||||
## Licence
|
## Licence
|
||||||
Most contributions require you to agree to a
|
Most contributions require you to agree to a
|
||||||
|
|||||||
136
docs/advanced/PIT.rst
Normal file
136
docs/advanced/PIT.rst
Normal file
@@ -0,0 +1,136 @@
|
|||||||
|
.. _pit:
|
||||||
|
|
||||||
|
===========================
|
||||||
|
(P)oint-(I)n-(T)ime Database
|
||||||
|
===========================
|
||||||
|
.. currentmodule:: qlib
|
||||||
|
|
||||||
|
|
||||||
|
Introduction
|
||||||
|
------------
|
||||||
|
Point-in-time data is a very important consideration when performing any sort of historical market analysis.
|
||||||
|
|
||||||
|
For example, let’s say we are backtesting a trading strategy and we are using the past five years of historical data as our input.
|
||||||
|
Our model is assumed to trade once a day, at the market close, and we’ll say we are calculating the trading signal for 1 January 2020 in our backtest. At that point, we should only have data for 1 January 2020, 31 December 2019, 30 December 2019 etc.
|
||||||
|
|
||||||
|
In financial data (especially financial reports), the same piece of data may be amended for multiple times overtime. If we only use the latest version for historical backtesting, data leakage will happen.
|
||||||
|
Point-in-time database is designed for solving this problem to make sure user get the right version of data at any historical timestamp. It will keep the performance of online trading and historical backtesting the same.
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
Data Preparation
|
||||||
|
----------------
|
||||||
|
|
||||||
|
Qlib provides a crawler to help users to download financial data and then a converter to dump the data in Qlib format.
|
||||||
|
Please follow `scripts/data_collector/pit/README.md <https://github.com/microsoft/qlib/tree/main/scripts/data_collector/pit/>`_ to download and convert data.
|
||||||
|
Besides, you can find some additional usage examples there.
|
||||||
|
|
||||||
|
|
||||||
|
File-based design for PIT data
|
||||||
|
------------------------------
|
||||||
|
|
||||||
|
Qlib provides a file-based storage for PIT data.
|
||||||
|
|
||||||
|
For each feature, it contains 4 columns, i.e. date, period, value, _next.
|
||||||
|
Each row corresponds to a statement.
|
||||||
|
|
||||||
|
The meaning of each feature with filename like `XXX_a.data`:
|
||||||
|
|
||||||
|
- `date`: the statement's date of publication.
|
||||||
|
- `period`: the period of the statement. (e.g. it will be quarterly frequency in most of the markets)
|
||||||
|
- If it is an annual period, it will be an integer corresponding to the year
|
||||||
|
- If it is an quarterly periods, it will be an integer like `<year><index of quarter>`. The last two decimal digits represents the index of quarter. Others represent the year.
|
||||||
|
- `value`: the described value
|
||||||
|
- `_next`: the byte index of the next occurance of the field.
|
||||||
|
|
||||||
|
Besides the feature data, an index `XXX_a.index` is included to speed up the querying performance
|
||||||
|
|
||||||
|
The statements are soted by the `date` in ascending order from the beginning of the file.
|
||||||
|
|
||||||
|
.. code-block:: python
|
||||||
|
|
||||||
|
# the data format from XXXX.data
|
||||||
|
array([(20070428, 200701, 0.090219 , 4294967295),
|
||||||
|
(20070817, 200702, 0.13933 , 4294967295),
|
||||||
|
(20071023, 200703, 0.24586301, 4294967295),
|
||||||
|
(20080301, 200704, 0.3479 , 80),
|
||||||
|
(20080313, 200704, 0.395989 , 4294967295),
|
||||||
|
(20080422, 200801, 0.100724 , 4294967295),
|
||||||
|
(20080828, 200802, 0.24996801, 4294967295),
|
||||||
|
(20081027, 200803, 0.33412001, 4294967295),
|
||||||
|
(20090325, 200804, 0.39011699, 4294967295),
|
||||||
|
(20090421, 200901, 0.102675 , 4294967295),
|
||||||
|
(20090807, 200902, 0.230712 , 4294967295),
|
||||||
|
(20091024, 200903, 0.30072999, 4294967295),
|
||||||
|
(20100402, 200904, 0.33546099, 4294967295),
|
||||||
|
(20100426, 201001, 0.083825 , 4294967295),
|
||||||
|
(20100812, 201002, 0.200545 , 4294967295),
|
||||||
|
(20101029, 201003, 0.260986 , 4294967295),
|
||||||
|
(20110321, 201004, 0.30739301, 4294967295),
|
||||||
|
(20110423, 201101, 0.097411 , 4294967295),
|
||||||
|
(20110831, 201102, 0.24825101, 4294967295),
|
||||||
|
(20111018, 201103, 0.318919 , 4294967295),
|
||||||
|
(20120323, 201104, 0.4039 , 420),
|
||||||
|
(20120411, 201104, 0.403925 , 4294967295),
|
||||||
|
(20120426, 201201, 0.112148 , 4294967295),
|
||||||
|
(20120810, 201202, 0.26484701, 4294967295),
|
||||||
|
(20121026, 201203, 0.370487 , 4294967295),
|
||||||
|
(20130329, 201204, 0.45004699, 4294967295),
|
||||||
|
(20130418, 201301, 0.099958 , 4294967295),
|
||||||
|
(20130831, 201302, 0.21044201, 4294967295),
|
||||||
|
(20131016, 201303, 0.30454299, 4294967295),
|
||||||
|
(20140325, 201304, 0.394328 , 4294967295),
|
||||||
|
(20140425, 201401, 0.083217 , 4294967295),
|
||||||
|
(20140829, 201402, 0.16450299, 4294967295),
|
||||||
|
(20141030, 201403, 0.23408499, 4294967295),
|
||||||
|
(20150421, 201404, 0.319612 , 4294967295),
|
||||||
|
(20150421, 201501, 0.078494 , 4294967295),
|
||||||
|
(20150828, 201502, 0.137504 , 4294967295),
|
||||||
|
(20151023, 201503, 0.201709 , 4294967295),
|
||||||
|
(20160324, 201504, 0.26420501, 4294967295),
|
||||||
|
(20160421, 201601, 0.073664 , 4294967295),
|
||||||
|
(20160827, 201602, 0.136576 , 4294967295),
|
||||||
|
(20161029, 201603, 0.188062 , 4294967295),
|
||||||
|
(20170415, 201604, 0.244385 , 4294967295),
|
||||||
|
(20170425, 201701, 0.080614 , 4294967295),
|
||||||
|
(20170728, 201702, 0.15151 , 4294967295),
|
||||||
|
(20171026, 201703, 0.25416601, 4294967295),
|
||||||
|
(20180328, 201704, 0.32954201, 4294967295),
|
||||||
|
(20180428, 201801, 0.088887 , 4294967295),
|
||||||
|
(20180802, 201802, 0.170563 , 4294967295),
|
||||||
|
(20181029, 201803, 0.25522 , 4294967295),
|
||||||
|
(20190329, 201804, 0.34464401, 4294967295),
|
||||||
|
(20190425, 201901, 0.094737 , 4294967295),
|
||||||
|
(20190713, 201902, 0. , 1040),
|
||||||
|
(20190718, 201902, 0.175322 , 4294967295),
|
||||||
|
(20191016, 201903, 0.25581899, 4294967295)],
|
||||||
|
dtype=[('date', '<u4'), ('period', '<u4'), ('value', '<f8'), ('_next', '<u4')])
|
||||||
|
# - each row contains 20 byte
|
||||||
|
|
||||||
|
|
||||||
|
# The data format from XXXX.index. It consists of two parts
|
||||||
|
# 1) the start index of the data. So the first part of the info will be like
|
||||||
|
2007
|
||||||
|
# 2) the remain index data will be like information below
|
||||||
|
# - The data indicate the **byte index** of first data update of a period.
|
||||||
|
# - e.g. Because the info at both byte 80 and 100 corresponds to 200704. The byte index of first occurance (i.e. 100) is recorded in the data.
|
||||||
|
array([ 0, 20, 40, 60, 100,
|
||||||
|
120, 140, 160, 180, 200,
|
||||||
|
220, 240, 260, 280, 300,
|
||||||
|
320, 340, 360, 380, 400,
|
||||||
|
440, 460, 480, 500, 520,
|
||||||
|
540, 560, 580, 600, 620,
|
||||||
|
640, 660, 680, 700, 720,
|
||||||
|
740, 760, 780, 800, 820,
|
||||||
|
840, 860, 880, 900, 920,
|
||||||
|
940, 960, 980, 1000, 1020,
|
||||||
|
1060, 4294967295], dtype=uint32)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
Known limitations:
|
||||||
|
|
||||||
|
- Currently, the PIT database is designed for quarterly or annually factors, which can handle fundamental data of financial reports in most markets.
|
||||||
|
- Qlib leverage the file name to identify the type of the data. File with name like `XXX_q.data` corresponds to quarterly data. File with name like `XXX_a.data` corresponds to annual data.
|
||||||
|
- The caclulation of PIT is not performed in the optimal way. There is great potential to boost the performance of PIT data calcuation.
|
||||||
@@ -52,7 +52,8 @@ Also, ``Qlib`` provides a high-frequency dataset. Users can run a high-frequency
|
|||||||
Qlib Format Dataset
|
Qlib Format Dataset
|
||||||
--------------------
|
--------------------
|
||||||
``Qlib`` has provided an off-the-shelf dataset in `.bin` format, users could use the script ``scripts/get_data.py`` to download the China-Stock dataset as follows.
|
``Qlib`` has provided an off-the-shelf dataset in `.bin` format, users could use the script ``scripts/get_data.py`` to download the China-Stock dataset as follows.
|
||||||
The price volume data look different from the actual dealling price because of they are **adjusted** (`adjusted price <https://www.investopedia.com/terms/a/adjusted_closing_price.asp>`_). And then you may find that the adjusted price may be different from different data sources. This is because different data sources may vary in the way of adjusting prices. Qlib normalize the price on first trading day of each stock to 1 when adjusting them.
|
The price volume data look different from the actual dealling price because of they are **adjusted** (`adjusted price <https://www.investopedia.com/terms/a/adjusted_closing_price.asp>`_). And then you may find that the adjusted price may be different from different data sources. This is because different data sources may vary in the way of adjusting prices. Qlib normalize the price on first trading day of each stock to 1 when adjusting them.
|
||||||
|
Users can leverage `$factor` to get the original trading price (e.g. `$close / $factor` to get the original close price).
|
||||||
|
|
||||||
.. code-block:: bash
|
.. code-block:: bash
|
||||||
|
|
||||||
@@ -436,7 +437,7 @@ Dataset
|
|||||||
|
|
||||||
The ``Dataset`` module in ``Qlib`` aims to prepare data for model training and inferencing.
|
The ``Dataset`` module in ``Qlib`` aims to prepare data for model training and inferencing.
|
||||||
|
|
||||||
The motivation of this module is that we want to maximize the flexibility of of different models to handle data that are suitable for themselves. This module gives the model the flexibility to process their data in an unique way. For instance, models such as ``GBDT`` may work well on data that contains `nan` or `None` value, while neural networks such as ``MLP`` will break down on such data.
|
The motivation of this module is that we want to maximize the flexibility of different models to handle data that are suitable for themselves. This module gives the model the flexibility to process their data in an unique way. For instance, models such as ``GBDT`` may work well on data that contains `nan` or `None` value, while neural networks such as ``MLP`` will break down on such data.
|
||||||
|
|
||||||
If user's model need process its data in a different way, user could implement his own ``Dataset`` class. If the model's
|
If user's model need process its data in a different way, user could implement his own ``Dataset`` class. If the model's
|
||||||
data processing is not special, ``DatasetH`` can be used directly.
|
data processing is not special, ``DatasetH`` can be used directly.
|
||||||
|
|||||||
@@ -28,4 +28,11 @@ The frequency of trading algorithm, decision content and execution environment c
|
|||||||
Example
|
Example
|
||||||
===========================
|
===========================
|
||||||
|
|
||||||
An example of nested decision execution framework for high-frequency can be found `here <https://github.com/microsoft/qlib/blob/main/examples/nested_decision_execution/workflow.py>`_.
|
An example of nested decision execution framework for high-frequency can be found `here <https://github.com/microsoft/qlib/blob/main/examples/nested_decision_execution/workflow.py>`_.
|
||||||
|
|
||||||
|
|
||||||
|
Besides, the above examples, here are some other related work about high-frequency trading in Qlib.
|
||||||
|
|
||||||
|
- `Prediction with high-frequency data <https://github.com/microsoft/qlib/tree/main/examples/highfreq#benchmarks-performance-predicting-the-price-trend-in-high-frequency-data>`_
|
||||||
|
- `Examples <https://github.com/microsoft/qlib/blob/main/examples/orderbook_data/>`_ to extract features form high-frequency data without fixed frequency.
|
||||||
|
- `A paper <https://github.com/microsoft/qlib/tree/high-freq-execution#high-frequency-execution>`_ for high-frequency trading.
|
||||||
|
|||||||
@@ -143,3 +143,9 @@ Here is a simple exampke of what is done in ``PortAnaRecord``, which users can r
|
|||||||
print(analysis_df)
|
print(analysis_df)
|
||||||
|
|
||||||
For more information about the APIs, please refer to `Record Template API <../reference/api.html#module-qlib.workflow.record_temp>`_.
|
For more information about the APIs, please refer to `Record Template API <../reference/api.html#module-qlib.workflow.record_temp>`_.
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
Known Limitations
|
||||||
|
=================
|
||||||
|
- The Python objects are saved based on pickle, which may results in issues when the environment dumping objects and loading objects are different.
|
||||||
|
|||||||
@@ -20,6 +20,9 @@ Introduction
|
|||||||
- model_performance_graph
|
- model_performance_graph
|
||||||
|
|
||||||
|
|
||||||
|
All of the accumulated profit metrics(e.g. return, max drawdown) in Qlib are calculated by summation.
|
||||||
|
This avoids the metrics or the plots being skewed exponentially over time.
|
||||||
|
|
||||||
Graphical Reports
|
Graphical Reports
|
||||||
===================
|
===================
|
||||||
|
|
||||||
@@ -101,7 +104,7 @@ Graphical Result
|
|||||||
- Axis Y:
|
- Axis Y:
|
||||||
- `ic`
|
- `ic`
|
||||||
The `Pearson correlation coefficient` series between `label` and `prediction score`.
|
The `Pearson correlation coefficient` series between `label` and `prediction score`.
|
||||||
In the above example, the `label` is formulated as `Ref($close, -1)/$close - 1`. Please refer to `Data Feature <data.html#feature>`_ for more details.
|
In the above example, the `label` is formulated as `Ref($close, -2)/Ref($close, -1)-1`. Please refer to `Data Feature <data.html#feature>`_ for more details.
|
||||||
|
|
||||||
- `rank_ic`
|
- `rank_ic`
|
||||||
The `Spearman's rank correlation coefficient` series between `label` and `prediction score`.
|
The `Spearman's rank correlation coefficient` series between `label` and `prediction score`.
|
||||||
|
|||||||
@@ -24,11 +24,8 @@ BaseStrategy
|
|||||||
|
|
||||||
Qlib provides a base class ``qlib.strategy.base.BaseStrategy``. All strategy classes need to inherit the base class and implement its interface.
|
Qlib provides a base class ``qlib.strategy.base.BaseStrategy``. All strategy classes need to inherit the base class and implement its interface.
|
||||||
|
|
||||||
- `get_risk_degree`
|
- `generate_trade_decision`
|
||||||
Return the proportion of your total value you will use in investment. Dynamically risk_degree will result in Market timing.
|
generate_trade_decision is a key interface that generates trade decisions in each trading bar.
|
||||||
|
|
||||||
- `generate_order_list`
|
|
||||||
Return the order list.
|
|
||||||
The frequency to call this method depends on the executor frequency("time_per_step"="day" by default). But the trading frequency can be decided by users' implementation.
|
The frequency to call this method depends on the executor frequency("time_per_step"="day" by default). But the trading frequency can be decided by users' implementation.
|
||||||
For example, if the user wants to trading in weekly while the `time_per_step` is "day" in executor, user can return non-empty TradeDecision weekly(otherwise return empty like `this <https://github.com/microsoft/qlib/blob/main/qlib/contrib/strategy/signal_strategy.py#L132>`_ ).
|
For example, if the user wants to trading in weekly while the `time_per_step` is "day" in executor, user can return non-empty TradeDecision weekly(otherwise return empty like `this <https://github.com/microsoft/qlib/blob/main/qlib/contrib/strategy/signal_strategy.py#L132>`_ ).
|
||||||
|
|
||||||
@@ -69,18 +66,24 @@ TopkDropoutStrategy
|
|||||||
- Adopt the ``Topk-Drop`` algorithm to calculate the target amount of each stock
|
- Adopt the ``Topk-Drop`` algorithm to calculate the target amount of each stock
|
||||||
|
|
||||||
.. note::
|
.. note::
|
||||||
``Topk-Drop`` algorithm:
|
There are two parameters for the ``Topk-Drop`` algorithm:
|
||||||
|
|
||||||
- `Topk`: The number of stocks held
|
- `Topk`: The number of stocks held
|
||||||
- `Drop`: The number of stocks sold on each trading day
|
- `Drop`: The number of stocks sold on each trading day
|
||||||
|
|
||||||
Currently, the number of held stocks is `Topk`.
|
In general, the number of stocks currently held is `Topk`, with the exception of being zero at the beginning period of trading.
|
||||||
On each trading day, the `Drop` number of held stocks with the worst `prediction score` will be sold, and the same number of unheld stocks with the best `prediction score` will be bought.
|
For each trading day, let $d$ be the number of the instruments currently held and with a rank $\gt K$ when ranked by the prediction scores from high to low.
|
||||||
|
Then `d` number of stocks currently held with the worst `prediction score` will be sold, and the same number of unheld stocks with the best `prediction score` will be bought.
|
||||||
|
|
||||||
|
In general, $d=$`Drop`, especially when the pool of the candidate instruments is large, $K$ is large, and `Drop` is small.
|
||||||
|
|
||||||
|
In most cases, ``TopkDrop`` algorithm sells and buys `Drop` stocks every trading day, which yields a turnover rate of 2$\times$`Drop`/$K$.
|
||||||
|
|
||||||
|
The following images illustrate a typical scenario.
|
||||||
.. image:: ../_static/img/topk_drop.png
|
.. image:: ../_static/img/topk_drop.png
|
||||||
:alt: Topk-Drop
|
:alt: Topk-Drop
|
||||||
|
|
||||||
``TopkDrop`` algorithm sells `Drop` stocks every trading day, which guarantees a fixed turnover rate.
|
|
||||||
|
|
||||||
- Generate the order list from the target amount
|
- Generate the order list from the target amount
|
||||||
|
|
||||||
@@ -126,7 +129,9 @@ A prediction sample is shown as follows.
|
|||||||
|
|
||||||
Normally, the prediction score is the output of the models. But some models are learned from a label with a different scale. So the scale of the prediction score may be different from your expectation(e.g. the return of instruments).
|
Normally, the prediction score is the output of the models. But some models are learned from a label with a different scale. So the scale of the prediction score may be different from your expectation(e.g. the return of instruments).
|
||||||
|
|
||||||
Qlib didn't add a step to scale the prediction score to a unified scale. Because not every trading strategy cares about the scale(e.g. TopkDropoutStrategy only cares about the order). So the strategy is responsible for rescaling the prediction score(e.g. some portfolio-optimization-based strategies may require a meaningful scale).
|
Qlib didn't add a step to scale the prediction score to a unified scale due to the following reasons.
|
||||||
|
- Because not every trading strategy cares about the scale(e.g. TopkDropoutStrategy only cares about the order). So the strategy is responsible for rescaling the prediction score(e.g. some portfolio-optimization-based strategies may require a meaningful scale).
|
||||||
|
- The model has the flexibility to define the target, loss, and data processing. So we don't think there is a silver bullet to rescale it back directly barely based on the model's outputs. If you want to scale it back to some meaningful values(e.g. stock returns.), an intuitive solution is to create a regression model for the model's recent outputs and your recent target values.
|
||||||
|
|
||||||
Running backtest
|
Running backtest
|
||||||
-----------------
|
-----------------
|
||||||
@@ -162,12 +167,9 @@ Running backtest
|
|||||||
start_time="2017-01-01", end_time="2020-08-01", strategy=strategy_obj
|
start_time="2017-01-01", end_time="2020-08-01", strategy=strategy_obj
|
||||||
)
|
)
|
||||||
analysis = dict()
|
analysis = dict()
|
||||||
analysis["excess_return_without_cost"] = risk_analysis(
|
# default frequency will be daily (i.e. "day")
|
||||||
report_normal["return"] - report_normal["bench"], freq=analysis_freq
|
analysis["excess_return_without_cost"] = risk_analysis(report_normal["return"] - report_normal["bench"])
|
||||||
)
|
analysis["excess_return_with_cost"] = risk_analysis(report_normal["return"] - report_normal["bench"] - report_normal["cost"])
|
||||||
analysis["excess_return_with_cost"] = risk_analysis(
|
|
||||||
report_normal["return"] - report_normal["bench"] - report_normal["cost"], freq=analysis_freq
|
|
||||||
)
|
|
||||||
|
|
||||||
analysis_df = pd.concat(analysis) # type: pd.DataFrame
|
analysis_df = pd.concat(analysis) # type: pd.DataFrame
|
||||||
pprint(analysis_df)
|
pprint(analysis_df)
|
||||||
@@ -192,6 +194,14 @@ Running backtest
|
|||||||
qlib.init(provider_uri=<qlib data dir>)
|
qlib.init(provider_uri=<qlib data dir>)
|
||||||
|
|
||||||
CSI300_BENCH = "SH000300"
|
CSI300_BENCH = "SH000300"
|
||||||
|
# Benchmark is for calculating the excess return of your strategy.
|
||||||
|
# Its data format will be like **ONE normal instrument**.
|
||||||
|
# For example, you can query its data with the code below
|
||||||
|
# `D.features(["SH000300"], ["$close"], start_time='2010-01-01', end_time='2017-12-31', freq='day')`
|
||||||
|
# It is different from the argument `market`, which indicates a universe of stocks (e.g. **A SET** of stocks like csi300)
|
||||||
|
# For example, you can query all data from a stock market with the code below.
|
||||||
|
# ` D.features(D.instruments(market='csi300'), ["$close"], start_time='2010-01-01', end_time='2017-12-31', freq='day')`
|
||||||
|
|
||||||
FREQ = "day"
|
FREQ = "day"
|
||||||
STRATEGY_CONFIG = {
|
STRATEGY_CONFIG = {
|
||||||
"topk": 50,
|
"topk": 50,
|
||||||
|
|||||||
@@ -233,7 +233,7 @@ The meaning of each field is as follows:
|
|||||||
Dataset Section
|
Dataset Section
|
||||||
~~~~~~~~~~~~~~~~~~~~
|
~~~~~~~~~~~~~~~~~~~~
|
||||||
|
|
||||||
The `dataset` field describes the parameters for the ``Dataset`` module in ``Qlib`` as well those for the module ``DataHandler``. For more information about the ``Dataset`` module, please refer to `Qlib Model <../component/data.html#dataset>`_.
|
The `dataset` field describes the parameters for the ``Dataset`` module in ``Qlib`` as well those for the module ``DataHandler``. For more information about the ``Dataset`` module, please refer to `Qlib Data <../component/data.html#dataset>`_.
|
||||||
|
|
||||||
The keywords arguments configuration of the ``DataHandler`` is as follows:
|
The keywords arguments configuration of the ``DataHandler`` is as follows:
|
||||||
|
|
||||||
@@ -248,7 +248,7 @@ The keywords arguments configuration of the ``DataHandler`` is as follows:
|
|||||||
|
|
||||||
Users can refer to the document of `DataHandler <../component/data.html#datahandler>`_ for more information about the meaning of each field in the configuration.
|
Users can refer to the document of `DataHandler <../component/data.html#datahandler>`_ for more information about the meaning of each field in the configuration.
|
||||||
|
|
||||||
Here is the configuration for the ``Dataset`` module which will take care of data preprossing and slicing during the training and testing phase.
|
Here is the configuration for the ``Dataset`` module which will take care of data preprocessing and slicing during the training and testing phase.
|
||||||
|
|
||||||
.. code-block:: YAML
|
.. code-block:: YAML
|
||||||
|
|
||||||
|
|||||||
12
docs/conf.py
12
docs/conf.py
@@ -54,9 +54,9 @@ master_doc = "index"
|
|||||||
|
|
||||||
|
|
||||||
# General information about the project.
|
# General information about the project.
|
||||||
project = u"QLib"
|
project = "QLib"
|
||||||
copyright = u"Microsoft"
|
copyright = "Microsoft"
|
||||||
author = u"Microsoft"
|
author = "Microsoft"
|
||||||
|
|
||||||
# The version info for the project you're documenting, acts as replacement for
|
# The version info for the project you're documenting, acts as replacement for
|
||||||
# |version| and |release|, also used in various other places throughout the
|
# |version| and |release|, also used in various other places throughout the
|
||||||
@@ -174,7 +174,7 @@ latex_elements = {
|
|||||||
# (source start file, target name, title,
|
# (source start file, target name, title,
|
||||||
# author, documentclass [howto, manual, or own class]).
|
# author, documentclass [howto, manual, or own class]).
|
||||||
latex_documents = [
|
latex_documents = [
|
||||||
(master_doc, "qlib.tex", u"QLib Documentation", u"Microsoft", "manual"),
|
(master_doc, "qlib.tex", "QLib Documentation", "Microsoft", "manual"),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
@@ -182,7 +182,7 @@ latex_documents = [
|
|||||||
|
|
||||||
# One entry per manual page. List of tuples
|
# One entry per manual page. List of tuples
|
||||||
# (source start file, name, description, authors, manual section).
|
# (source start file, name, description, authors, manual section).
|
||||||
man_pages = [(master_doc, "qlib", u"QLib Documentation", [author], 1)]
|
man_pages = [(master_doc, "qlib", "QLib Documentation", [author], 1)]
|
||||||
|
|
||||||
|
|
||||||
# -- Options for Texinfo output -------------------------------------------
|
# -- Options for Texinfo output -------------------------------------------
|
||||||
@@ -194,7 +194,7 @@ texinfo_documents = [
|
|||||||
(
|
(
|
||||||
master_doc,
|
master_doc,
|
||||||
"QLib",
|
"QLib",
|
||||||
u"QLib Documentation",
|
"QLib Documentation",
|
||||||
author,
|
author,
|
||||||
"QLib",
|
"QLib",
|
||||||
"One line description of project.",
|
"One line description of project.",
|
||||||
|
|||||||
@@ -14,9 +14,35 @@ Continuous Integration (CI) tools help you stick to the quality standards by run
|
|||||||
|
|
||||||
When you submit a PR request, you can check whether your code passes the CI tests in the "check" section at the bottom of the web page.
|
When you submit a PR request, you can check whether your code passes the CI tests in the "check" section at the bottom of the web page.
|
||||||
|
|
||||||
A common error is the mixed use of space and tab. You can fix the bug by inputing the following code in the command line.
|
1. Qlib will check the code format with black. The PR will raise error if your code does not align to the standard of Qlib(e.g. a common error is the mixed use of space and tab).
|
||||||
|
You can fix the bug by inputing the following code in the command line.
|
||||||
|
|
||||||
.. code-block:: python
|
.. code-block:: bash
|
||||||
|
|
||||||
pip install black
|
pip install black
|
||||||
python -m black . -l 120
|
python -m black . -l 120
|
||||||
|
|
||||||
|
|
||||||
|
2. Qlib will check your code style pylint. The checking command is implemented in [github action workflow](https://github.com/microsoft/qlib/blob/0e8b94a552f1c457cfa6cd2c1bb3b87ebb3fb279/.github/workflows/test.yml#L66).
|
||||||
|
Sometime pylint's restrictions are not that reasonable. You can ignore specific errors like this
|
||||||
|
|
||||||
|
.. code-block:: python
|
||||||
|
|
||||||
|
return -ICLoss()(pred, target, index) # pylint: disable=E1130
|
||||||
|
|
||||||
|
|
||||||
|
3. Qlib will check your code style flake8. The checking command is implemented in [github action workflow](https://github.com/microsoft/qlib/blob/0e8b94a552f1c457cfa6cd2c1bb3b87ebb3fb279/.github/workflows/test.yml#L73).
|
||||||
|
You can fix the bug by inputing the following code in the command line.
|
||||||
|
|
||||||
|
.. code-block:: bash
|
||||||
|
|
||||||
|
flake8 --ignore E501,F541,E402,F401,W503,E741,E266,E203,E302,E731,E262,F523,F821,F811,F841,E713,E265,W291,E712,E722,W293 qlib
|
||||||
|
|
||||||
|
|
||||||
|
4. Qlib has integrated pre-commit, which will make it easier for developers to format their code.
|
||||||
|
Just run the following two commands, and the code will be automatically formatted using black and flake8 when the git commit command is executed.
|
||||||
|
|
||||||
|
.. code-block:: bash
|
||||||
|
|
||||||
|
pip install -e .[dev]
|
||||||
|
pre-commit install
|
||||||
@@ -53,6 +53,7 @@ Document Structure
|
|||||||
Online & Offline mode <advanced/server.rst>
|
Online & Offline mode <advanced/server.rst>
|
||||||
Serialization <advanced/serial.rst>
|
Serialization <advanced/serial.rst>
|
||||||
Task Management <advanced/task_management.rst>
|
Task Management <advanced/task_management.rst>
|
||||||
|
Point-In-Time database <advanced/PIT.rst>
|
||||||
|
|
||||||
.. toctree::
|
.. toctree::
|
||||||
:maxdepth: 3
|
:maxdepth: 3
|
||||||
|
|||||||
@@ -120,6 +120,32 @@ For more details about features, please refer `Feature API <../component/data.ht
|
|||||||
|
|
||||||
.. note:: When calling `D.features()` at the client, use parameter `disk_cache=0` to skip dataset cache, use `disk_cache=1` to generate and use dataset cache. In addition, when calling at the server, users can use `disk_cache=2` to update the dataset cache.
|
.. note:: When calling `D.features()` at the client, use parameter `disk_cache=0` to skip dataset cache, use `disk_cache=1` to generate and use dataset cache. In addition, when calling at the server, users can use `disk_cache=2` to update the dataset cache.
|
||||||
|
|
||||||
|
|
||||||
|
When you are building complicated expressions, implementing all the expressions in a single string may not be easy.
|
||||||
|
For example, it looks quite long and complicated:
|
||||||
|
|
||||||
|
.. code-block:: python
|
||||||
|
|
||||||
|
>> from qlib.data import D
|
||||||
|
>> data = D.features(["sh600519"], ["(($high / $close) + ($open / $close)) * (($high / $close) + ($open / $close)) / ($high / $close) + ($open / $close)"], start_time="20200101")
|
||||||
|
|
||||||
|
|
||||||
|
But using string is not the only way to implement the expression. You can also implement expression by code.
|
||||||
|
Here is an exmaple which does the same thing as above examples.
|
||||||
|
|
||||||
|
|
||||||
|
.. code-block:: python
|
||||||
|
|
||||||
|
>> from qlib.data.ops import *
|
||||||
|
>> f1 = Feature("high") / Feature("close")
|
||||||
|
>> f2 = Feature("open") / Feature("close")
|
||||||
|
>> f3 = f1 + f2
|
||||||
|
>> f4 = f3 * f3 / f3
|
||||||
|
|
||||||
|
>> data = D.features(["sh600519"], [f4], start_time="20200101")
|
||||||
|
>> data.head()
|
||||||
|
|
||||||
|
|
||||||
API
|
API
|
||||||
====================
|
====================
|
||||||
To know more about how to use the Data, go to API Reference: `Data API <../reference/api.html#data>`_
|
To know more about how to use the Data, go to API Reference: `Data API <../reference/api.html#data>`_
|
||||||
|
|||||||
@@ -37,7 +37,8 @@ Initialize Qlib before calling other APIs: run following code in python.
|
|||||||
Parameters
|
Parameters
|
||||||
-------------------
|
-------------------
|
||||||
|
|
||||||
Besides `provider_uri` and `region`, `qlib.init` has other parameters. The following are several important parameters of `qlib.init`:
|
Besides `provider_uri` and `region`, `qlib.init` has other parameters.
|
||||||
|
The following are several important parameters of `qlib.init` (`Qlib` has a lot of config. Only part of parameters are limited here. More detailed setting can be found `here <https://github.com/microsoft/qlib/blob/main/qlib/config.py>`_):
|
||||||
|
|
||||||
- `provider_uri`
|
- `provider_uri`
|
||||||
Type: str. The URI of the Qlib data. For example, it could be the location where the data loaded by ``get_data.py`` are stored.
|
Type: str. The URI of the Qlib data. For example, it could be the location where the data loaded by ``get_data.py`` are stored.
|
||||||
@@ -48,7 +49,7 @@ Besides `provider_uri` and `region`, `qlib.init` has other parameters. The follo
|
|||||||
- ``qlib.constant.REG_CN``: China stock market.
|
- ``qlib.constant.REG_CN``: China stock market.
|
||||||
|
|
||||||
Different modes will result in different trading limitations and costs.
|
Different modes will result in different trading limitations and costs.
|
||||||
The region is just `shortcuts for defining a batch of configurations <https://github.com/microsoft/qlib/blob/main/qlib/config.py#L239>`_. Users can set the key configurations manually if the existing region setting can't meet their requirements.
|
The region is just `shortcuts for defining a batch of configurations <https://github.com/microsoft/qlib/blob/528f74af099bf6156e9480bcd2bb28e453231212/qlib/config.py#L249>`_, which include minimal trading order unit (``trade_unit``), trading limitation (``limit_threshold``) , etc. It is not a necessary part and users can set the key configurations manually if the existing region setting can't meet their requirements.
|
||||||
- `redis_host`
|
- `redis_host`
|
||||||
Type: str, optional parameter(default: "127.0.0.1"), host of `redis`
|
Type: str, optional parameter(default: "127.0.0.1"), host of `redis`
|
||||||
The lock and cache mechanism relies on redis.
|
The lock and cache mechanism relies on redis.
|
||||||
@@ -88,3 +89,9 @@ Besides `provider_uri` and `region`, `qlib.init` has other parameters. The follo
|
|||||||
"task_url": "mongodb://localhost:27017/", # your mongo url
|
"task_url": "mongodb://localhost:27017/", # your mongo url
|
||||||
"task_db_name": "rolling_db", # the database name of Task Management
|
"task_db_name": "rolling_db", # the database name of Task Management
|
||||||
})
|
})
|
||||||
|
|
||||||
|
- `logging_level`
|
||||||
|
The logging level for the system.
|
||||||
|
|
||||||
|
- `kernels`
|
||||||
|
The number of processes used when calculating features in Qlib's expression engine. It is very helpful to set it to 1 when you are debuggin an expression calculating exception
|
||||||
|
|||||||
@@ -6,3 +6,4 @@
|
|||||||
|
|
||||||
[https://www.ijcai.org/Proceedings/2017/0366.pdf](https://www.ijcai.org/Proceedings/2017/0366.pdf)
|
[https://www.ijcai.org/Proceedings/2017/0366.pdf](https://www.ijcai.org/Proceedings/2017/0366.pdf)
|
||||||
|
|
||||||
|
- NOTE: Current version of implementation is just a simplified version of ALSTM. It is an LSTM with attention.
|
||||||
|
|||||||
3
examples/benchmarks/HIST/README.md
Normal file
3
examples/benchmarks/HIST/README.md
Normal file
@@ -0,0 +1,3 @@
|
|||||||
|
# HIST
|
||||||
|
* Code: [https://github.com/Wentao-Xu/HIST](https://github.com/Wentao-Xu/HIST)
|
||||||
|
* Paper: [HIST: A Graph-based Framework for Stock Trend Forecasting via Mining Concept-Oriented Shared InformationAdaRNN: Adaptive Learning and Forecasting for Time Series](https://arxiv.org/abs/2110.13716).
|
||||||
BIN
examples/benchmarks/HIST/qlib_csi300_stock_index.npy
Normal file
BIN
examples/benchmarks/HIST/qlib_csi300_stock_index.npy
Normal file
Binary file not shown.
4
examples/benchmarks/HIST/requirements.txt
Normal file
4
examples/benchmarks/HIST/requirements.txt
Normal file
@@ -0,0 +1,4 @@
|
|||||||
|
pandas==1.1.2
|
||||||
|
numpy==1.21.0
|
||||||
|
scikit_learn==0.23.2
|
||||||
|
torch==1.7.0
|
||||||
92
examples/benchmarks/HIST/workflow_config_hist_Alpha360.yaml
Normal file
92
examples/benchmarks/HIST/workflow_config_hist_Alpha360.yaml
Normal file
@@ -0,0 +1,92 @@
|
|||||||
|
qlib_init:
|
||||||
|
provider_uri: "~/.qlib/qlib_data/cn_data"
|
||||||
|
region: cn
|
||||||
|
market: &market csi300
|
||||||
|
benchmark: &benchmark SH000300
|
||||||
|
data_handler_config: &data_handler_config
|
||||||
|
start_time: 2008-01-01
|
||||||
|
end_time: 2020-08-01
|
||||||
|
fit_start_time: 2008-01-01
|
||||||
|
fit_end_time: 2014-12-31
|
||||||
|
instruments: *market
|
||||||
|
infer_processors:
|
||||||
|
- class: RobustZScoreNorm
|
||||||
|
kwargs:
|
||||||
|
fields_group: feature
|
||||||
|
clip_outlier: true
|
||||||
|
- class: Fillna
|
||||||
|
kwargs:
|
||||||
|
fields_group: feature
|
||||||
|
learn_processors:
|
||||||
|
- class: DropnaLabel
|
||||||
|
- class: CSRankNorm
|
||||||
|
kwargs:
|
||||||
|
fields_group: label
|
||||||
|
label: ["Ref($close, -2) / Ref($close, -1) - 1"]
|
||||||
|
port_analysis_config: &port_analysis_config
|
||||||
|
strategy:
|
||||||
|
class: TopkDropoutStrategy
|
||||||
|
module_path: qlib.contrib.strategy
|
||||||
|
kwargs:
|
||||||
|
signal:
|
||||||
|
- <MODEL>
|
||||||
|
- <DATASET>
|
||||||
|
topk: 50
|
||||||
|
n_drop: 5
|
||||||
|
backtest:
|
||||||
|
start_time: 2017-01-01
|
||||||
|
end_time: 2020-08-01
|
||||||
|
account: 100000000
|
||||||
|
benchmark: *benchmark
|
||||||
|
exchange_kwargs:
|
||||||
|
limit_threshold: 0.095
|
||||||
|
deal_price: close
|
||||||
|
open_cost: 0.0005
|
||||||
|
close_cost: 0.0015
|
||||||
|
min_cost: 5
|
||||||
|
task:
|
||||||
|
model:
|
||||||
|
class: HIST
|
||||||
|
module_path: qlib.contrib.model.pytorch_hist
|
||||||
|
kwargs:
|
||||||
|
d_feat: 6
|
||||||
|
hidden_size: 64
|
||||||
|
num_layers: 2
|
||||||
|
dropout: 0
|
||||||
|
n_epochs: 200
|
||||||
|
lr: 1e-4
|
||||||
|
early_stop: 20
|
||||||
|
metric: ic
|
||||||
|
loss: mse
|
||||||
|
base_model: LSTM
|
||||||
|
model_path: "benchmarks/LSTM/model_lstm_csi300.pkl"
|
||||||
|
stock2concept: "benchmarks/HIST/qlib_csi300_stock2concept.npy"
|
||||||
|
stock_index: "benchmarks/HIST/qlib_csi300_stock_index.npy"
|
||||||
|
GPU: 0
|
||||||
|
dataset:
|
||||||
|
class: DatasetH
|
||||||
|
module_path: qlib.data.dataset
|
||||||
|
kwargs:
|
||||||
|
handler:
|
||||||
|
class: Alpha360
|
||||||
|
module_path: qlib.contrib.data.handler
|
||||||
|
kwargs: *data_handler_config
|
||||||
|
segments:
|
||||||
|
train: [2008-01-01, 2014-12-31]
|
||||||
|
valid: [2015-01-01, 2016-12-31]
|
||||||
|
test: [2017-01-01, 2020-08-01]
|
||||||
|
record:
|
||||||
|
- class: SignalRecord
|
||||||
|
module_path: qlib.workflow.record_temp
|
||||||
|
kwargs:
|
||||||
|
model: <MODEL>
|
||||||
|
dataset: <DATASET>
|
||||||
|
- class: SigAnaRecord
|
||||||
|
module_path: qlib.workflow.record_temp
|
||||||
|
kwargs:
|
||||||
|
ana_long_short: False
|
||||||
|
ann_scaler: 252
|
||||||
|
- class: PortAnaRecord
|
||||||
|
module_path: qlib.workflow.record_temp
|
||||||
|
kwargs:
|
||||||
|
config: *port_analysis_config
|
||||||
4
examples/benchmarks/IGMTF/README.md
Normal file
4
examples/benchmarks/IGMTF/README.md
Normal file
@@ -0,0 +1,4 @@
|
|||||||
|
# IGMTF
|
||||||
|
* Code: [https://github.com/Wentao-Xu/IGMTF](https://github.com/Wentao-Xu/IGMTF)
|
||||||
|
* Paper: [IGMTF: An Instance-wise Graph-based Framework for
|
||||||
|
Multivariate Time Series Forecasting](https://arxiv.org/abs/2109.06489).
|
||||||
4
examples/benchmarks/IGMTF/requirements.txt
Normal file
4
examples/benchmarks/IGMTF/requirements.txt
Normal file
@@ -0,0 +1,4 @@
|
|||||||
|
pandas==1.1.2
|
||||||
|
numpy==1.21.0
|
||||||
|
scikit_learn==0.23.2
|
||||||
|
torch==1.7.0
|
||||||
@@ -0,0 +1,89 @@
|
|||||||
|
qlib_init:
|
||||||
|
provider_uri: "~/.qlib/qlib_data/cn_data"
|
||||||
|
region: cn
|
||||||
|
market: &market csi300
|
||||||
|
benchmark: &benchmark SH000300
|
||||||
|
data_handler_config: &data_handler_config
|
||||||
|
start_time: 2008-01-01
|
||||||
|
end_time: 2020-08-01
|
||||||
|
fit_start_time: 2008-01-01
|
||||||
|
fit_end_time: 2014-12-31
|
||||||
|
instruments: *market
|
||||||
|
infer_processors:
|
||||||
|
- class: RobustZScoreNorm
|
||||||
|
kwargs:
|
||||||
|
fields_group: feature
|
||||||
|
clip_outlier: true
|
||||||
|
- class: Fillna
|
||||||
|
kwargs:
|
||||||
|
fields_group: feature
|
||||||
|
learn_processors:
|
||||||
|
- class: DropnaLabel
|
||||||
|
- class: CSRankNorm
|
||||||
|
kwargs:
|
||||||
|
fields_group: label
|
||||||
|
label: ["Ref($close, -2) / Ref($close, -1) - 1"]
|
||||||
|
port_analysis_config: &port_analysis_config
|
||||||
|
strategy:
|
||||||
|
class: TopkDropoutStrategy
|
||||||
|
module_path: qlib.contrib.strategy
|
||||||
|
kwargs:
|
||||||
|
model: <MODEL>
|
||||||
|
dataset: <DATASET>
|
||||||
|
topk: 50
|
||||||
|
n_drop: 5
|
||||||
|
backtest:
|
||||||
|
start_time: 2017-01-01
|
||||||
|
end_time: 2020-08-01
|
||||||
|
account: 100000000
|
||||||
|
benchmark: *benchmark
|
||||||
|
exchange_kwargs:
|
||||||
|
limit_threshold: 0.095
|
||||||
|
deal_price: close
|
||||||
|
open_cost: 0.0005
|
||||||
|
close_cost: 0.0015
|
||||||
|
min_cost: 5
|
||||||
|
task:
|
||||||
|
model:
|
||||||
|
class: IGMTF
|
||||||
|
module_path: qlib.contrib.model.pytorch_igmtf
|
||||||
|
kwargs:
|
||||||
|
d_feat: 6
|
||||||
|
hidden_size: 64
|
||||||
|
num_layers: 2
|
||||||
|
dropout: 0
|
||||||
|
n_epochs: 200
|
||||||
|
lr: 1e-4
|
||||||
|
early_stop: 20
|
||||||
|
metric: ic
|
||||||
|
loss: mse
|
||||||
|
base_model: LSTM
|
||||||
|
model_path: "benchmarks/LSTM/model_lstm_csi300.pkl"
|
||||||
|
GPU: 0
|
||||||
|
dataset:
|
||||||
|
class: DatasetH
|
||||||
|
module_path: qlib.data.dataset
|
||||||
|
kwargs:
|
||||||
|
handler:
|
||||||
|
class: Alpha360
|
||||||
|
module_path: qlib.contrib.data.handler
|
||||||
|
kwargs: *data_handler_config
|
||||||
|
segments:
|
||||||
|
train: [2008-01-01, 2014-12-31]
|
||||||
|
valid: [2015-01-01, 2016-12-31]
|
||||||
|
test: [2017-01-01, 2020-08-01]
|
||||||
|
record:
|
||||||
|
- class: SignalRecord
|
||||||
|
module_path: qlib.workflow.record_temp
|
||||||
|
kwargs:
|
||||||
|
model: <MODEL>
|
||||||
|
dataset: <DATASET>
|
||||||
|
- class: SigAnaRecord
|
||||||
|
module_path: qlib.workflow.record_temp
|
||||||
|
kwargs:
|
||||||
|
ana_long_short: False
|
||||||
|
ann_scaler: 252
|
||||||
|
- class: PortAnaRecord
|
||||||
|
module_path: qlib.workflow.record_temp
|
||||||
|
kwargs:
|
||||||
|
config: *port_analysis_config
|
||||||
@@ -63,8 +63,6 @@ task:
|
|||||||
module_path: qlib.contrib.model.pytorch_nn
|
module_path: qlib.contrib.model.pytorch_nn
|
||||||
kwargs:
|
kwargs:
|
||||||
loss: mse
|
loss: mse
|
||||||
input_dim: 157
|
|
||||||
output_dim: 1
|
|
||||||
lr: 0.002
|
lr: 0.002
|
||||||
lr_decay: 0.96
|
lr_decay: 0.96
|
||||||
lr_decay_steps: 100
|
lr_decay_steps: 100
|
||||||
@@ -73,6 +71,8 @@ task:
|
|||||||
batch_size: 8192
|
batch_size: 8192
|
||||||
GPU: 0
|
GPU: 0
|
||||||
weight_decay: 0.0002
|
weight_decay: 0.0002
|
||||||
|
pt_model_kwargs:
|
||||||
|
input_dim: 157
|
||||||
dataset:
|
dataset:
|
||||||
class: DatasetH
|
class: DatasetH
|
||||||
module_path: qlib.data.dataset
|
module_path: qlib.data.dataset
|
||||||
|
|||||||
@@ -51,8 +51,6 @@ task:
|
|||||||
module_path: qlib.contrib.model.pytorch_nn
|
module_path: qlib.contrib.model.pytorch_nn
|
||||||
kwargs:
|
kwargs:
|
||||||
loss: mse
|
loss: mse
|
||||||
input_dim: 360
|
|
||||||
output_dim: 1
|
|
||||||
lr: 0.002
|
lr: 0.002
|
||||||
lr_decay: 0.96
|
lr_decay: 0.96
|
||||||
lr_decay_steps: 100
|
lr_decay_steps: 100
|
||||||
@@ -60,6 +58,8 @@ task:
|
|||||||
max_steps: 8000
|
max_steps: 8000
|
||||||
batch_size: 4096
|
batch_size: 4096
|
||||||
GPU: 0
|
GPU: 0
|
||||||
|
pt_model_kwargs:
|
||||||
|
input_dim: 360
|
||||||
dataset:
|
dataset:
|
||||||
class: DatasetH
|
class: DatasetH
|
||||||
module_path: qlib.data.dataset
|
module_path: qlib.data.dataset
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ This page lists a batch of methods designed for alpha seeking. Each method tries
|
|||||||
The alpha is evaluated in two ways.
|
The alpha is evaluated in two ways.
|
||||||
1. The correlation between the alpha and future return.
|
1. The correlation between the alpha and future return.
|
||||||
1. Constructing portfolio based on the alpha and evaluating the final total return.
|
1. Constructing portfolio based on the alpha and evaluating the final total return.
|
||||||
|
- The explanation of metrics can be found [here](https://qlib.readthedocs.io/en/latest/component/report.html#id4)
|
||||||
|
|
||||||
Here are the results of each benchmark model running on Qlib's `Alpha360` and `Alpha158` dataset with China's A shared-stock & CSI300 data respectively. The values of each metric are the mean and std calculated based on 20 runs with different random seeds.
|
Here are the results of each benchmark model running on Qlib's `Alpha360` and `Alpha158` dataset with China's A shared-stock & CSI300 data respectively. The values of each metric are the mean and std calculated based on 20 runs with different random seeds.
|
||||||
|
|
||||||
@@ -16,6 +17,8 @@ The numbers shown below demonstrate the performance of the entire `workflow` of
|
|||||||
> NOTE:
|
> NOTE:
|
||||||
> The backtest start from 0.8.0 is quite different from previous version. Please check out the changelog for the difference.
|
> The backtest start from 0.8.0 is quite different from previous version. Please check out the changelog for the difference.
|
||||||
|
|
||||||
|
> NOTE:
|
||||||
|
> We have very limited resources to implement and finetune the models. We tried our best effort to fairly compare these models. But some models may have greater potential than what it looks like in the table below. Your contribution is highly welcomed to explore their potential.
|
||||||
|
|
||||||
## Alpha158 dataset
|
## Alpha158 dataset
|
||||||
|
|
||||||
@@ -62,7 +65,33 @@ The numbers shown below demonstrate the performance of the entire `workflow` of
|
|||||||
| GATs (Petar Velickovic, et al.) | Alpha360 | 0.0476±0.00 | 0.3508±0.02 | 0.0598±0.00 | 0.4604±0.01 | 0.0824±0.02 | 1.1079±0.26 | -0.0894±0.03 |
|
| GATs (Petar Velickovic, et al.) | Alpha360 | 0.0476±0.00 | 0.3508±0.02 | 0.0598±0.00 | 0.4604±0.01 | 0.0824±0.02 | 1.1079±0.26 | -0.0894±0.03 |
|
||||||
| TCTS(Xueqing Wu, et al.) | Alpha360 | 0.0508±0.00 | 0.3931±0.04 | 0.0599±0.00 | 0.4756±0.03 | 0.0893±0.03 | 1.2256±0.36 | -0.0857±0.02 |
|
| TCTS(Xueqing Wu, et al.) | Alpha360 | 0.0508±0.00 | 0.3931±0.04 | 0.0599±0.00 | 0.4756±0.03 | 0.0893±0.03 | 1.2256±0.36 | -0.0857±0.02 |
|
||||||
| TRA(Hengxu Lin, et al.) | Alpha360 | 0.0485±0.00 | 0.3787±0.03 | 0.0587±0.00 | 0.4756±0.03 | 0.0920±0.03 | 1.2789±0.42 | -0.0834±0.02 |
|
| TRA(Hengxu Lin, et al.) | Alpha360 | 0.0485±0.00 | 0.3787±0.03 | 0.0587±0.00 | 0.4756±0.03 | 0.0920±0.03 | 1.2789±0.42 | -0.0834±0.02 |
|
||||||
|
| IGMTF(Wentao Xu, et al.) | Alpha360 | 0.0480±0.00 | 0.3589±0.02 | 0.0606±0.00 | 0.4773±0.01 | 0.0946±0.02 | 1.3509±0.25 | -0.0716±0.02 |
|
||||||
|
| HIST(Wentao Xu, et al.) | Alpha360 | 0.0522±0.00 | 0.3530±0.01 | 0.0667±0.00 | 0.4576±0.01 | 0.0987±0.02 | 1.3726±0.27 | -0.0681±0.01 |
|
||||||
|
|
||||||
|
|
||||||
- The selected 20 features are based on the feature importance of a lightgbm-based model.
|
- The selected 20 features are based on the feature importance of a lightgbm-based model.
|
||||||
- The base model of DoubleEnsemble is LGBM.
|
- The base model of DoubleEnsemble is LGBM.
|
||||||
- The base model of TCTS is GRU.
|
- The base model of TCTS is GRU.
|
||||||
|
- About the datasets
|
||||||
|
- Alpha158 is a tabular dataset. There are less spatial relationships between different features. Each feature are carefully desgined by human (a.k.a feature engineering)
|
||||||
|
- Alpha360 contains raw price and volue data without much feature engineering. There are strong strong spatial relationships between the features in the time dimension.
|
||||||
|
- The metrics can be categorized into two
|
||||||
|
- Signal-based evaluation: IC, ICIR, Rank IC, Rank ICIR
|
||||||
|
- Portfolio-based metrics: Annualized Return, Information Ratio, Max Drawdown
|
||||||
|
|
||||||
|
|
||||||
|
# Contributing
|
||||||
|
|
||||||
|
Your contributions to new models are highly welcome!
|
||||||
|
|
||||||
|
If you want to contribute your new models, you can follow the steps below.
|
||||||
|
1. Create a folder for your model
|
||||||
|
2. The folder contains following items(you can refer to [this example](https://github.com/microsoft/qlib/tree/main/examples/benchmarks/TCTS)).
|
||||||
|
- `requirements.txt`: required dependencies.
|
||||||
|
- `README.md`: a brief introduction to your models
|
||||||
|
- `workflow_config_<model name>_<dataset>.yaml`: a configuration which can read by `qrun`. You are encouraged to run your model in all datasets.
|
||||||
|
3. You can integrate your model as a module [in this folder](https://github.com/microsoft/qlib/tree/main/qlib/contrib/model).
|
||||||
|
4. Please updated your results in the benchmark tables, e.g. [Alpha360](#alpha158-dataset), [Alpha158](#alpha158-dataset)(the values of each metric are the mean and std calculated based on 20 runs with different random seeds, if you don't have enough computational resource, you can ask for help in the PR).
|
||||||
|
5. Update the info in the index page in the [news list](https://github.com/microsoft/qlib#newspaper-whats-new----sparkling_heart) and [model list](https://github.com/microsoft/qlib#quant-model-paper-zoo).
|
||||||
|
|
||||||
|
Finally, you can send PR for review. ([here is an example](https://github.com/microsoft/qlib/pull/1040))
|
||||||
|
|||||||
@@ -6,8 +6,7 @@ import torch
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from qlib.utils import init_instance_by_config
|
from qlib.data.dataset import DatasetH
|
||||||
from qlib.data.dataset import DatasetH, DataHandler
|
|
||||||
|
|
||||||
|
|
||||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
@@ -95,7 +94,7 @@ class MTSDatasetH(DatasetH):
|
|||||||
shuffle=True,
|
shuffle=True,
|
||||||
pin_memory=False,
|
pin_memory=False,
|
||||||
drop_last=False,
|
drop_last=False,
|
||||||
**kwargs
|
**kwargs,
|
||||||
):
|
):
|
||||||
|
|
||||||
assert horizon > 0, "please specify `horizon` to avoid data leakage"
|
assert horizon > 0, "please specify `horizon` to avoid data leakage"
|
||||||
@@ -150,8 +149,15 @@ class MTSDatasetH(DatasetH):
|
|||||||
|
|
||||||
def _prepare_seg(self, slc, **kwargs):
|
def _prepare_seg(self, slc, **kwargs):
|
||||||
fn = _get_date_parse_fn(self._index[0][1])
|
fn = _get_date_parse_fn(self._index[0][1])
|
||||||
start_date = fn(slc.start)
|
|
||||||
end_date = fn(slc.stop)
|
if isinstance(slc, slice):
|
||||||
|
start, stop = slc.start, slc.stop
|
||||||
|
elif isinstance(slc, (list, tuple)):
|
||||||
|
start, stop = slc
|
||||||
|
else:
|
||||||
|
raise NotImplementedError(f"This type of input is not supported")
|
||||||
|
start_date = fn(start)
|
||||||
|
end_date = fn(stop)
|
||||||
obj = copy.copy(self) # shallow copy
|
obj = copy.copy(self) # shallow copy
|
||||||
# NOTE: Seriable will disable copy `self._data` so we manually assign them here
|
# NOTE: Seriable will disable copy `self._data` so we manually assign them here
|
||||||
obj._data = self._data
|
obj._data = self._data
|
||||||
|
|||||||
@@ -130,7 +130,7 @@ class TRAModel(Model):
|
|||||||
|
|
||||||
if prob is not None:
|
if prob is not None:
|
||||||
P = sinkhorn(-L, epsilon=0.01) # sample assignment matrix
|
P = sinkhorn(-L, epsilon=0.01) # sample assignment matrix
|
||||||
lamb = self.lamb * (self.rho ** self.global_step)
|
lamb = self.lamb * (self.rho**self.global_step)
|
||||||
reg = prob.log().mul(P).sum(dim=-1).mean()
|
reg = prob.log().mul(P).sum(dim=-1).mean()
|
||||||
loss = loss - lamb * reg
|
loss = loss - lamb * reg
|
||||||
|
|
||||||
@@ -547,7 +547,7 @@ def evaluate(pred):
|
|||||||
score = pred.score
|
score = pred.score
|
||||||
label = pred.label
|
label = pred.label
|
||||||
diff = score - label
|
diff = score - label
|
||||||
MSE = (diff ** 2).mean()
|
MSE = (diff**2).mean()
|
||||||
MAE = (diff.abs()).mean()
|
MAE = (diff.abs()).mean()
|
||||||
IC = score.corr(label)
|
IC = score.corr(label)
|
||||||
return {"MSE": MSE, "MAE": MAE, "IC": IC}
|
return {"MSE": MSE, "MAE": MAE, "IC": IC}
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
pandas==1.1.2
|
pandas==1.1.2
|
||||||
numpy==1.17.4
|
numpy==1.21.0
|
||||||
scikit_learn==0.23.2
|
scikit_learn==0.23.2
|
||||||
torch==1.7.0
|
torch==1.7.0
|
||||||
@@ -1,3 +1,3 @@
|
|||||||
numpy==1.17.4
|
numpy==1.21.0
|
||||||
pandas==1.1.2
|
pandas==1.1.2
|
||||||
torch==1.2.0
|
torch==1.2.0
|
||||||
@@ -1,3 +1,3 @@
|
|||||||
numpy==1.17.4
|
numpy==1.21.0
|
||||||
pandas==1.1.2
|
pandas==1.1.2
|
||||||
xgboost==1.2.1
|
xgboost==1.2.1
|
||||||
@@ -4,16 +4,16 @@ This is the implementation of `DDG-DA` based on `Meta Controller` component prov
|
|||||||
Please refer to the paper for more details: *DDG-DA: Data Distribution Generation for Predictable Concept Drift Adaptation* [[arXiv](https://arxiv.org/abs/2201.04038)]
|
Please refer to the paper for more details: *DDG-DA: Data Distribution Generation for Predictable Concept Drift Adaptation* [[arXiv](https://arxiv.org/abs/2201.04038)]
|
||||||
|
|
||||||
|
|
||||||
## Background
|
# Background
|
||||||
In many real-world scenarios, we often deal with streaming data that is sequentially collected over time. Due to the non-stationary nature of the environment, the streaming data distribution may change in unpredictable ways, which is known as concept drift. To handle concept drift, previous methods first detect when/where the concept drift happens and then adapt models to fit the distribution of the latest data. However, there are still many cases that some underlying factors of environment evolution are predictable, making it possible to model the future concept drift trend of the streaming data, while such cases are not fully explored in previous work.
|
In many real-world scenarios, we often deal with streaming data that is sequentially collected over time. Due to the non-stationary nature of the environment, the streaming data distribution may change in unpredictable ways, which is known as concept drift. To handle concept drift, previous methods first detect when/where the concept drift happens and then adapt models to fit the distribution of the latest data. However, there are still many cases that some underlying factors of environment evolution are predictable, making it possible to model the future concept drift trend of the streaming data, while such cases are not fully explored in previous work.
|
||||||
|
|
||||||
Therefore, we propose a novel method `DDG-DA`, that can effectively forecast the evolution of data distribution and improve the performance of models. Specifically, we first train a predictor to estimate the future data distribution, then leverage it to generate training samples, and finally train models on the generated data.
|
Therefore, we propose a novel method `DDG-DA`, that can effectively forecast the evolution of data distribution and improve the performance of models. Specifically, we first train a predictor to estimate the future data distribution, then leverage it to generate training samples, and finally train models on the generated data.
|
||||||
|
|
||||||
## Dataset
|
# Dataset
|
||||||
The data in the paper are private. So we conduct experiments on Qlib's public dataset.
|
The data in the paper are private. So we conduct experiments on Qlib's public dataset.
|
||||||
Though the dataset is different, the conclusion remains the same. By applying `DDG-DA`, users can see rising trends at the test phase both in the proxy models' ICs and the performances of the forecasting models.
|
Though the dataset is different, the conclusion remains the same. By applying `DDG-DA`, users can see rising trends at the test phase both in the proxy models' ICs and the performances of the forecasting models.
|
||||||
|
|
||||||
## Run the Code
|
# Run the Code
|
||||||
Users can try `DDG-DA` by running the following command:
|
Users can try `DDG-DA` by running the following command:
|
||||||
```bash
|
```bash
|
||||||
python workflow.py run_all
|
python workflow.py run_all
|
||||||
@@ -24,7 +24,10 @@ The default forecasting models are `Linear`. Users can choose other forecasting
|
|||||||
python workflow.py --forecast_model="gbdt" run_all
|
python workflow.py --forecast_model="gbdt" run_all
|
||||||
```
|
```
|
||||||
|
|
||||||
|
# Results
|
||||||
## Results
|
|
||||||
|
|
||||||
The results of related methods in Qlib's public dataset can be found [here](../)
|
The results of related methods in Qlib's public dataset can be found [here](../)
|
||||||
|
|
||||||
|
# Requirements
|
||||||
|
Here is the minimal hardware requirements to run the ``workflow.py`` of DDG-DA.
|
||||||
|
* Memory: 45G
|
||||||
|
* Disk: 4G
|
||||||
|
|||||||
@@ -9,13 +9,10 @@ from qlib.data.dataset.handler import DataHandlerLP
|
|||||||
import pandas as pd
|
import pandas as pd
|
||||||
import fire
|
import fire
|
||||||
import sys
|
import sys
|
||||||
from tqdm.auto import tqdm
|
|
||||||
import yaml
|
|
||||||
import pickle
|
import pickle
|
||||||
from qlib import auto_init
|
from qlib import auto_init
|
||||||
from qlib.model.trainer import TrainerR, task_train
|
from qlib.model.trainer import TrainerR
|
||||||
from qlib.utils import init_instance_by_config
|
from qlib.utils import init_instance_by_config
|
||||||
from qlib.workflow.task.gen import RollingGen, task_generator
|
|
||||||
from qlib.workflow import R
|
from qlib.workflow import R
|
||||||
from qlib.tests.data import GetData
|
from qlib.tests.data import GetData
|
||||||
|
|
||||||
@@ -47,9 +44,10 @@ class DDGDA:
|
|||||||
rb = RollingBenchmark(model_type="gbdt")
|
rb = RollingBenchmark(model_type="gbdt")
|
||||||
task = rb.basic_task()
|
task = rb.basic_task()
|
||||||
|
|
||||||
model = init_instance_by_config(task["model"])
|
with R.start(experiment_name="feature_importance"):
|
||||||
dataset = init_instance_by_config(task["dataset"])
|
model = init_instance_by_config(task["model"])
|
||||||
model.fit(dataset)
|
dataset = init_instance_by_config(task["dataset"])
|
||||||
|
model.fit(dataset)
|
||||||
|
|
||||||
fi = model.get_feature_importance()
|
fi = model.get_feature_importance()
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
pandas==1.1.2
|
pandas==1.1.2
|
||||||
numpy==1.17.4
|
numpy==1.21.0
|
||||||
lightgbm==3.1.0
|
lightgbm==3.1.0
|
||||||
optuna==2.7.0
|
optuna==2.7.0
|
||||||
optuna-dashboard==0.4.1
|
optuna-dashboard==0.4.1
|
||||||
|
|||||||
@@ -2,7 +2,7 @@
|
|||||||
# Licensed under the MIT License.
|
# Licensed under the MIT License.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
This example is about how can simulate the OnlineManager based on rolling tasks.
|
This example is about how can simulate the OnlineManager based on rolling tasks.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from pprint import pprint
|
from pprint import pprint
|
||||||
@@ -15,6 +15,10 @@ from qlib.workflow.online.strategy import RollingStrategy
|
|||||||
from qlib.workflow.task.gen import RollingGen
|
from qlib.workflow.task.gen import RollingGen
|
||||||
from qlib.workflow.task.manage import TaskManager
|
from qlib.workflow.task.manage import TaskManager
|
||||||
from qlib.tests.config import CSI100_RECORD_LGB_TASK_CONFIG_ONLINE, CSI100_RECORD_XGBOOST_TASK_CONFIG_ONLINE
|
from qlib.tests.config import CSI100_RECORD_LGB_TASK_CONFIG_ONLINE, CSI100_RECORD_XGBOOST_TASK_CONFIG_ONLINE
|
||||||
|
import pandas as pd
|
||||||
|
from qlib.contrib.evaluate import backtest_daily
|
||||||
|
from qlib.contrib.evaluate import risk_analysis
|
||||||
|
from qlib.contrib.strategy import TopkDropoutStrategy
|
||||||
|
|
||||||
|
|
||||||
class OnlineSimulationExample:
|
class OnlineSimulationExample:
|
||||||
@@ -30,6 +34,7 @@ class OnlineSimulationExample:
|
|||||||
start_time="2018-09-10",
|
start_time="2018-09-10",
|
||||||
end_time="2018-10-31",
|
end_time="2018-10-31",
|
||||||
tasks=None,
|
tasks=None,
|
||||||
|
trainer="TrainerR",
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Init OnlineManagerExample.
|
Init OnlineManagerExample.
|
||||||
@@ -60,7 +65,13 @@ class OnlineSimulationExample:
|
|||||||
self.rolling_gen = RollingGen(
|
self.rolling_gen = RollingGen(
|
||||||
step=rolling_step, rtype=RollingGen.ROLL_SD, ds_extra_mod_func=None
|
step=rolling_step, rtype=RollingGen.ROLL_SD, ds_extra_mod_func=None
|
||||||
) # The rolling tasks generator, ds_extra_mod_func is None because we just need to simulate to 2018-10-31 and needn't change the handler end time.
|
) # The rolling tasks generator, ds_extra_mod_func is None because we just need to simulate to 2018-10-31 and needn't change the handler end time.
|
||||||
self.trainer = TrainerRM(self.exp_name, self.task_pool) # Also can be TrainerR, TrainerRM, DelayTrainerR
|
if trainer == "TrainerRM":
|
||||||
|
self.trainer = TrainerRM(self.exp_name, self.task_pool)
|
||||||
|
elif trainer == "TrainerR":
|
||||||
|
self.trainer = TrainerR(self.exp_name)
|
||||||
|
else:
|
||||||
|
# TODO: support all the trainers: TrainerR, TrainerRM, DelayTrainerR
|
||||||
|
raise NotImplementedError(f"This type of input is not supported")
|
||||||
self.rolling_online_manager = OnlineManager(
|
self.rolling_online_manager = OnlineManager(
|
||||||
RollingStrategy(exp_name, task_template=tasks, rolling_gen=self.rolling_gen),
|
RollingStrategy(exp_name, task_template=tasks, rolling_gen=self.rolling_gen),
|
||||||
trainer=self.trainer,
|
trainer=self.trainer,
|
||||||
@@ -70,7 +81,8 @@ class OnlineSimulationExample:
|
|||||||
|
|
||||||
# Reset all things to the first status, be careful to save important data
|
# Reset all things to the first status, be careful to save important data
|
||||||
def reset(self):
|
def reset(self):
|
||||||
TaskManager(self.task_pool).remove()
|
if isinstance(self.trainer, TrainerRM):
|
||||||
|
TaskManager(self.task_pool).remove()
|
||||||
exp = R.get_exp(experiment_name=self.exp_name)
|
exp = R.get_exp(experiment_name=self.exp_name)
|
||||||
for rid in exp.list_recorders():
|
for rid in exp.list_recorders():
|
||||||
exp.delete_recorder(rid)
|
exp.delete_recorder(rid)
|
||||||
@@ -84,7 +96,30 @@ class OnlineSimulationExample:
|
|||||||
print("========== collect results ==========")
|
print("========== collect results ==========")
|
||||||
print(self.rolling_online_manager.get_collector()())
|
print(self.rolling_online_manager.get_collector()())
|
||||||
print("========== signals ==========")
|
print("========== signals ==========")
|
||||||
print(self.rolling_online_manager.get_signals())
|
signals = self.rolling_online_manager.get_signals()
|
||||||
|
print(signals)
|
||||||
|
# Backtesting
|
||||||
|
# - the code is based on this example https://qlib.readthedocs.io/en/latest/component/strategy.html
|
||||||
|
CSI300_BENCH = "SH000903"
|
||||||
|
STRATEGY_CONFIG = {
|
||||||
|
"topk": 30,
|
||||||
|
"n_drop": 3,
|
||||||
|
"signal": signals.to_frame("score"),
|
||||||
|
}
|
||||||
|
strategy_obj = TopkDropoutStrategy(**STRATEGY_CONFIG)
|
||||||
|
report_normal, positions_normal = backtest_daily(
|
||||||
|
start_time=signals.index.get_level_values("datetime").min(),
|
||||||
|
end_time=signals.index.get_level_values("datetime").max(),
|
||||||
|
strategy=strategy_obj,
|
||||||
|
)
|
||||||
|
analysis = dict()
|
||||||
|
analysis["excess_return_without_cost"] = risk_analysis(report_normal["return"] - report_normal["bench"])
|
||||||
|
analysis["excess_return_with_cost"] = risk_analysis(
|
||||||
|
report_normal["return"] - report_normal["bench"] - report_normal["cost"]
|
||||||
|
)
|
||||||
|
|
||||||
|
analysis_df = pd.concat(analysis) # type: pd.DataFrame
|
||||||
|
pprint(analysis_df)
|
||||||
|
|
||||||
def worker(self):
|
def worker(self):
|
||||||
# train tasks by other progress or machines for multiprocessing
|
# train tasks by other progress or machines for multiprocessing
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ class TestClass(unittest.TestCase):
|
|||||||
provider_uri = "~/.qlib/qlib_data/yahoo_cn_1min"
|
provider_uri = "~/.qlib/qlib_data/yahoo_cn_1min"
|
||||||
qlib.init(
|
qlib.init(
|
||||||
provider_uri=provider_uri,
|
provider_uri=provider_uri,
|
||||||
mem_cache_size_limit=1024 ** 3 * 2,
|
mem_cache_size_limit=1024**3 * 2,
|
||||||
mem_cache_type="sizeof",
|
mem_cache_type="sizeof",
|
||||||
kernels=1,
|
kernels=1,
|
||||||
expression_provider={"class": "LocalExpressionProvider", "kwargs": {"time2idx": False}},
|
expression_provider={"class": "LocalExpressionProvider", "kwargs": {"time2idx": False}},
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ We use China stock market data for our example.
|
|||||||
unzip -d ~/.qlib/qlib_data/cn_data csi300_weight.zip
|
unzip -d ~/.qlib/qlib_data/cn_data csi300_weight.zip
|
||||||
rm -f csi300_weight.zip
|
rm -f csi300_weight.zip
|
||||||
```
|
```
|
||||||
|
NOTE: We don't find any public free resource to get the weight in the benchmark. To run the example, we manually create this weight data.
|
||||||
|
|
||||||
2. Prepare risk model data:
|
2. Prepare risk model data:
|
||||||
|
|
||||||
|
|||||||
1218
examples/tutorial/detailed_workflow.ipynb
Normal file
1218
examples/tutorial/detailed_workflow.ipynb
Normal file
File diff suppressed because it is too large
Load Diff
@@ -256,7 +256,6 @@
|
|||||||
"recorder = R.get_recorder(recorder_id=ba_rid, experiment_name=\"backtest_analysis\")\n",
|
"recorder = R.get_recorder(recorder_id=ba_rid, experiment_name=\"backtest_analysis\")\n",
|
||||||
"print(recorder)\n",
|
"print(recorder)\n",
|
||||||
"pred_df = recorder.load_object(\"pred.pkl\")\n",
|
"pred_df = recorder.load_object(\"pred.pkl\")\n",
|
||||||
"pred_df_dates = pred_df.index.get_level_values(level='datetime')\n",
|
|
||||||
"report_normal_df = recorder.load_object(\"portfolio_analysis/report_normal_1day.pkl\")\n",
|
"report_normal_df = recorder.load_object(\"portfolio_analysis/report_normal_1day.pkl\")\n",
|
||||||
"positions = recorder.load_object(\"portfolio_analysis/positions_normal_1day.pkl\")\n",
|
"positions = recorder.load_object(\"portfolio_analysis/positions_normal_1day.pkl\")\n",
|
||||||
"analysis_df = recorder.load_object(\"portfolio_analysis/port_analysis_1day.pkl\")"
|
"analysis_df = recorder.load_object(\"portfolio_analysis/port_analysis_1day.pkl\")"
|
||||||
|
|||||||
@@ -2,7 +2,7 @@
|
|||||||
# Licensed under the MIT License.
|
# Licensed under the MIT License.
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
__version__ = "0.8.2"
|
__version__ = "0.8.5.99"
|
||||||
__version__bak = __version__ # This version is backup for QlibConfig.reset_qlib_version
|
__version__bak = __version__ # This version is backup for QlibConfig.reset_qlib_version
|
||||||
import os
|
import os
|
||||||
from typing import Union
|
from typing import Union
|
||||||
@@ -12,6 +12,7 @@ import platform
|
|||||||
import subprocess
|
import subprocess
|
||||||
from .log import get_module_logger
|
from .log import get_module_logger
|
||||||
|
|
||||||
|
|
||||||
# init qlib
|
# init qlib
|
||||||
def init(default_conf="client", **kwargs):
|
def init(default_conf="client", **kwargs):
|
||||||
"""
|
"""
|
||||||
@@ -30,8 +31,8 @@ def init(default_conf="client", **kwargs):
|
|||||||
When using the recorder, skip_if_reg can set to True to avoid loss of recorder.
|
When using the recorder, skip_if_reg can set to True to avoid loss of recorder.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
from .config import C
|
from .config import C # pylint: disable=C0415
|
||||||
from .data.cache import H
|
from .data.cache import H # pylint: disable=C0415
|
||||||
|
|
||||||
# FIXME: this logger ignored the level in config
|
# FIXME: this logger ignored the level in config
|
||||||
logger = get_module_logger("Initialization", level=logging.INFO)
|
logger = get_module_logger("Initialization", level=logging.INFO)
|
||||||
@@ -85,7 +86,7 @@ def _mount_nfs_uri(provider_uri, mount_path, auto_mount: bool = False):
|
|||||||
mount_command = "sudo mount.nfs %s %s" % (provider_uri, mount_path)
|
mount_command = "sudo mount.nfs %s %s" % (provider_uri, mount_path)
|
||||||
# If the provider uri looks like this 172.23.233.89//data/csdesign'
|
# If the provider uri looks like this 172.23.233.89//data/csdesign'
|
||||||
# It will be a nfs path. The client provider will be used
|
# It will be a nfs path. The client provider will be used
|
||||||
if not auto_mount:
|
if not auto_mount: # pylint: disable=R1702
|
||||||
if not Path(mount_path).exists():
|
if not Path(mount_path).exists():
|
||||||
raise FileNotFoundError(
|
raise FileNotFoundError(
|
||||||
f"Invalid mount path: {mount_path}! Please mount manually: {mount_command} or Set init parameter `auto_mount=True`"
|
f"Invalid mount path: {mount_path}! Please mount manually: {mount_command} or Set init parameter `auto_mount=True`"
|
||||||
@@ -139,8 +140,10 @@ def _mount_nfs_uri(provider_uri, mount_path, auto_mount: bool = False):
|
|||||||
if not _is_mount:
|
if not _is_mount:
|
||||||
try:
|
try:
|
||||||
Path(mount_path).mkdir(parents=True, exist_ok=True)
|
Path(mount_path).mkdir(parents=True, exist_ok=True)
|
||||||
except Exception:
|
except Exception as e:
|
||||||
raise OSError(f"Failed to create directory {mount_path}, please create {mount_path} manually!")
|
raise OSError(
|
||||||
|
f"Failed to create directory {mount_path}, please create {mount_path} manually!"
|
||||||
|
) from e
|
||||||
|
|
||||||
# check nfs-common
|
# check nfs-common
|
||||||
command_res = os.popen("dpkg -l | grep nfs-common")
|
command_res = os.popen("dpkg -l | grep nfs-common")
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
# Copyright (c) Microsoft Corporation.
|
# Copyright (c) Microsoft Corporation.
|
||||||
# Licensed under the MIT License.
|
# Licensed under the MIT License.
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
import copy
|
import copy
|
||||||
from typing import List, Tuple, Union, TYPE_CHECKING
|
from typing import List, Tuple, Union, TYPE_CHECKING
|
||||||
@@ -171,8 +172,8 @@ def get_strategy_executor(
|
|||||||
# NOTE:
|
# NOTE:
|
||||||
# - for avoiding recursive import
|
# - for avoiding recursive import
|
||||||
# - typing annotations is not reliable
|
# - typing annotations is not reliable
|
||||||
from ..strategy.base import BaseStrategy
|
from ..strategy.base import BaseStrategy # pylint: disable=C0415
|
||||||
from .executor import BaseExecutor
|
from .executor import BaseExecutor # pylint: disable=C0415
|
||||||
|
|
||||||
trade_account = create_account_instance(
|
trade_account = create_account_instance(
|
||||||
start_time=start_time, end_time=end_time, benchmark=benchmark, account=account, pos_type=pos_type
|
start_time=start_time, end_time=end_time, benchmark=benchmark, account=account, pos_type=pos_type
|
||||||
@@ -323,3 +324,6 @@ def format_decisions(
|
|||||||
last_dec_idx = i
|
last_dec_idx = i
|
||||||
res[1].append((decisions[last_dec_idx], format_decisions(decisions[last_dec_idx + 1 :])))
|
res[1].append((decisions[last_dec_idx], format_decisions(decisions[last_dec_idx + 1 :])))
|
||||||
return res
|
return res
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["Order"]
|
||||||
|
|||||||
@@ -2,11 +2,11 @@
|
|||||||
# Licensed under the MIT License.
|
# Licensed under the MIT License.
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
import copy
|
import copy
|
||||||
from typing import Dict, List, Tuple, TYPE_CHECKING
|
from typing import Dict, List, Tuple
|
||||||
from qlib.utils import init_instance_by_config
|
from qlib.utils import init_instance_by_config
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from .position import BasePosition, InfPosition, Position
|
from .position import BasePosition
|
||||||
from .report import PortfolioMetrics, Indicator
|
from .report import PortfolioMetrics, Indicator
|
||||||
from .decision import BaseTradeDecision, Order
|
from .decision import BaseTradeDecision, Order
|
||||||
from .exchange import Exchange
|
from .exchange import Exchange
|
||||||
|
|||||||
@@ -7,19 +7,18 @@ from qlib.data.data import Cal
|
|||||||
from qlib.utils.time import concat_date_time, epsilon_change
|
from qlib.utils.time import concat_date_time, epsilon_change
|
||||||
from qlib.log import get_module_logger
|
from qlib.log import get_module_logger
|
||||||
|
|
||||||
|
from typing import ClassVar, Optional, Union, List, Tuple
|
||||||
|
|
||||||
# try to fix circular imports when enabling type hints
|
# try to fix circular imports when enabling type hints
|
||||||
from typing import Callable, TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from qlib.strategy.base import BaseStrategy
|
from qlib.strategy.base import BaseStrategy
|
||||||
from qlib.backtest.exchange import Exchange
|
from qlib.backtest.exchange import Exchange
|
||||||
from qlib.backtest.utils import TradeCalendarManager
|
from qlib.backtest.utils import TradeCalendarManager
|
||||||
import warnings
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import numpy as np
|
from dataclasses import dataclass
|
||||||
from dataclasses import dataclass, field
|
|
||||||
from typing import ClassVar, Optional, Union, List, Set, Tuple
|
|
||||||
|
|
||||||
|
|
||||||
class OrderDir(IntEnum):
|
class OrderDir(IntEnum):
|
||||||
@@ -418,7 +417,7 @@ class BaseTradeDecision:
|
|||||||
return kwargs["default_value"]
|
return kwargs["default_value"]
|
||||||
else:
|
else:
|
||||||
# Default to get full index
|
# Default to get full index
|
||||||
raise NotImplementedError(f"The decision didn't provide an index range")
|
raise NotImplementedError(f"The decision didn't provide an index range") from NotImplementedError
|
||||||
|
|
||||||
# clip index
|
# clip index
|
||||||
if getattr(self, "total_step", None) is not None:
|
if getattr(self, "total_step", None) is not None:
|
||||||
|
|||||||
@@ -3,13 +3,13 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
from typing import List, Tuple, Union
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from .account import Account
|
from .account import Account
|
||||||
|
|
||||||
from qlib.backtest.position import BasePosition, Position
|
from qlib.backtest.position import BasePosition, Position
|
||||||
import random
|
import random
|
||||||
from typing import List, Tuple, Union
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
@@ -18,7 +18,7 @@ from ..config import C
|
|||||||
from ..constant import REG_CN
|
from ..constant import REG_CN
|
||||||
from ..log import get_module_logger
|
from ..log import get_module_logger
|
||||||
from .decision import Order, OrderDir, OrderHelper
|
from .decision import Order, OrderDir, OrderHelper
|
||||||
from .high_performance_ds import BaseQuote, PandasQuote, NumpyQuote
|
from .high_performance_ds import BaseQuote, NumpyQuote
|
||||||
|
|
||||||
|
|
||||||
class Exchange:
|
class Exchange:
|
||||||
|
|||||||
@@ -1,22 +1,18 @@
|
|||||||
from abc import abstractclassmethod, abstractmethod
|
from abc import abstractmethod
|
||||||
import copy
|
import copy
|
||||||
from qlib.backtest.position import BasePosition
|
from qlib.backtest.position import BasePosition
|
||||||
from qlib.log import get_module_logger
|
from qlib.log import get_module_logger
|
||||||
from types import GeneratorType
|
from types import GeneratorType
|
||||||
from qlib.backtest.account import Account
|
from qlib.backtest.account import Account
|
||||||
import warnings
|
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from typing import List, Tuple, Union
|
from typing import List, Tuple, Union
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
|
|
||||||
from qlib.backtest.report import Indicator
|
from .decision import Order, BaseTradeDecision
|
||||||
|
|
||||||
from .decision import EmptyTradeDecision, Order, BaseTradeDecision
|
|
||||||
from .exchange import Exchange
|
from .exchange import Exchange
|
||||||
from .utils import TradeCalendarManager, CommonInfrastructure, LevelInfrastructure, get_start_end_idx
|
from .utils import TradeCalendarManager, CommonInfrastructure, LevelInfrastructure, get_start_end_idx
|
||||||
|
|
||||||
from ..utils import init_instance_by_config
|
from ..utils import init_instance_by_config
|
||||||
from ..utils.time import Freq
|
|
||||||
from ..strategy.base import BaseStrategy
|
from ..strategy.base import BaseStrategy
|
||||||
|
|
||||||
|
|
||||||
@@ -193,7 +189,8 @@ class BaseExecutor:
|
|||||||
pass
|
pass
|
||||||
return return_value.get("execute_result")
|
return return_value.get("execute_result")
|
||||||
|
|
||||||
@abstractclassmethod
|
@classmethod
|
||||||
|
@abstractmethod
|
||||||
def _collect_data(cls, trade_decision: BaseTradeDecision, level: int = 0) -> Tuple[List[object], dict]:
|
def _collect_data(cls, trade_decision: BaseTradeDecision, level: int = 0) -> Tuple[List[object], dict]:
|
||||||
"""
|
"""
|
||||||
Please refer to the doc of collect_data
|
Please refer to the doc of collect_data
|
||||||
@@ -245,7 +242,7 @@ class BaseExecutor:
|
|||||||
if self.track_data:
|
if self.track_data:
|
||||||
yield trade_decision
|
yield trade_decision
|
||||||
|
|
||||||
atomic = not issubclass(self.__class__, NestedExecutor) # issubclass(A, A) is True
|
atomic = not issubclass(self.__class__, NestedExecutor) # issubclass(A, A) is True
|
||||||
|
|
||||||
if atomic and trade_decision.get_range_limit(default_value=None) is not None:
|
if atomic and trade_decision.get_range_limit(default_value=None) is not None:
|
||||||
raise ValueError("atomic executor doesn't support specify `range_limit`")
|
raise ValueError("atomic executor doesn't support specify `range_limit`")
|
||||||
@@ -453,7 +450,6 @@ class NestedExecutor(BaseExecutor):
|
|||||||
inner_exe_res :
|
inner_exe_res :
|
||||||
the execution result of inner task
|
the execution result of inner task
|
||||||
"""
|
"""
|
||||||
pass
|
|
||||||
|
|
||||||
def get_all_executors(self):
|
def get_all_executors(self):
|
||||||
"""get all executors, including self and inner_executor.get_all_executors()"""
|
"""get all executors, including self and inner_executor.get_all_executors()"""
|
||||||
|
|||||||
@@ -2,8 +2,6 @@
|
|||||||
# Licensed under the MIT License.
|
# Licensed under the MIT License.
|
||||||
|
|
||||||
|
|
||||||
import copy
|
|
||||||
import pathlib
|
|
||||||
from typing import Dict, List, Union
|
from typing import Dict, List, Union
|
||||||
|
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
@@ -362,7 +360,9 @@ class Position(BasePosition):
|
|||||||
# check if to delete
|
# check if to delete
|
||||||
if self.position[stock_id]["amount"] < -1e-5:
|
if self.position[stock_id]["amount"] < -1e-5:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"only have {} {}, require {}".format(self.position[stock_id]["amount"], stock_id, trade_amount)
|
"only have {} {}, require {}".format(
|
||||||
|
self.position[stock_id]["amount"] + trade_amount, stock_id, trade_amount
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
new_cash = trade_val - cost
|
new_cash = trade_val - cost
|
||||||
@@ -538,7 +538,7 @@ class InfPosition(BasePosition):
|
|||||||
def get_stock_amount_dict(self) -> Dict:
|
def get_stock_amount_dict(self) -> Dict:
|
||||||
raise NotImplementedError(f"InfPosition doesn't support get_stock_amount_dict")
|
raise NotImplementedError(f"InfPosition doesn't support get_stock_amount_dict")
|
||||||
|
|
||||||
def get_stock_weight_dict(self, only_stock: bool) -> Dict:
|
def get_stock_weight_dict(self, only_stock: bool = False) -> Dict:
|
||||||
raise NotImplementedError(f"InfPosition doesn't support get_stock_weight_dict")
|
raise NotImplementedError(f"InfPosition doesn't support get_stock_weight_dict")
|
||||||
|
|
||||||
def add_count_all(self, bar):
|
def add_count_all(self, bar):
|
||||||
|
|||||||
@@ -10,11 +10,8 @@ import numpy as np
|
|||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from qlib.backtest.exchange import Exchange
|
from qlib.backtest.exchange import Exchange
|
||||||
from .decision import IdxTradeRange
|
|
||||||
from qlib.backtest.decision import BaseTradeDecision, Order, OrderDir
|
from qlib.backtest.decision import BaseTradeDecision, Order, OrderDir
|
||||||
from qlib.backtest.utils import TradeCalendarManager
|
from .high_performance_ds import BaseOrderIndicator, NumpyOrderIndicator, SingleMetric
|
||||||
from .high_performance_ds import BaseOrderIndicator, PandasOrderIndicator, NumpyOrderIndicator, SingleMetric
|
|
||||||
from ..data import D
|
|
||||||
from ..tests.config import CSI300_BENCH
|
from ..tests.config import CSI300_BENCH
|
||||||
from ..utils.resam import get_higher_eq_freq_feature, resam_ts_data
|
from ..utils.resam import get_higher_eq_freq_feature, resam_ts_data
|
||||||
import qlib.utils.index_data as idd
|
import qlib.utils.index_data as idd
|
||||||
|
|||||||
@@ -28,7 +28,6 @@ class Signal(metaclass=abc.ABCMeta):
|
|||||||
Union[pd.Series, pd.DataFrame, None]:
|
Union[pd.Series, pd.DataFrame, None]:
|
||||||
returns None if no signal in the specific day
|
returns None if no signal in the specific day
|
||||||
"""
|
"""
|
||||||
...
|
|
||||||
|
|
||||||
|
|
||||||
class SignalWCache(Signal):
|
class SignalWCache(Signal):
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ from pathlib import Path
|
|||||||
from typing import Callable, Optional, Union
|
from typing import Callable, Optional, Union
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from qlib.constant import REG_CN, REG_US
|
from qlib.constant import REG_CN, REG_US, REG_TW
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from qlib.utils.time import Freq
|
from qlib.utils.time import Freq
|
||||||
@@ -92,6 +92,7 @@ _default_config = {
|
|||||||
"calendar_provider": "LocalCalendarProvider",
|
"calendar_provider": "LocalCalendarProvider",
|
||||||
"instrument_provider": "LocalInstrumentProvider",
|
"instrument_provider": "LocalInstrumentProvider",
|
||||||
"feature_provider": "LocalFeatureProvider",
|
"feature_provider": "LocalFeatureProvider",
|
||||||
|
"pit_provider": "LocalPITProvider",
|
||||||
"expression_provider": "LocalExpressionProvider",
|
"expression_provider": "LocalExpressionProvider",
|
||||||
"dataset_provider": "LocalDatasetProvider",
|
"dataset_provider": "LocalDatasetProvider",
|
||||||
"provider": "LocalProvider",
|
"provider": "LocalProvider",
|
||||||
@@ -108,7 +109,6 @@ _default_config = {
|
|||||||
"provider_uri": "",
|
"provider_uri": "",
|
||||||
# cache
|
# cache
|
||||||
"expression_cache": None,
|
"expression_cache": None,
|
||||||
"dataset_cache": None,
|
|
||||||
"calendar_cache": None,
|
"calendar_cache": None,
|
||||||
# for simple dataset cache
|
# for simple dataset cache
|
||||||
"local_cache_path": None,
|
"local_cache_path": None,
|
||||||
@@ -171,6 +171,18 @@ _default_config = {
|
|||||||
"default_exp_name": "Experiment",
|
"default_exp_name": "Experiment",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
|
"pit_record_type": {
|
||||||
|
"date": "I", # uint32
|
||||||
|
"period": "I", # uint32
|
||||||
|
"value": "d", # float64
|
||||||
|
"index": "I", # uint32
|
||||||
|
},
|
||||||
|
"pit_record_nan": {
|
||||||
|
"date": 0,
|
||||||
|
"period": 0,
|
||||||
|
"value": float("NAN"),
|
||||||
|
"index": 0xFFFFFFFF,
|
||||||
|
},
|
||||||
# Default config for MongoDB
|
# Default config for MongoDB
|
||||||
"mongo": {
|
"mongo": {
|
||||||
"task_url": "mongodb://localhost:27017/",
|
"task_url": "mongodb://localhost:27017/",
|
||||||
@@ -184,20 +196,12 @@ _default_config = {
|
|||||||
|
|
||||||
MODE_CONF = {
|
MODE_CONF = {
|
||||||
"server": {
|
"server": {
|
||||||
# data provider config
|
|
||||||
"calendar_provider": "LocalCalendarProvider",
|
|
||||||
"instrument_provider": "LocalInstrumentProvider",
|
|
||||||
"feature_provider": "LocalFeatureProvider",
|
|
||||||
"expression_provider": "LocalExpressionProvider",
|
|
||||||
"dataset_provider": "LocalDatasetProvider",
|
|
||||||
"provider": "LocalProvider",
|
|
||||||
# config it in qlib.init()
|
# config it in qlib.init()
|
||||||
"provider_uri": "",
|
"provider_uri": "",
|
||||||
# redis
|
# redis
|
||||||
"redis_host": "127.0.0.1",
|
"redis_host": "127.0.0.1",
|
||||||
"redis_port": 6379,
|
"redis_port": 6379,
|
||||||
"redis_task_db": 1,
|
"redis_task_db": 1,
|
||||||
"kernels": NUM_USABLE_CPU,
|
|
||||||
# cache
|
# cache
|
||||||
"expression_cache": DISK_EXPRESSION_CACHE,
|
"expression_cache": DISK_EXPRESSION_CACHE,
|
||||||
"dataset_cache": DISK_DATASET_CACHE,
|
"dataset_cache": DISK_DATASET_CACHE,
|
||||||
@@ -205,25 +209,15 @@ MODE_CONF = {
|
|||||||
"mount_path": None,
|
"mount_path": None,
|
||||||
},
|
},
|
||||||
"client": {
|
"client": {
|
||||||
# data provider config
|
|
||||||
"calendar_provider": "LocalCalendarProvider",
|
|
||||||
"instrument_provider": "LocalInstrumentProvider",
|
|
||||||
"feature_provider": "LocalFeatureProvider",
|
|
||||||
"expression_provider": "LocalExpressionProvider",
|
|
||||||
"dataset_provider": "LocalDatasetProvider",
|
|
||||||
"provider": "LocalProvider",
|
|
||||||
# config it in user's own code
|
# config it in user's own code
|
||||||
"provider_uri": "~/.qlib/qlib_data/cn_data",
|
"provider_uri": "~/.qlib/qlib_data/cn_data",
|
||||||
# cache
|
# cache
|
||||||
# Using parameter 'remote' to announce the client is using server_cache, and the writing access will be disabled.
|
# Using parameter 'remote' to announce the client is using server_cache, and the writing access will be disabled.
|
||||||
# Disable cache by default. Avoid introduce advanced features for beginners
|
# Disable cache by default. Avoid introduce advanced features for beginners
|
||||||
"expression_cache": None,
|
|
||||||
"dataset_cache": None,
|
"dataset_cache": None,
|
||||||
# SimpleDatasetCache directory
|
# SimpleDatasetCache directory
|
||||||
"local_cache_path": Path("~/.cache/qlib_simple_cache").expanduser().resolve(),
|
"local_cache_path": Path("~/.cache/qlib_simple_cache").expanduser().resolve(),
|
||||||
"calendar_cache": None,
|
|
||||||
# client config
|
# client config
|
||||||
"kernels": NUM_USABLE_CPU,
|
|
||||||
"mount_path": None,
|
"mount_path": None,
|
||||||
"auto_mount": False, # The nfs is already mounted on our server[auto_mount: False].
|
"auto_mount": False, # The nfs is already mounted on our server[auto_mount: False].
|
||||||
# The nfs should be auto-mounted by qlib on other
|
# The nfs should be auto-mounted by qlib on other
|
||||||
@@ -257,6 +251,11 @@ _default_region_config = {
|
|||||||
"limit_threshold": None,
|
"limit_threshold": None,
|
||||||
"deal_price": "close",
|
"deal_price": "close",
|
||||||
},
|
},
|
||||||
|
REG_TW: {
|
||||||
|
"trade_unit": 1000,
|
||||||
|
"limit_threshold": 0.1,
|
||||||
|
"deal_price": "close",
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -388,13 +387,11 @@ class QlibConfig(Config):
|
|||||||
default_conf : str
|
default_conf : str
|
||||||
the default config template chosen by user: "server", "client"
|
the default config template chosen by user: "server", "client"
|
||||||
"""
|
"""
|
||||||
from .utils import set_log_with_config, get_module_logger, can_use_cache
|
from .utils import set_log_with_config, get_module_logger, can_use_cache # pylint: disable=C0415
|
||||||
|
|
||||||
self.reset()
|
self.reset()
|
||||||
|
|
||||||
_logging_config = self.logging_config
|
_logging_config = kwargs.get("logging_config", self.logging_config)
|
||||||
if "logging_config" in kwargs:
|
|
||||||
_logging_config = kwargs["logging_config"]
|
|
||||||
|
|
||||||
# set global config
|
# set global config
|
||||||
if _logging_config:
|
if _logging_config:
|
||||||
@@ -433,11 +430,11 @@ class QlibConfig(Config):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def register(self):
|
def register(self):
|
||||||
from .utils import init_instance_by_config
|
from .utils import init_instance_by_config # pylint: disable=C0415
|
||||||
from .data.ops import register_all_ops
|
from .data.ops import register_all_ops # pylint: disable=C0415
|
||||||
from .data.data import register_all_wrappers
|
from .data.data import register_all_wrappers # pylint: disable=C0415
|
||||||
from .workflow import R, QlibRecorder
|
from .workflow import R, QlibRecorder # pylint: disable=C0415
|
||||||
from .workflow.utils import experiment_exit_handler
|
from .workflow.utils import experiment_exit_handler # pylint: disable=C0415
|
||||||
|
|
||||||
register_all_ops(self)
|
register_all_ops(self)
|
||||||
register_all_wrappers(self)
|
register_all_wrappers(self)
|
||||||
@@ -454,7 +451,7 @@ class QlibConfig(Config):
|
|||||||
self._registered = True
|
self._registered = True
|
||||||
|
|
||||||
def reset_qlib_version(self):
|
def reset_qlib_version(self):
|
||||||
import qlib
|
import qlib # pylint: disable=C0415
|
||||||
|
|
||||||
reset_version = self.get("qlib_reset_version", None)
|
reset_version = self.get("qlib_reset_version", None)
|
||||||
if reset_version is not None:
|
if reset_version is not None:
|
||||||
|
|||||||
@@ -4,6 +4,10 @@
|
|||||||
# REGION CONST
|
# REGION CONST
|
||||||
REG_CN = "cn"
|
REG_CN = "cn"
|
||||||
REG_US = "us"
|
REG_US = "us"
|
||||||
|
REG_TW = "tw"
|
||||||
|
|
||||||
# Epsilon for avoiding division by zero.
|
# Epsilon for avoiding division by zero.
|
||||||
EPS = 1e-12
|
EPS = 1e-12
|
||||||
|
|
||||||
|
# Infinity in integer
|
||||||
|
INF = 10**18
|
||||||
|
|||||||
@@ -7,8 +7,7 @@ import warnings
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from qlib.utils import init_instance_by_config
|
from qlib.data.dataset import DatasetH
|
||||||
from qlib.data.dataset import DatasetH, DataHandler
|
|
||||||
|
|
||||||
|
|
||||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
@@ -16,7 +15,7 @@ device = "cuda" if torch.cuda.is_available() else "cpu"
|
|||||||
|
|
||||||
def _to_tensor(x):
|
def _to_tensor(x):
|
||||||
if not isinstance(x, torch.Tensor):
|
if not isinstance(x, torch.Tensor):
|
||||||
return torch.tensor(x, dtype=torch.float, device=device)
|
return torch.tensor(x, dtype=torch.float, device=device) # pylint: disable=E1101
|
||||||
return x
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -5,9 +5,7 @@ from ...data.dataset.handler import DataHandlerLP
|
|||||||
from ...data.dataset.processor import Processor
|
from ...data.dataset.processor import Processor
|
||||||
from ...utils import get_callable_kwargs
|
from ...utils import get_callable_kwargs
|
||||||
from ...data.dataset import processor as processor_module
|
from ...data.dataset import processor as processor_module
|
||||||
from ...log import TimeInspector
|
|
||||||
from inspect import getfullargspec
|
from inspect import getfullargspec
|
||||||
import copy
|
|
||||||
|
|
||||||
|
|
||||||
def check_transform_proc(proc_l, fit_start_time, fit_end_time):
|
def check_transform_proc(proc_l, fit_start_time, fit_end_time):
|
||||||
|
|||||||
164
qlib/contrib/data/highfreq_handler.py
Normal file
164
qlib/contrib/data/highfreq_handler.py
Normal file
@@ -0,0 +1,164 @@
|
|||||||
|
from qlib.data.dataset.handler import DataHandler, DataHandlerLP
|
||||||
|
|
||||||
|
EPSILON = 1e-4
|
||||||
|
|
||||||
|
|
||||||
|
class HighFreqHandler(DataHandlerLP):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
instruments="csi300",
|
||||||
|
start_time=None,
|
||||||
|
end_time=None,
|
||||||
|
infer_processors=[],
|
||||||
|
learn_processors=[],
|
||||||
|
fit_start_time=None,
|
||||||
|
fit_end_time=None,
|
||||||
|
drop_raw=True,
|
||||||
|
):
|
||||||
|
def check_transform_proc(proc_l):
|
||||||
|
new_l = []
|
||||||
|
for p in proc_l:
|
||||||
|
p["kwargs"].update(
|
||||||
|
{
|
||||||
|
"fit_start_time": fit_start_time,
|
||||||
|
"fit_end_time": fit_end_time,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
new_l.append(p)
|
||||||
|
return new_l
|
||||||
|
|
||||||
|
infer_processors = check_transform_proc(infer_processors)
|
||||||
|
learn_processors = check_transform_proc(learn_processors)
|
||||||
|
|
||||||
|
data_loader = {
|
||||||
|
"class": "QlibDataLoader",
|
||||||
|
"kwargs": {
|
||||||
|
"config": self.get_feature_config(),
|
||||||
|
"swap_level": False,
|
||||||
|
"freq": "1min",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
super().__init__(
|
||||||
|
instruments=instruments,
|
||||||
|
start_time=start_time,
|
||||||
|
end_time=end_time,
|
||||||
|
data_loader=data_loader,
|
||||||
|
infer_processors=infer_processors,
|
||||||
|
learn_processors=learn_processors,
|
||||||
|
drop_raw=drop_raw,
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_feature_config(self):
|
||||||
|
fields = []
|
||||||
|
names = []
|
||||||
|
|
||||||
|
template_if = "If(IsNull({1}), {0}, {1})"
|
||||||
|
template_paused = "Select(Gt($hx_paused_num, 1.001), {0})"
|
||||||
|
|
||||||
|
def get_normalized_price_feature(price_field, shift=0):
|
||||||
|
# norm with the close price of 237th minute of yesterday.
|
||||||
|
if shift == 0:
|
||||||
|
template_norm = "{0}/DayLast(Ref({1}, 243))"
|
||||||
|
else:
|
||||||
|
template_norm = "Ref({0}, " + str(shift) + ")/DayLast(Ref({1}, 243))"
|
||||||
|
|
||||||
|
template_fillnan = "FFillNan({0})"
|
||||||
|
# calculate -> ffill -> remove paused
|
||||||
|
feature_ops = template_paused.format(
|
||||||
|
template_fillnan.format(
|
||||||
|
template_norm.format(template_if.format("$close", price_field), template_fillnan.format("$close"))
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return feature_ops
|
||||||
|
|
||||||
|
fields += [get_normalized_price_feature("$open", 0)]
|
||||||
|
fields += [get_normalized_price_feature("$high", 0)]
|
||||||
|
fields += [get_normalized_price_feature("$low", 0)]
|
||||||
|
fields += [get_normalized_price_feature("$close", 0)]
|
||||||
|
fields += [get_normalized_price_feature("$vwap", 0)]
|
||||||
|
names += ["$open", "$high", "$low", "$close", "$vwap"]
|
||||||
|
|
||||||
|
fields += [get_normalized_price_feature("$open", 240)]
|
||||||
|
fields += [get_normalized_price_feature("$high", 240)]
|
||||||
|
fields += [get_normalized_price_feature("$low", 240)]
|
||||||
|
fields += [get_normalized_price_feature("$close", 240)]
|
||||||
|
fields += [get_normalized_price_feature("$vwap", 240)]
|
||||||
|
names += ["$open_1", "$high_1", "$low_1", "$close_1", "$vwap_1"]
|
||||||
|
|
||||||
|
# calculate and fill nan with 0
|
||||||
|
template_gzero = "If(Ge({0}, 0), {0}, 0)"
|
||||||
|
fields += [
|
||||||
|
template_gzero.format(
|
||||||
|
template_paused.format(
|
||||||
|
"If(IsNull({0}), 0, {0})".format("{0}/Ref(DayLast(Mean({0}, 7200)), 240)".format("$volume"))
|
||||||
|
)
|
||||||
|
)
|
||||||
|
]
|
||||||
|
names += ["$volume"]
|
||||||
|
|
||||||
|
fields += [
|
||||||
|
template_gzero.format(
|
||||||
|
template_paused.format(
|
||||||
|
"If(IsNull({0}), 0, {0})".format(
|
||||||
|
"Ref({0}, 240)/Ref(DayLast(Mean({0}, 7200)), 240)".format("$volume")
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
]
|
||||||
|
names += ["$volume_1"]
|
||||||
|
|
||||||
|
return fields, names
|
||||||
|
|
||||||
|
|
||||||
|
class HighFreqBacktestHandler(DataHandler):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
instruments="csi300",
|
||||||
|
start_time=None,
|
||||||
|
end_time=None,
|
||||||
|
):
|
||||||
|
data_loader = {
|
||||||
|
"class": "QlibDataLoader",
|
||||||
|
"kwargs": {
|
||||||
|
"config": self.get_feature_config(),
|
||||||
|
"swap_level": False,
|
||||||
|
"freq": "1min",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
super().__init__(
|
||||||
|
instruments=instruments,
|
||||||
|
start_time=start_time,
|
||||||
|
end_time=end_time,
|
||||||
|
data_loader=data_loader,
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_feature_config(self):
|
||||||
|
fields = []
|
||||||
|
names = []
|
||||||
|
|
||||||
|
template_if = "If(IsNull({1}), {0}, {1})"
|
||||||
|
template_paused = "Select(Gt($hx_paused_num, 1.001), {0})"
|
||||||
|
# template_paused = "{0}"
|
||||||
|
template_fillnan = "FFillNan({0})"
|
||||||
|
fields += [
|
||||||
|
template_fillnan.format(template_paused.format("$close")),
|
||||||
|
]
|
||||||
|
names += ["$close0"]
|
||||||
|
|
||||||
|
fields += [
|
||||||
|
template_paused.format(
|
||||||
|
template_if.format(
|
||||||
|
template_fillnan.format("$close"),
|
||||||
|
"$vwap",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
]
|
||||||
|
names += ["$vwap0"]
|
||||||
|
|
||||||
|
fields += [template_paused.format("If(IsNull({0}), 0, {0})".format("$volume"))]
|
||||||
|
names += ["$volume0"]
|
||||||
|
|
||||||
|
fields += [template_paused.format("If(IsNull({0}), 0, {0})".format("$factor"))]
|
||||||
|
names += ["$factor0"]
|
||||||
|
|
||||||
|
return fields, names
|
||||||
81
qlib/contrib/data/highfreq_processor.py
Normal file
81
qlib/contrib/data/highfreq_processor.py
Normal file
@@ -0,0 +1,81 @@
|
|||||||
|
import os
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pandas as pd
|
||||||
|
from qlib.data.dataset.processor import Processor
|
||||||
|
from qlib.data.dataset.utils import fetch_df_by_index
|
||||||
|
from typing import Dict
|
||||||
|
|
||||||
|
|
||||||
|
class HighFreqTrans(Processor):
|
||||||
|
def __init__(self, dtype: str = "bool"):
|
||||||
|
self.dtype = dtype
|
||||||
|
|
||||||
|
def fit(self, df_features):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def __call__(self, df_features):
|
||||||
|
if self.dtype == "bool":
|
||||||
|
return df_features.astype(np.int8)
|
||||||
|
else:
|
||||||
|
return df_features.astype(np.float32)
|
||||||
|
|
||||||
|
|
||||||
|
class HighFreqNorm(Processor):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
fit_start_time: pd.Timestamp,
|
||||||
|
fit_end_time: pd.Timestamp,
|
||||||
|
feature_save_dir: str,
|
||||||
|
norm_groups: Dict[str, int],
|
||||||
|
):
|
||||||
|
|
||||||
|
self.fit_start_time = fit_start_time
|
||||||
|
self.fit_end_time = fit_end_time
|
||||||
|
self.feature_save_dir = feature_save_dir
|
||||||
|
self.norm_groups = norm_groups
|
||||||
|
|
||||||
|
def fit(self, df_features) -> None:
|
||||||
|
if os.path.exists(self.feature_save_dir) and len(os.listdir(self.feature_save_dir)) != 0:
|
||||||
|
return
|
||||||
|
os.makedirs(self.feature_save_dir)
|
||||||
|
fetch_df = fetch_df_by_index(df_features, slice(self.fit_start_time, self.fit_end_time), level="datetime")
|
||||||
|
del df_features
|
||||||
|
index = 0
|
||||||
|
names = {}
|
||||||
|
for name, dim in self.norm_groups.items():
|
||||||
|
names[name] = slice(index, index + dim)
|
||||||
|
index += dim
|
||||||
|
for name, name_val in names.items():
|
||||||
|
df_values = fetch_df.iloc(axis=1)[name_val].values
|
||||||
|
if name.endswith("volume"):
|
||||||
|
df_values = np.log1p(df_values)
|
||||||
|
self.feature_mean = np.nanmean(df_values)
|
||||||
|
np.save(self.feature_save_dir + name + "_mean.npy", self.feature_mean)
|
||||||
|
df_values = df_values - self.feature_mean
|
||||||
|
self.feature_std = np.nanstd(np.absolute(df_values))
|
||||||
|
np.save(self.feature_save_dir + name + "_std.npy", self.feature_std)
|
||||||
|
df_values = df_values / self.feature_std
|
||||||
|
np.save(self.feature_save_dir + name + "_vmax.npy", np.nanmax(df_values))
|
||||||
|
np.save(self.feature_save_dir + name + "_vmin.npy", np.nanmin(df_values))
|
||||||
|
return
|
||||||
|
|
||||||
|
def __call__(self, df_features):
|
||||||
|
if "date" in df_features:
|
||||||
|
df_features.droplevel("date", inplace=True)
|
||||||
|
df_values = df_features.values
|
||||||
|
index = 0
|
||||||
|
names = {}
|
||||||
|
for name, dim in self.norm_groups.items():
|
||||||
|
names[name] = slice(index, index + dim)
|
||||||
|
index += dim
|
||||||
|
for name, name_val in names.items():
|
||||||
|
feature_mean = np.load(self.feature_save_dir + name + "_mean.npy")
|
||||||
|
feature_std = np.load(self.feature_save_dir + name + "_std.npy")
|
||||||
|
|
||||||
|
if name.endswith("volume"):
|
||||||
|
df_values[:, name_val] = np.log1p(df_values[:, name_val])
|
||||||
|
df_values[:, name_val] -= feature_mean
|
||||||
|
df_values[:, name_val] /= feature_std
|
||||||
|
df_features = pd.DataFrame(data=df_values, index=df_features.index, columns=df_features.columns)
|
||||||
|
return df_features.fillna(0)
|
||||||
301
qlib/contrib/data/highfreq_provider.py
Normal file
301
qlib/contrib/data/highfreq_provider.py
Normal file
@@ -0,0 +1,301 @@
|
|||||||
|
import os
|
||||||
|
import time
|
||||||
|
import datetime
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import qlib
|
||||||
|
from qlib.data import D
|
||||||
|
from qlib.config import REG_CN
|
||||||
|
from qlib.utils import init_instance_by_config
|
||||||
|
from qlib.data.dataset.handler import DataHandlerLP
|
||||||
|
from qlib.data.data import Cal
|
||||||
|
from qlib.contrib.ops.high_freq import get_calendar_day, DayLast, FFillNan, BFillNan, Date, Select, IsNull, IsInf, Cut
|
||||||
|
import pickle as pkl
|
||||||
|
from joblib import Parallel, delayed
|
||||||
|
from utilsd.logging import print_log
|
||||||
|
|
||||||
|
|
||||||
|
class HighFreqProvider:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
start_time: str,
|
||||||
|
end_time: str,
|
||||||
|
train_end_time: str,
|
||||||
|
valid_start_time: str,
|
||||||
|
valid_end_time: str,
|
||||||
|
test_start_time: str,
|
||||||
|
qlib_conf: dict,
|
||||||
|
feature_conf: dict,
|
||||||
|
label_conf: Optional[dict] = None,
|
||||||
|
backtest_conf: dict = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> None:
|
||||||
|
self.start_time = start_time
|
||||||
|
self.end_time = end_time
|
||||||
|
self.test_start_time = test_start_time
|
||||||
|
self.train_end_time = train_end_time
|
||||||
|
self.valid_start_time = valid_start_time
|
||||||
|
self.valid_end_time = valid_end_time
|
||||||
|
self._init_qlib(qlib_conf)
|
||||||
|
self.feature_conf = feature_conf
|
||||||
|
self.label_conf = label_conf
|
||||||
|
self.backtest_conf = backtest_conf
|
||||||
|
self.qlib_conf = qlib_conf
|
||||||
|
|
||||||
|
def get_pre_datasets(self):
|
||||||
|
"""Generate the training, validation and test datasets for prediction
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple[BaseDataset, BaseDataset, BaseDataset]: The training and test datasets
|
||||||
|
"""
|
||||||
|
|
||||||
|
dict_feature_path = self.feature_conf["path"]
|
||||||
|
train_feature_path = dict_feature_path[:-4] + "_train.pkl"
|
||||||
|
valid_feature_path = dict_feature_path[:-4] + "_valid.pkl"
|
||||||
|
test_feature_path = dict_feature_path[:-4] + "_test.pkl"
|
||||||
|
|
||||||
|
dict_label_path = self.label_conf["path"]
|
||||||
|
train_label_path = dict_label_path[:-4] + "_train.pkl"
|
||||||
|
valid_label_path = dict_label_path[:-4] + "_valid.pkl"
|
||||||
|
test_label_path = dict_label_path[:-4] + "_test.pkl"
|
||||||
|
|
||||||
|
if (
|
||||||
|
not os.path.isfile(train_feature_path)
|
||||||
|
or not os.path.isfile(valid_feature_path)
|
||||||
|
or not os.path.isfile(test_feature_path)
|
||||||
|
):
|
||||||
|
xtrain, xvalid, xtest = self._gen_data(self.feature_conf)
|
||||||
|
xtrain.to_pickle(train_feature_path)
|
||||||
|
xvalid.to_pickle(valid_feature_path)
|
||||||
|
xtest.to_pickle(test_feature_path)
|
||||||
|
del xtrain, xvalid, xtest
|
||||||
|
|
||||||
|
if (
|
||||||
|
not os.path.isfile(train_label_path)
|
||||||
|
or not os.path.isfile(valid_label_path)
|
||||||
|
or not os.path.isfile(test_label_path)
|
||||||
|
):
|
||||||
|
ytrain, yvalid, ytest = self._gen_data(self.label_conf)
|
||||||
|
ytrain.to_pickle(train_label_path)
|
||||||
|
yvalid.to_pickle(valid_label_path)
|
||||||
|
ytest.to_pickle(test_label_path)
|
||||||
|
del ytrain, yvalid, ytest
|
||||||
|
|
||||||
|
feature = {
|
||||||
|
"train": train_feature_path,
|
||||||
|
"valid": valid_feature_path,
|
||||||
|
"test": test_feature_path,
|
||||||
|
}
|
||||||
|
|
||||||
|
label = {
|
||||||
|
"train": train_label_path,
|
||||||
|
"valid": valid_label_path,
|
||||||
|
"test": test_label_path,
|
||||||
|
}
|
||||||
|
|
||||||
|
return feature, label
|
||||||
|
|
||||||
|
def get_backtest(self, **kwargs) -> None:
|
||||||
|
self._gen_data(self.backtest_conf)
|
||||||
|
|
||||||
|
def _init_qlib(self, qlib_conf):
|
||||||
|
"""initialize qlib"""
|
||||||
|
|
||||||
|
qlib.init(
|
||||||
|
region=REG_CN,
|
||||||
|
auto_mount=False,
|
||||||
|
custom_ops=[DayLast, FFillNan, BFillNan, Date, Select, IsNull, IsInf, Cut],
|
||||||
|
expression_cache=None,
|
||||||
|
**qlib_conf,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _prepare_calender_cache(self):
|
||||||
|
"""preload the calendar for cache"""
|
||||||
|
|
||||||
|
# This code used the copy-on-write feature of Linux
|
||||||
|
# to avoid calculating the calendar multiple times in the subprocess.
|
||||||
|
# This code may accelerate, but may be not useful on Windows and Mac Os
|
||||||
|
Cal.calendar(freq="1min")
|
||||||
|
get_calendar_day(freq="1min")
|
||||||
|
|
||||||
|
def _gen_dataframe(self, config, datasets=["train", "valid", "test"]):
|
||||||
|
try:
|
||||||
|
path = config.pop("path")
|
||||||
|
except KeyError as e:
|
||||||
|
raise ValueError("Must specify the path to save the dataset.") from e
|
||||||
|
if os.path.isfile(path):
|
||||||
|
start = time.time()
|
||||||
|
print_log("Dataset exists, load from disk.", __name__)
|
||||||
|
|
||||||
|
# res = dataset.prepare(['train', 'valid', 'test'])
|
||||||
|
with open(path, "rb") as f:
|
||||||
|
data = pkl.load(f)
|
||||||
|
if isinstance(data, dict):
|
||||||
|
res = [data[i] for i in datasets]
|
||||||
|
else:
|
||||||
|
res = data.prepare(datasets)
|
||||||
|
print_log(f"Data loaded, time cost: {time.time() - start:.2f}", __name__)
|
||||||
|
else:
|
||||||
|
if not os.path.exists(os.path.dirname(path)):
|
||||||
|
os.makedirs(os.path.dirname(path))
|
||||||
|
print_log("Generating dataset", __name__)
|
||||||
|
start_time = time.time()
|
||||||
|
self._prepare_calender_cache()
|
||||||
|
dataset = init_instance_by_config(config)
|
||||||
|
trainset, validset, testset = dataset.prepare(["train", "valid", "test"])
|
||||||
|
data = {
|
||||||
|
"train": trainset,
|
||||||
|
"valid": validset,
|
||||||
|
"test": testset,
|
||||||
|
}
|
||||||
|
with open(path, "wb") as f:
|
||||||
|
pkl.dump(data, f)
|
||||||
|
with open(path[:-4] + "train.pkl", "wb") as f:
|
||||||
|
pkl.dump(trainset, f)
|
||||||
|
with open(path[:-4] + "valid.pkl", "wb") as f:
|
||||||
|
pkl.dump(validset, f)
|
||||||
|
with open(path[:-4] + "test.pkl", "wb") as f:
|
||||||
|
pkl.dump(testset, f)
|
||||||
|
res = [data[i] for i in datasets]
|
||||||
|
print_log(f"Data generated, time cost: {(time.time() - start_time):.2f}", __name__)
|
||||||
|
return res
|
||||||
|
|
||||||
|
def _gen_data(self, config, datasets=["train", "valid", "test"]):
|
||||||
|
try:
|
||||||
|
path = config.pop("path")
|
||||||
|
except KeyError as e:
|
||||||
|
raise ValueError("Must specify the path to save the dataset.") from e
|
||||||
|
if os.path.isfile(path):
|
||||||
|
start = time.time()
|
||||||
|
print_log("Dataset exists, load from disk.", __name__)
|
||||||
|
|
||||||
|
# res = dataset.prepare(['train', 'valid', 'test'])
|
||||||
|
with open(path, "rb") as f:
|
||||||
|
data = pkl.load(f)
|
||||||
|
if isinstance(data, dict):
|
||||||
|
res = [data[i] for i in datasets]
|
||||||
|
else:
|
||||||
|
res = data.prepare(datasets)
|
||||||
|
print_log(f"Data loaded, time cost: {time.time() - start:.2f}", __name__)
|
||||||
|
else:
|
||||||
|
if not os.path.exists(os.path.dirname(path)):
|
||||||
|
os.makedirs(os.path.dirname(path))
|
||||||
|
print_log("Generating dataset", __name__)
|
||||||
|
start_time = time.time()
|
||||||
|
self._prepare_calender_cache()
|
||||||
|
dataset = init_instance_by_config(config)
|
||||||
|
dataset.config(dump_all=True, recursive=True)
|
||||||
|
dataset.to_pickle(path)
|
||||||
|
res = dataset.prepare(datasets)
|
||||||
|
print_log(f"Data generated, time cost: {(time.time() - start_time):.2f}", __name__)
|
||||||
|
return res
|
||||||
|
|
||||||
|
def _gen_dataset(self, config):
|
||||||
|
try:
|
||||||
|
path = config.pop("path")
|
||||||
|
except KeyError as e:
|
||||||
|
raise ValueError("Must specify the path to save the dataset.") from e
|
||||||
|
if os.path.isfile(path):
|
||||||
|
start = time.time()
|
||||||
|
print_log("Dataset exists, load from disk.", __name__)
|
||||||
|
|
||||||
|
with open(path, "rb") as f:
|
||||||
|
dataset = pkl.load(f)
|
||||||
|
print_log(f"Data loaded, time cost: {time.time() - start:.2f}", __name__)
|
||||||
|
else:
|
||||||
|
start = time.time()
|
||||||
|
if not os.path.exists(os.path.dirname(path)):
|
||||||
|
os.makedirs(os.path.dirname(path))
|
||||||
|
print_log("Generating dataset", __name__)
|
||||||
|
self._prepare_calender_cache()
|
||||||
|
dataset = init_instance_by_config(config)
|
||||||
|
print_log(f"Dataset init, time cost: {time.time() - start:.2f}", __name__)
|
||||||
|
dataset.prepare(["train", "valid", "test"])
|
||||||
|
print_log(f"Dataset prepared, time cost: {time.time() - start:.2f}", __name__)
|
||||||
|
dataset.config(dump_all=True, recursive=True)
|
||||||
|
dataset.to_pickle(path)
|
||||||
|
return dataset
|
||||||
|
|
||||||
|
def _gen_day_dataset(self, config, conf_type):
|
||||||
|
try:
|
||||||
|
path = config.pop("path")
|
||||||
|
except KeyError as e:
|
||||||
|
raise ValueError("Must specify the path to save the dataset.") from e
|
||||||
|
|
||||||
|
if os.path.isfile(path + "tmp_dataset.pkl"):
|
||||||
|
start = time.time()
|
||||||
|
print_log("Dataset exists, load from disk.", __name__)
|
||||||
|
else:
|
||||||
|
start = time.time()
|
||||||
|
if not os.path.exists(os.path.dirname(path)):
|
||||||
|
os.makedirs(os.path.dirname(path))
|
||||||
|
print_log("Generating dataset", __name__)
|
||||||
|
self._prepare_calender_cache()
|
||||||
|
dataset = init_instance_by_config(config)
|
||||||
|
print_log(f"Dataset init, time cost: {time.time() - start:.2f}", __name__)
|
||||||
|
dataset.config(dump_all=False, recursive=True)
|
||||||
|
dataset.to_pickle(path + "tmp_dataset.pkl")
|
||||||
|
|
||||||
|
with open(path + "tmp_dataset.pkl", "rb") as f:
|
||||||
|
new_dataset = pkl.load(f)
|
||||||
|
|
||||||
|
time_list = D.calendar(start_time=self.start_time, end_time=self.end_time, freq="1min")[::240]
|
||||||
|
|
||||||
|
def generate_dataset(times):
|
||||||
|
if os.path.isfile(path + times.strftime("%Y-%m-%d") + ".pkl"):
|
||||||
|
print("exist " + times.strftime("%Y-%m-%d"))
|
||||||
|
return
|
||||||
|
self._init_qlib(self.qlib_conf)
|
||||||
|
end_times = times + datetime.timedelta(days=1)
|
||||||
|
new_dataset.handler.config(**{"start_time": times, "end_time": end_times})
|
||||||
|
if conf_type == "backtest":
|
||||||
|
new_dataset.handler.setup_data()
|
||||||
|
else:
|
||||||
|
new_dataset.handler.setup_data(init_type=DataHandlerLP.IT_LS)
|
||||||
|
new_dataset.config(dump_all=True, recursive=True)
|
||||||
|
new_dataset.to_pickle(path + times.strftime("%Y-%m-%d") + ".pkl")
|
||||||
|
|
||||||
|
Parallel(n_jobs=8)(delayed(generate_dataset)(times) for times in time_list)
|
||||||
|
|
||||||
|
def _gen_stock_dataset(self, config, conf_type):
|
||||||
|
try:
|
||||||
|
path = config.pop("path")
|
||||||
|
except KeyError as e:
|
||||||
|
raise ValueError("Must specify the path to save the dataset.") from e
|
||||||
|
|
||||||
|
if os.path.isfile(path + "tmp_dataset.pkl"):
|
||||||
|
start = time.time()
|
||||||
|
print_log("Dataset exists, load from disk.", __name__)
|
||||||
|
else:
|
||||||
|
start = time.time()
|
||||||
|
if not os.path.exists(os.path.dirname(path)):
|
||||||
|
os.makedirs(os.path.dirname(path))
|
||||||
|
print_log("Generating dataset", __name__)
|
||||||
|
self._prepare_calender_cache()
|
||||||
|
dataset = init_instance_by_config(config)
|
||||||
|
print_log(f"Dataset init, time cost: {time.time() - start:.2f}", __name__)
|
||||||
|
dataset.config(dump_all=False, recursive=True)
|
||||||
|
dataset.to_pickle(path + "tmp_dataset.pkl")
|
||||||
|
|
||||||
|
with open(path + "tmp_dataset.pkl", "rb") as f:
|
||||||
|
new_dataset = pkl.load(f)
|
||||||
|
|
||||||
|
instruments = D.instruments(market="all")
|
||||||
|
stock_list = D.list_instruments(
|
||||||
|
instruments=instruments, start_time=self.start_time, end_time=self.end_time, freq="1min", as_list=True
|
||||||
|
)
|
||||||
|
|
||||||
|
def generate_dataset(stock):
|
||||||
|
if os.path.isfile(path + stock + ".pkl"):
|
||||||
|
print("exist " + stock)
|
||||||
|
return
|
||||||
|
self._init_qlib(self.qlib_conf)
|
||||||
|
new_dataset.handler.config(**{"instruments": [stock]})
|
||||||
|
if conf_type == "backtest":
|
||||||
|
new_dataset.handler.setup_data()
|
||||||
|
else:
|
||||||
|
new_dataset.handler.setup_data(init_type=DataHandlerLP.IT_LS)
|
||||||
|
new_dataset.config(dump_all=True, recursive=True)
|
||||||
|
new_dataset.to_pickle(path + stock + ".pkl")
|
||||||
|
|
||||||
|
Parallel(n_jobs=32)(delayed(generate_dataset)(stock) for stock in stock_list)
|
||||||
@@ -1,9 +1,6 @@
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
|
||||||
import copy
|
|
||||||
|
|
||||||
from ...log import TimeInspector
|
from ...log import TimeInspector
|
||||||
from ...utils.serial import Serializable
|
|
||||||
from ...data.dataset.processor import Processor, get_group_columns
|
from ...data.dataset.processor import Processor, get_group_columns
|
||||||
|
|
||||||
|
|
||||||
@@ -62,10 +59,10 @@ class ConfigSectionProcessor(Processor):
|
|||||||
|
|
||||||
# Features
|
# Features
|
||||||
cols = df_focus.columns[df_focus.columns.str.contains("^KLEN|^KLOW|^KUP")]
|
cols = df_focus.columns[df_focus.columns.str.contains("^KLEN|^KLOW|^KUP")]
|
||||||
df_focus[cols] = df_focus[cols].apply(lambda x: x ** 0.25).groupby(level="datetime").apply(_feature_norm)
|
df_focus[cols] = df_focus[cols].apply(lambda x: x**0.25).groupby(level="datetime").apply(_feature_norm)
|
||||||
|
|
||||||
cols = df_focus.columns[df_focus.columns.str.contains("^KLOW2|^KUP2")]
|
cols = df_focus.columns[df_focus.columns.str.contains("^KLOW2|^KUP2")]
|
||||||
df_focus[cols] = df_focus[cols].apply(lambda x: x ** 0.5).groupby(level="datetime").apply(_feature_norm)
|
df_focus[cols] = df_focus[cols].apply(lambda x: x**0.5).groupby(level="datetime").apply(_feature_norm)
|
||||||
|
|
||||||
_cols = [
|
_cols = [
|
||||||
"KMID",
|
"KMID",
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
# Copyright (c) Microsoft Corporation.
|
# Copyright (c) Microsoft Corporation.
|
||||||
# Licensed under the MIT License.
|
# Licensed under the MIT License.
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from typing import Dict, Iterable
|
from typing import Dict, Iterable, Union
|
||||||
|
|
||||||
|
|
||||||
def align_index(df_dict, join):
|
def align_index(df_dict, join):
|
||||||
@@ -24,6 +24,10 @@ class SepDataFrame:
|
|||||||
SepDataFrame tries to act like a DataFrame whose column with multiindex
|
SepDataFrame tries to act like a DataFrame whose column with multiindex
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
# TODO:
|
||||||
|
# SepDataFrame try to behave like pandas dataframe, but it is still not them same
|
||||||
|
# Contributions are welcome to make it more complete.
|
||||||
|
|
||||||
def __init__(self, df_dict: Dict[str, pd.DataFrame], join: str, skip_align=False):
|
def __init__(self, df_dict: Dict[str, pd.DataFrame], join: str, skip_align=False):
|
||||||
"""
|
"""
|
||||||
initialize the data based on the dataframe dictionary
|
initialize the data based on the dataframe dictionary
|
||||||
@@ -77,14 +81,37 @@ class SepDataFrame:
|
|||||||
|
|
||||||
def _update_join(self):
|
def _update_join(self):
|
||||||
if self.join not in self:
|
if self.join not in self:
|
||||||
self.join = next(iter(self._df_dict.keys()))
|
if len(self._df_dict) > 0:
|
||||||
|
self.join = next(iter(self._df_dict.keys()))
|
||||||
|
else:
|
||||||
|
# NOTE: this will change the behavior of previous reindex when all the keys are empty
|
||||||
|
self.join = None
|
||||||
|
|
||||||
def __getitem__(self, item):
|
def __getitem__(self, item):
|
||||||
|
# TODO: behave more like pandas when multiindex
|
||||||
return self._df_dict[item]
|
return self._df_dict[item]
|
||||||
|
|
||||||
def __setitem__(self, item: str, df: pd.DataFrame):
|
def __setitem__(self, item: str, df: Union[pd.DataFrame, pd.Series]):
|
||||||
# TODO: consider the join behavior
|
# TODO: consider the join behavior
|
||||||
self._df_dict[item] = df
|
if not isinstance(item, tuple):
|
||||||
|
self._df_dict[item] = df
|
||||||
|
else:
|
||||||
|
# NOTE: corner case of MultiIndex
|
||||||
|
_df_dict_key, *col_name = item
|
||||||
|
col_name = tuple(col_name)
|
||||||
|
if _df_dict_key in self._df_dict:
|
||||||
|
if len(col_name) == 1:
|
||||||
|
col_name = col_name[0]
|
||||||
|
self._df_dict[_df_dict_key][col_name] = df
|
||||||
|
else:
|
||||||
|
if isinstance(df, pd.Series):
|
||||||
|
if len(col_name) == 1:
|
||||||
|
col_name = col_name[0]
|
||||||
|
self._df_dict[_df_dict_key] = df.to_frame(col_name)
|
||||||
|
else:
|
||||||
|
df_copy = df.copy() # avoid changing df
|
||||||
|
df_copy.columns = pd.MultiIndex.from_tuples([(*col_name, *idx) for idx in df.columns.to_list()])
|
||||||
|
self._df_dict[_df_dict_key] = df_copy
|
||||||
|
|
||||||
def __delitem__(self, item: str):
|
def __delitem__(self, item: str):
|
||||||
del self._df_dict[item]
|
del self._df_dict[item]
|
||||||
@@ -164,14 +191,14 @@ import builtins
|
|||||||
|
|
||||||
|
|
||||||
def _isinstance(instance, cls):
|
def _isinstance(instance, cls):
|
||||||
if isinstance_orig(instance, SepDataFrame): # pylint: disable=E0602
|
if isinstance_orig(instance, SepDataFrame): # pylint: disable=E0602 # noqa: F821
|
||||||
if isinstance(cls, Iterable):
|
if isinstance(cls, Iterable):
|
||||||
for c in cls:
|
for c in cls:
|
||||||
if c is pd.DataFrame:
|
if c is pd.DataFrame:
|
||||||
return True
|
return True
|
||||||
elif cls is pd.DataFrame:
|
elif cls is pd.DataFrame:
|
||||||
return True
|
return True
|
||||||
return isinstance_orig(instance, cls) # pylint: disable=E0602
|
return isinstance_orig(instance, cls) # pylint: disable=E0602 # noqa: F821
|
||||||
|
|
||||||
|
|
||||||
builtins.isinstance_orig = builtins.isinstance
|
builtins.isinstance_orig = builtins.isinstance
|
||||||
|
|||||||
@@ -4,8 +4,10 @@ Here is a batch of evaluation functions.
|
|||||||
The interface should be redesigned carefully in the future.
|
The interface should be redesigned carefully in the future.
|
||||||
"""
|
"""
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from typing import Tuple
|
from typing import Tuple
|
||||||
|
from qlib import get_module_logger
|
||||||
|
from qlib.utils.paral import complex_parallel, DelayedDict
|
||||||
|
from joblib import Parallel, delayed
|
||||||
|
|
||||||
|
|
||||||
def calc_long_short_prec(
|
def calc_long_short_prec(
|
||||||
@@ -61,32 +63,6 @@ def calc_long_short_prec(
|
|||||||
return (l_dom.groupby(date_col).sum() / l_c), (s_dom.groupby(date_col).sum() / s_c)
|
return (l_dom.groupby(date_col).sum() / l_c), (s_dom.groupby(date_col).sum() / s_c)
|
||||||
|
|
||||||
|
|
||||||
def calc_ic(pred: pd.Series, label: pd.Series, date_col="datetime", dropna=False) -> Tuple[pd.Series, pd.Series]:
|
|
||||||
"""calc_ic.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
pred :
|
|
||||||
pred
|
|
||||||
label :
|
|
||||||
label
|
|
||||||
date_col :
|
|
||||||
date_col
|
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
(pd.Series, pd.Series)
|
|
||||||
ic and rank ic
|
|
||||||
"""
|
|
||||||
df = pd.DataFrame({"pred": pred, "label": label})
|
|
||||||
ic = df.groupby(date_col).apply(lambda df: df["pred"].corr(df["label"]))
|
|
||||||
ric = df.groupby(date_col).apply(lambda df: df["pred"].corr(df["label"], method="spearman"))
|
|
||||||
if dropna:
|
|
||||||
return ic.dropna(), ric.dropna()
|
|
||||||
else:
|
|
||||||
return ic, ric
|
|
||||||
|
|
||||||
|
|
||||||
def calc_long_short_return(
|
def calc_long_short_return(
|
||||||
pred: pd.Series,
|
pred: pd.Series,
|
||||||
label: pd.Series,
|
label: pd.Series,
|
||||||
@@ -127,3 +103,105 @@ def calc_long_short_return(
|
|||||||
r_short = group.apply(lambda x: x.nsmallest(N(x), columns="pred").label.mean())
|
r_short = group.apply(lambda x: x.nsmallest(N(x), columns="pred").label.mean())
|
||||||
r_avg = group.label.mean()
|
r_avg = group.label.mean()
|
||||||
return (r_long - r_short) / 2, r_avg
|
return (r_long - r_short) / 2, r_avg
|
||||||
|
|
||||||
|
|
||||||
|
def pred_autocorr(pred: pd.Series, lag=1, inst_col="instrument", date_col="datetime"):
|
||||||
|
"""pred_autocorr.
|
||||||
|
|
||||||
|
Limitation:
|
||||||
|
- If the datetime is not sequential densely, the correlation will be calulated based on adjacent dates. (some users may expected NaN)
|
||||||
|
|
||||||
|
:param pred: pd.Series with following format
|
||||||
|
instrument datetime
|
||||||
|
SH600000 2016-01-04 -0.000403
|
||||||
|
2016-01-05 -0.000753
|
||||||
|
2016-01-06 -0.021801
|
||||||
|
2016-01-07 -0.065230
|
||||||
|
2016-01-08 -0.062465
|
||||||
|
:type pred: pd.Series
|
||||||
|
:param lag:
|
||||||
|
"""
|
||||||
|
if isinstance(pred, pd.DataFrame):
|
||||||
|
pred = pred.iloc[:, 0]
|
||||||
|
get_module_logger("pred_autocorr").warning(f"Only the first column in {pred.columns} of `pred` is kept")
|
||||||
|
pred_ustk = pred.sort_index().unstack(inst_col)
|
||||||
|
corr_s = {}
|
||||||
|
for (idx, cur), (_, prev) in zip(pred_ustk.iterrows(), pred_ustk.shift(lag).iterrows()):
|
||||||
|
corr_s[idx] = cur.corr(prev)
|
||||||
|
corr_s = pd.Series(corr_s).sort_index()
|
||||||
|
return corr_s
|
||||||
|
|
||||||
|
|
||||||
|
def pred_autocorr_all(pred_dict, n_jobs=-1, **kwargs):
|
||||||
|
"""
|
||||||
|
calculate auto correlation for pred_dict
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
pred_dict : dict
|
||||||
|
A dict like {<method_name>: <prediction>}
|
||||||
|
kwargs :
|
||||||
|
all these arguments will be passed into pred_autocorr
|
||||||
|
"""
|
||||||
|
ac_dict = {}
|
||||||
|
for k, pred in pred_dict.items():
|
||||||
|
ac_dict[k] = delayed(pred_autocorr)(pred, **kwargs)
|
||||||
|
return complex_parallel(Parallel(n_jobs=n_jobs, verbose=10), ac_dict)
|
||||||
|
|
||||||
|
|
||||||
|
def calc_ic(pred: pd.Series, label: pd.Series, date_col="datetime", dropna=False) -> (pd.Series, pd.Series):
|
||||||
|
"""calc_ic.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
pred :
|
||||||
|
pred
|
||||||
|
label :
|
||||||
|
label
|
||||||
|
date_col :
|
||||||
|
date_col
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
(pd.Series, pd.Series)
|
||||||
|
ic and rank ic
|
||||||
|
"""
|
||||||
|
df = pd.DataFrame({"pred": pred, "label": label})
|
||||||
|
ic = df.groupby(date_col).apply(lambda df: df["pred"].corr(df["label"]))
|
||||||
|
ric = df.groupby(date_col).apply(lambda df: df["pred"].corr(df["label"], method="spearman"))
|
||||||
|
if dropna:
|
||||||
|
return ic.dropna(), ric.dropna()
|
||||||
|
else:
|
||||||
|
return ic, ric
|
||||||
|
|
||||||
|
|
||||||
|
def calc_all_ic(pred_dict_all, label, date_col="datetime", dropna=False, n_jobs=-1):
|
||||||
|
"""calc_all_ic.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
pred_dict_all :
|
||||||
|
A dict like {<method_name>: <prediction>}
|
||||||
|
label:
|
||||||
|
A pd.Series of label values
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
{'Q2+IND_z': {'ic': <ic series like>
|
||||||
|
2016-01-04 -0.057407
|
||||||
|
...
|
||||||
|
2020-05-28 0.183470
|
||||||
|
2020-05-29 0.171393
|
||||||
|
'ric': <rank ic series like>
|
||||||
|
2016-01-04 -0.040888
|
||||||
|
...
|
||||||
|
2020-05-28 0.236665
|
||||||
|
2020-05-29 0.183886
|
||||||
|
}
|
||||||
|
...}
|
||||||
|
"""
|
||||||
|
pred_all_ics = {}
|
||||||
|
for k, pred in pred_dict_all.items():
|
||||||
|
pred_all_ics[k] = DelayedDict(["ic", "ric"], delayed(calc_ic)(pred, label, date_col=date_col, dropna=dropna))
|
||||||
|
pred_all_ics = complex_parallel(Parallel(n_jobs=n_jobs, verbose=10), pred_all_ics)
|
||||||
|
return pred_all_ics
|
||||||
|
|||||||
@@ -5,12 +5,10 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import copy
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from scipy.stats import spearmanr, pearsonr
|
from scipy.stats import spearmanr, pearsonr
|
||||||
|
|
||||||
|
|
||||||
from ..data import D
|
from ..data import D
|
||||||
|
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
@@ -243,4 +241,4 @@ def get_rank_ic(a, b):
|
|||||||
|
|
||||||
|
|
||||||
def get_normal_ic(a, b):
|
def get_normal_ic(a, b):
|
||||||
return pearsonr(a, b).correlation
|
return pearsonr(a, b)[0]
|
||||||
|
|||||||
@@ -2,3 +2,6 @@
|
|||||||
# Licensed under the MIT License.
|
# Licensed under the MIT License.
|
||||||
|
|
||||||
from .data_selection import MetaTaskDS, MetaDatasetDS, MetaModelDS
|
from .data_selection import MetaTaskDS, MetaDatasetDS, MetaModelDS
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["MetaTaskDS", "MetaDatasetDS", "MetaModelDS"]
|
||||||
|
|||||||
@@ -3,3 +3,6 @@
|
|||||||
|
|
||||||
from .dataset import MetaDatasetDS, MetaTaskDS
|
from .dataset import MetaDatasetDS, MetaTaskDS
|
||||||
from .model import MetaModelDS
|
from .model import MetaModelDS
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["MetaDatasetDS", "MetaTaskDS", "MetaModelDS"]
|
||||||
|
|||||||
@@ -1,24 +1,23 @@
|
|||||||
# Copyright (c) Microsoft Corporation.
|
# Copyright (c) Microsoft Corporation.
|
||||||
# Licensed under the MIT License.
|
# Licensed under the MIT License.
|
||||||
from copy import deepcopy
|
|
||||||
from qlib.data.dataset.utils import init_task_handler
|
|
||||||
from qlib.utils.data import deepcopy_basic_type
|
|
||||||
from qlib.contrib.torch import data_to_tensor
|
|
||||||
from qlib.workflow.task.utils import TimeAdjuster
|
|
||||||
from qlib.model.meta.task import MetaTask
|
|
||||||
from typing import Dict, List, Union, Text, Tuple
|
|
||||||
from qlib.data.dataset.handler import DataHandler
|
|
||||||
from qlib.log import get_module_logger
|
|
||||||
from qlib.utils import auto_filter_kwargs, get_date_by_shift, init_instance_by_config
|
|
||||||
from qlib.workflow import R
|
|
||||||
from qlib.workflow.task.gen import RollingGen, task_generator
|
|
||||||
from joblib import Parallel, delayed
|
|
||||||
from qlib.model.meta.dataset import MetaTaskDataset
|
|
||||||
from qlib.model.trainer import task_train, TrainerR
|
|
||||||
from qlib.data.dataset import DatasetH
|
|
||||||
from tqdm.auto import tqdm
|
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
from copy import deepcopy
|
||||||
|
from joblib import Parallel, delayed # pylint: disable=E0401
|
||||||
|
from typing import Dict, List, Union, Text, Tuple
|
||||||
|
from qlib.data.dataset.utils import init_task_handler
|
||||||
|
from qlib.data.dataset import DatasetH
|
||||||
|
from qlib.contrib.torch import data_to_tensor
|
||||||
|
from qlib.model.meta.task import MetaTask
|
||||||
|
from qlib.model.meta.dataset import MetaTaskDataset
|
||||||
|
from qlib.model.trainer import TrainerR
|
||||||
|
from qlib.log import get_module_logger
|
||||||
|
from qlib.utils import auto_filter_kwargs, get_date_by_shift, init_instance_by_config
|
||||||
|
from qlib.utils.data import deepcopy_basic_type
|
||||||
|
from qlib.workflow import R
|
||||||
|
from qlib.workflow.task.gen import RollingGen, task_generator
|
||||||
|
from qlib.workflow.task.utils import TimeAdjuster
|
||||||
|
from tqdm.auto import tqdm
|
||||||
|
|
||||||
|
|
||||||
class InternalData:
|
class InternalData:
|
||||||
|
|||||||
@@ -1,28 +1,25 @@
|
|||||||
# Copyright (c) Microsoft Corporation.
|
# Copyright (c) Microsoft Corporation.
|
||||||
# Licensed under the MIT License.
|
# Licensed under the MIT License.
|
||||||
|
|
||||||
from qlib.log import get_module_logger
|
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from qlib.model.meta.task import MetaTask
|
|
||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
from torch import optim
|
from torch import optim
|
||||||
from tqdm.auto import tqdm
|
from tqdm.auto import tqdm
|
||||||
import collections
|
|
||||||
import copy
|
import copy
|
||||||
from typing import Union, List, Tuple, Dict
|
from typing import Union, List
|
||||||
|
|
||||||
from ....data.dataset.weight import Reweighter
|
|
||||||
from ....model.meta.dataset import MetaTaskDataset
|
from ....model.meta.dataset import MetaTaskDataset
|
||||||
from ....model.meta.model import MetaModel, MetaTaskModel
|
from ....model.meta.model import MetaTaskModel
|
||||||
from ....workflow import R
|
from ....workflow import R
|
||||||
|
|
||||||
from .utils import ICLoss
|
from .utils import ICLoss
|
||||||
from .dataset import MetaDatasetDS
|
from .dataset import MetaDatasetDS
|
||||||
from qlib.contrib.meta.data_selection.net import PredNet
|
|
||||||
from qlib.data.dataset.weight import Reweighter
|
|
||||||
from qlib.log import get_module_logger
|
from qlib.log import get_module_logger
|
||||||
|
from qlib.model.meta.task import MetaTask
|
||||||
|
from qlib.data.dataset.weight import Reweighter
|
||||||
|
from qlib.contrib.meta.data_selection.net import PredNet
|
||||||
|
|
||||||
logger = get_module_logger("data selection")
|
logger = get_module_logger("data selection")
|
||||||
|
|
||||||
@@ -100,7 +97,6 @@ class MetaModelDS(MetaTaskModel):
|
|||||||
|
|
||||||
if phase == "train":
|
if phase == "train":
|
||||||
opt.zero_grad()
|
opt.zero_grad()
|
||||||
norm_loss = nn.MSELoss()
|
|
||||||
loss.backward()
|
loss.backward()
|
||||||
opt.step()
|
opt.step()
|
||||||
elif phase == "test":
|
elif phase == "test":
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
# Copyright (c) Microsoft Corporation.
|
# Copyright (c) Microsoft Corporation.
|
||||||
# Licensed under the MIT License.
|
# Licensed under the MIT License.
|
||||||
|
|
||||||
import pandas as pd
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
|||||||
@@ -1,16 +1,17 @@
|
|||||||
# Copyright (c) Microsoft Corporation.
|
# Copyright (c) Microsoft Corporation.
|
||||||
# Licensed under the MIT License.
|
# Licensed under the MIT License.
|
||||||
|
|
||||||
import pandas as pd
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
from qlib.contrib.torch import data_to_tensor
|
|
||||||
|
|
||||||
|
|
||||||
class ICLoss(nn.Module):
|
class ICLoss(nn.Module):
|
||||||
def forward(self, pred, y, idx, skip_size=50):
|
def forward(self, pred, y, idx, skip_size=50):
|
||||||
"""forward.
|
"""forward.
|
||||||
|
FIXME:
|
||||||
|
- Some times it will be a slightly different from the result from `pandas.corr()`
|
||||||
|
- It may be caused by the precision problem of model;
|
||||||
|
|
||||||
:param pred:
|
:param pred:
|
||||||
:param y:
|
:param y:
|
||||||
|
|||||||
@@ -10,17 +10,19 @@ try:
|
|||||||
from .gbdt import LGBModel
|
from .gbdt import LGBModel
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
DEnsembleModel, LGBModel = None, None
|
DEnsembleModel, LGBModel = None, None
|
||||||
print("Please install necessary libs for DEnsembleModel and LGBModel, such as lightgbm.")
|
print(
|
||||||
|
"ModuleNotFoundError. DEnsembleModel and LGBModel are skipped. (optional: maybe installing lightgbm can fix it.)"
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
from .xgboost import XGBModel
|
from .xgboost import XGBModel
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
XGBModel = None
|
XGBModel = None
|
||||||
print("Please install necessary libs for XGBModel, such as xgboost.")
|
print("ModuleNotFoundError. XGBModel is skipped(optional: maybe installing xgboost can fix it).")
|
||||||
try:
|
try:
|
||||||
from .linear import LinearModel
|
from .linear import LinearModel
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
LinearModel = None
|
LinearModel = None
|
||||||
print("Please install necessary libs for LinearModel, such as scipy and sklearn.")
|
print("ModuleNotFoundError. LinearModel is skipped(optional: maybe installing scipy and sklearn can fix it).")
|
||||||
# import pytorch models
|
# import pytorch models
|
||||||
try:
|
try:
|
||||||
from .pytorch_alstm import ALSTM
|
from .pytorch_alstm import ALSTM
|
||||||
@@ -36,6 +38,6 @@ try:
|
|||||||
pytorch_classes = (ALSTM, GATs, GRU, LSTM, DNNModelPytorch, TabnetModel, SFM_Model, TCN, ADD)
|
pytorch_classes = (ALSTM, GATs, GRU, LSTM, DNNModelPytorch, TabnetModel, SFM_Model, TCN, ADD)
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
pytorch_classes = ()
|
pytorch_classes = ()
|
||||||
print("Please install necessary libs for PyTorch models.")
|
print("ModuleNotFoundError. PyTorch models are skipped (optional: maybe installing pytorch can fix it).")
|
||||||
|
|
||||||
all_model_classes = (CatBoostModel, DEnsembleModel, LGBModel, XGBModel, LinearModel) + pytorch_classes
|
all_model_classes = (CatBoostModel, DEnsembleModel, LGBModel, XGBModel, LinearModel) + pytorch_classes
|
||||||
|
|||||||
@@ -160,7 +160,7 @@ class DEnsembleModel(Model, FeatureInt):
|
|||||||
h_avg = h.groupby("bins")["h_value"].mean()
|
h_avg = h.groupby("bins")["h_value"].mean()
|
||||||
weights = pd.Series(np.zeros(N, dtype=float))
|
weights = pd.Series(np.zeros(N, dtype=float))
|
||||||
for i_b, b in enumerate(h_avg.index):
|
for i_b, b in enumerate(h_avg.index):
|
||||||
weights[h["bins"] == b] = 1.0 / (self.decay ** k_th * h_avg[i_b] + 0.1)
|
weights[h["bins"] == b] = 1.0 / (self.decay**k_th * h_avg[i_b] + 0.1)
|
||||||
return weights
|
return weights
|
||||||
|
|
||||||
def feature_selection(self, df_train, loss_values):
|
def feature_selection(self, df_train, loss_values):
|
||||||
@@ -249,7 +249,7 @@ class DEnsembleModel(Model, FeatureInt):
|
|||||||
return pred
|
return pred
|
||||||
|
|
||||||
def predict_sub(self, submodel, df_data, features):
|
def predict_sub(self, submodel, df_data, features):
|
||||||
x_data, y_data = df_data["feature"].loc[:, features], df_data["label"]
|
x_data = df_data["feature"].loc[:, features]
|
||||||
pred_sub = pd.Series(submodel.predict(x_data.values), index=x_data.index)
|
pred_sub = pd.Series(submodel.predict(x_data.values), index=x_data.index)
|
||||||
return pred_sub
|
return pred_sub
|
||||||
|
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ from ...data.dataset import DatasetH
|
|||||||
from ...data.dataset.handler import DataHandlerLP
|
from ...data.dataset.handler import DataHandlerLP
|
||||||
from ...model.interpret.base import LightGBMFInt
|
from ...model.interpret.base import LightGBMFInt
|
||||||
from ...data.dataset.weight import Reweighter
|
from ...data.dataset.weight import Reweighter
|
||||||
|
from qlib.workflow import R
|
||||||
|
|
||||||
|
|
||||||
class LGBModel(ModelFT, LightGBMFInt):
|
class LGBModel(ModelFT, LightGBMFInt):
|
||||||
@@ -59,27 +60,34 @@ class LGBModel(ModelFT, LightGBMFInt):
|
|||||||
num_boost_round=None,
|
num_boost_round=None,
|
||||||
early_stopping_rounds=None,
|
early_stopping_rounds=None,
|
||||||
verbose_eval=20,
|
verbose_eval=20,
|
||||||
evals_result=dict(),
|
evals_result=None,
|
||||||
reweighter=None,
|
reweighter=None,
|
||||||
**kwargs
|
**kwargs,
|
||||||
):
|
):
|
||||||
|
if evals_result is None:
|
||||||
|
evals_result = {} # in case of unsafety of Python default values
|
||||||
ds_l = self._prepare_data(dataset, reweighter)
|
ds_l = self._prepare_data(dataset, reweighter)
|
||||||
ds, names = list(zip(*ds_l))
|
ds, names = list(zip(*ds_l))
|
||||||
|
early_stopping_callback = lgb.early_stopping(
|
||||||
|
self.early_stopping_rounds if early_stopping_rounds is None else early_stopping_rounds
|
||||||
|
)
|
||||||
|
# NOTE: if you encounter error here. Please upgrade your lightgbm
|
||||||
|
verbose_eval_callback = lgb.log_evaluation(period=verbose_eval)
|
||||||
|
evals_result_callback = lgb.record_evaluation(evals_result)
|
||||||
self.model = lgb.train(
|
self.model = lgb.train(
|
||||||
self.params,
|
self.params,
|
||||||
ds[0], # training dataset
|
ds[0], # training dataset
|
||||||
num_boost_round=self.num_boost_round if num_boost_round is None else num_boost_round,
|
num_boost_round=self.num_boost_round if num_boost_round is None else num_boost_round,
|
||||||
valid_sets=ds,
|
valid_sets=ds,
|
||||||
valid_names=names,
|
valid_names=names,
|
||||||
early_stopping_rounds=(
|
callbacks=[early_stopping_callback, verbose_eval_callback, evals_result_callback],
|
||||||
self.early_stopping_rounds if early_stopping_rounds is None else early_stopping_rounds
|
**kwargs,
|
||||||
),
|
|
||||||
verbose_eval=verbose_eval,
|
|
||||||
evals_result=evals_result,
|
|
||||||
**kwargs
|
|
||||||
)
|
)
|
||||||
for k in names:
|
for k in names:
|
||||||
evals_result[k] = list(evals_result[k].values())[0]
|
for key, val in evals_result[k].items():
|
||||||
|
name = f"{key}.{k}"
|
||||||
|
for epoch, m in enumerate(val):
|
||||||
|
R.log_metrics(**{name.replace("@", "_"): m}, step=epoch)
|
||||||
|
|
||||||
def predict(self, dataset: DatasetH, segment: Union[Text, slice] = "test"):
|
def predict(self, dataset: DatasetH, segment: Union[Text, slice] = "test"):
|
||||||
if self.model is None:
|
if self.model is None:
|
||||||
@@ -101,9 +109,10 @@ class LGBModel(ModelFT, LightGBMFInt):
|
|||||||
verbose level
|
verbose level
|
||||||
"""
|
"""
|
||||||
# Based on existing model and finetune by train more rounds
|
# Based on existing model and finetune by train more rounds
|
||||||
dtrain, _ = self._prepare_data(dataset, reweighter)
|
dtrain, _ = self._prepare_data(dataset, reweighter) # pylint: disable=W0632
|
||||||
if dtrain.empty:
|
if dtrain.empty:
|
||||||
raise ValueError("Empty data from dataset, please check your dataset config.")
|
raise ValueError("Empty data from dataset, please check your dataset config.")
|
||||||
|
verbose_eval_callback = lgb.log_evaluation(period=verbose_eval)
|
||||||
self.model = lgb.train(
|
self.model = lgb.train(
|
||||||
self.params,
|
self.params,
|
||||||
dtrain,
|
dtrain,
|
||||||
@@ -111,5 +120,5 @@ class LGBModel(ModelFT, LightGBMFInt):
|
|||||||
init_model=self.model,
|
init_model=self.model,
|
||||||
valid_sets=[dtrain],
|
valid_sets=[dtrain],
|
||||||
valid_names=["train"],
|
valid_names=["train"],
|
||||||
verbose_eval=verbose_eval,
|
callbacks=[verbose_eval_callback],
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -58,7 +58,7 @@ class HFLGBModel(ModelFT, LightGBMFInt):
|
|||||||
"""
|
"""
|
||||||
Test the signal in high frequency test set
|
Test the signal in high frequency test set
|
||||||
"""
|
"""
|
||||||
if self.model == None:
|
if self.model is None:
|
||||||
raise ValueError("Model hasn't been trained yet")
|
raise ValueError("Model hasn't been trained yet")
|
||||||
df_test = dataset.prepare("test", col_set=["feature", "label"], data_key=DataHandlerLP.DK_I)
|
df_test = dataset.prepare("test", col_set=["feature", "label"], data_key=DataHandlerLP.DK_I)
|
||||||
df_test.dropna(inplace=True)
|
df_test.dropna(inplace=True)
|
||||||
@@ -110,20 +110,21 @@ class HFLGBModel(ModelFT, LightGBMFInt):
|
|||||||
num_boost_round=1000,
|
num_boost_round=1000,
|
||||||
early_stopping_rounds=50,
|
early_stopping_rounds=50,
|
||||||
verbose_eval=20,
|
verbose_eval=20,
|
||||||
evals_result=dict(),
|
evals_result=None,
|
||||||
**kwargs
|
|
||||||
):
|
):
|
||||||
|
if evals_result is None:
|
||||||
|
evals_result = dict()
|
||||||
dtrain, dvalid = self._prepare_data(dataset)
|
dtrain, dvalid = self._prepare_data(dataset)
|
||||||
|
early_stopping_callback = lgb.early_stopping(early_stopping_rounds)
|
||||||
|
verbose_eval_callback = lgb.log_evaluation(period=verbose_eval)
|
||||||
|
evals_result_callback = lgb.record_evaluation(evals_result)
|
||||||
self.model = lgb.train(
|
self.model = lgb.train(
|
||||||
self.params,
|
self.params,
|
||||||
dtrain,
|
dtrain,
|
||||||
num_boost_round=num_boost_round,
|
num_boost_round=num_boost_round,
|
||||||
valid_sets=[dtrain, dvalid],
|
valid_sets=[dtrain, dvalid],
|
||||||
valid_names=["train", "valid"],
|
valid_names=["train", "valid"],
|
||||||
early_stopping_rounds=early_stopping_rounds,
|
callbacks=[early_stopping_callback, verbose_eval_callback, evals_result_callback],
|
||||||
verbose_eval=verbose_eval,
|
|
||||||
evals_result=evals_result,
|
|
||||||
**kwargs
|
|
||||||
)
|
)
|
||||||
evals_result["train"] = list(evals_result["train"].values())[0]
|
evals_result["train"] = list(evals_result["train"].values())[0]
|
||||||
evals_result["valid"] = list(evals_result["valid"].values())[0]
|
evals_result["valid"] = list(evals_result["valid"].values())[0]
|
||||||
@@ -149,6 +150,7 @@ class HFLGBModel(ModelFT, LightGBMFInt):
|
|||||||
"""
|
"""
|
||||||
# Based on existing model and finetune by train more rounds
|
# Based on existing model and finetune by train more rounds
|
||||||
dtrain, _ = self._prepare_data(dataset)
|
dtrain, _ = self._prepare_data(dataset)
|
||||||
|
verbose_eval_callback = lgb.log_evaluation(period=verbose_eval)
|
||||||
self.model = lgb.train(
|
self.model = lgb.train(
|
||||||
self.params,
|
self.params,
|
||||||
dtrain,
|
dtrain,
|
||||||
@@ -156,5 +158,5 @@ class HFLGBModel(ModelFT, LightGBMFInt):
|
|||||||
init_model=self.model,
|
init_model=self.model,
|
||||||
valid_sets=[dtrain],
|
valid_sets=[dtrain],
|
||||||
valid_names=["train"],
|
valid_names=["train"],
|
||||||
verbose_eval=verbose_eval,
|
callbacks=[verbose_eval_callback],
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,12 +1,10 @@
|
|||||||
# Copyright (c) Microsoft Corporation.
|
# Copyright (c) Microsoft Corporation.
|
||||||
import os
|
import os
|
||||||
from pdb import set_trace
|
|
||||||
from torch.utils.data import Dataset, DataLoader
|
from torch.utils.data import Dataset, DataLoader
|
||||||
|
|
||||||
import copy
|
import copy
|
||||||
from typing import Text, Union
|
from typing import Text, Union
|
||||||
|
|
||||||
import math
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import torch
|
import torch
|
||||||
@@ -146,7 +144,7 @@ class ADARNN(Model):
|
|||||||
raise NotImplementedError("optimizer {} is not supported!".format(optimizer))
|
raise NotImplementedError("optimizer {} is not supported!".format(optimizer))
|
||||||
|
|
||||||
self.fitted = False
|
self.fitted = False
|
||||||
self.model.cuda()
|
self.model.to(self.device)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def use_gpu(self):
|
def use_gpu(self):
|
||||||
@@ -155,7 +153,7 @@ class ADARNN(Model):
|
|||||||
def train_AdaRNN(self, train_loader_list, epoch, dist_old=None, weight_mat=None):
|
def train_AdaRNN(self, train_loader_list, epoch, dist_old=None, weight_mat=None):
|
||||||
self.model.train()
|
self.model.train()
|
||||||
criterion = nn.MSELoss()
|
criterion = nn.MSELoss()
|
||||||
dist_mat = torch.zeros(self.num_layers, self.len_seq).cuda()
|
dist_mat = torch.zeros(self.num_layers, self.len_seq).to(self.device)
|
||||||
len_loader = np.inf
|
len_loader = np.inf
|
||||||
for loader in train_loader_list:
|
for loader in train_loader_list:
|
||||||
if len(loader) < len_loader:
|
if len(loader) < len_loader:
|
||||||
@@ -167,7 +165,7 @@ class ADARNN(Model):
|
|||||||
list_label = []
|
list_label = []
|
||||||
for data in data_all:
|
for data in data_all:
|
||||||
# feature :[36, 24, 6]
|
# feature :[36, 24, 6]
|
||||||
feature, label_reg = data[0].cuda().float(), data[1].cuda().float()
|
feature, label_reg = data[0].to(self.device).float(), data[1].to(self.device).float()
|
||||||
list_feat.append(feature)
|
list_feat.append(feature)
|
||||||
list_label.append(label_reg)
|
list_label.append(label_reg)
|
||||||
flag = False
|
flag = False
|
||||||
@@ -181,12 +179,12 @@ class ADARNN(Model):
|
|||||||
if flag:
|
if flag:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
total_loss = torch.zeros(1).cuda()
|
total_loss = torch.zeros(1).to(self.device)
|
||||||
for i in range(len(index)):
|
for i, n in enumerate(index):
|
||||||
feature_s = list_feat[index[i][0]]
|
feature_s = list_feat[n[0]]
|
||||||
feature_t = list_feat[index[i][1]]
|
feature_t = list_feat[n[1]]
|
||||||
label_reg_s = list_label[index[i][0]]
|
label_reg_s = list_label[n[0]]
|
||||||
label_reg_t = list_label[index[i][1]]
|
label_reg_t = list_label[n[1]]
|
||||||
feature_all = torch.cat((feature_s, feature_t), 0)
|
feature_all = torch.cat((feature_s, feature_t), 0)
|
||||||
|
|
||||||
if epoch < self.pre_epoch:
|
if epoch < self.pre_epoch:
|
||||||
@@ -327,7 +325,7 @@ class ADARNN(Model):
|
|||||||
else:
|
else:
|
||||||
end = begin + self.batch_size
|
end = begin + self.batch_size
|
||||||
|
|
||||||
x_batch = torch.from_numpy(x_values[begin:end]).float().cuda()
|
x_batch = torch.from_numpy(x_values[begin:end]).float().to(self.device)
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
pred = self.model.predict(x_batch).detach().cpu().numpy()
|
pred = self.model.predict(x_batch).detach().cpu().numpy()
|
||||||
@@ -337,7 +335,7 @@ class ADARNN(Model):
|
|||||||
return pd.Series(np.concatenate(preds), index=index)
|
return pd.Series(np.concatenate(preds), index=index)
|
||||||
|
|
||||||
def transform_type(self, init_weight):
|
def transform_type(self, init_weight):
|
||||||
weight = torch.ones(self.num_layers, self.len_seq).cuda()
|
weight = torch.ones(self.num_layers, self.len_seq).to(self.device)
|
||||||
for i in range(self.num_layers):
|
for i in range(self.num_layers):
|
||||||
for j in range(self.len_seq):
|
for j in range(self.len_seq):
|
||||||
weight[i, j] = init_weight[i][j].item()
|
weight[i, j] = init_weight[i][j].item()
|
||||||
@@ -391,6 +389,7 @@ class AdaRNN(nn.Module):
|
|||||||
len_seq=9,
|
len_seq=9,
|
||||||
model_type="AdaRNN",
|
model_type="AdaRNN",
|
||||||
trans_loss="mmd",
|
trans_loss="mmd",
|
||||||
|
GPU=0,
|
||||||
):
|
):
|
||||||
super(AdaRNN, self).__init__()
|
super(AdaRNN, self).__init__()
|
||||||
self.use_bottleneck = use_bottleneck
|
self.use_bottleneck = use_bottleneck
|
||||||
@@ -401,6 +400,7 @@ class AdaRNN(nn.Module):
|
|||||||
self.model_type = model_type
|
self.model_type = model_type
|
||||||
self.trans_loss = trans_loss
|
self.trans_loss = trans_loss
|
||||||
self.len_seq = len_seq
|
self.len_seq = len_seq
|
||||||
|
self.device = torch.device("cuda:%d" % (GPU) if torch.cuda.is_available() and GPU >= 0 else "cpu")
|
||||||
in_size = self.n_input
|
in_size = self.n_input
|
||||||
|
|
||||||
features = nn.ModuleList()
|
features = nn.ModuleList()
|
||||||
@@ -410,7 +410,7 @@ class AdaRNN(nn.Module):
|
|||||||
in_size = hidden
|
in_size = hidden
|
||||||
self.features = nn.Sequential(*features)
|
self.features = nn.Sequential(*features)
|
||||||
|
|
||||||
if use_bottleneck == True: # finance
|
if use_bottleneck is True: # finance
|
||||||
self.bottleneck = nn.Sequential(
|
self.bottleneck = nn.Sequential(
|
||||||
nn.Linear(n_hiddens[-1], bottleneck_width),
|
nn.Linear(n_hiddens[-1], bottleneck_width),
|
||||||
nn.Linear(bottleneck_width, bottleneck_width),
|
nn.Linear(bottleneck_width, bottleneck_width),
|
||||||
@@ -449,7 +449,7 @@ class AdaRNN(nn.Module):
|
|||||||
def forward_pre_train(self, x, len_win=0):
|
def forward_pre_train(self, x, len_win=0):
|
||||||
out = self.gru_features(x)
|
out = self.gru_features(x)
|
||||||
fea = out[0] # [2N,L,H]
|
fea = out[0] # [2N,L,H]
|
||||||
if self.use_bottleneck == True:
|
if self.use_bottleneck is True:
|
||||||
fea_bottleneck = self.bottleneck(fea[:, -1, :])
|
fea_bottleneck = self.bottleneck(fea[:, -1, :])
|
||||||
fc_out = self.fc(fea_bottleneck).squeeze()
|
fc_out = self.fc(fea_bottleneck).squeeze()
|
||||||
else:
|
else:
|
||||||
@@ -457,9 +457,9 @@ class AdaRNN(nn.Module):
|
|||||||
|
|
||||||
out_list_all, out_weight_list = out[1], out[2]
|
out_list_all, out_weight_list = out[1], out[2]
|
||||||
out_list_s, out_list_t = self.get_features(out_list_all)
|
out_list_s, out_list_t = self.get_features(out_list_all)
|
||||||
loss_transfer = torch.zeros((1,)).cuda()
|
loss_transfer = torch.zeros((1,)).to(self.device)
|
||||||
for i in range(len(out_list_s)):
|
for i, n in enumerate(out_list_s):
|
||||||
criterion_transder = TransferLoss(loss_type=self.trans_loss, input_dim=out_list_s[i].shape[2])
|
criterion_transder = TransferLoss(loss_type=self.trans_loss, input_dim=n.shape[2])
|
||||||
h_start = 0
|
h_start = 0
|
||||||
for j in range(h_start, self.len_seq, 1):
|
for j in range(h_start, self.len_seq, 1):
|
||||||
i_start = j - len_win if j - len_win >= 0 else 0
|
i_start = j - len_win if j - len_win >= 0 else 0
|
||||||
@@ -471,7 +471,7 @@ class AdaRNN(nn.Module):
|
|||||||
else 1 / (self.len_seq - h_start) * (2 * len_win + 1)
|
else 1 / (self.len_seq - h_start) * (2 * len_win + 1)
|
||||||
)
|
)
|
||||||
loss_transfer = loss_transfer + weight * criterion_transder.compute(
|
loss_transfer = loss_transfer + weight * criterion_transder.compute(
|
||||||
out_list_s[i][:, j, :], out_list_t[i][:, k, :]
|
n[:, j, :], out_list_t[i][:, k, :]
|
||||||
)
|
)
|
||||||
return fc_out, loss_transfer, out_weight_list
|
return fc_out, loss_transfer, out_weight_list
|
||||||
|
|
||||||
@@ -484,7 +484,7 @@ class AdaRNN(nn.Module):
|
|||||||
out, _ = self.features[i](x_input.float())
|
out, _ = self.features[i](x_input.float())
|
||||||
x_input = out
|
x_input = out
|
||||||
out_lis.append(out)
|
out_lis.append(out)
|
||||||
if self.model_type == "AdaRNN" and predict == False:
|
if self.model_type == "AdaRNN" and predict is False:
|
||||||
out_gate = self.process_gate_weight(x_input, i)
|
out_gate = self.process_gate_weight(x_input, i)
|
||||||
out_weight_list.append(out_gate)
|
out_weight_list.append(out_gate)
|
||||||
return out, out_lis, out_weight_list
|
return out, out_lis, out_weight_list
|
||||||
@@ -518,16 +518,16 @@ class AdaRNN(nn.Module):
|
|||||||
|
|
||||||
out_list_all = out[1]
|
out_list_all = out[1]
|
||||||
out_list_s, out_list_t = self.get_features(out_list_all)
|
out_list_s, out_list_t = self.get_features(out_list_all)
|
||||||
loss_transfer = torch.zeros((1,)).cuda()
|
loss_transfer = torch.zeros((1,)).to(self.device)
|
||||||
if weight_mat is None:
|
if weight_mat is None:
|
||||||
weight = (1.0 / self.len_seq * torch.ones(self.num_layers, self.len_seq)).cuda()
|
weight = (1.0 / self.len_seq * torch.ones(self.num_layers, self.len_seq)).to(self.device)
|
||||||
else:
|
else:
|
||||||
weight = weight_mat
|
weight = weight_mat
|
||||||
dist_mat = torch.zeros(self.num_layers, self.len_seq).cuda()
|
dist_mat = torch.zeros(self.num_layers, self.len_seq).to(self.device)
|
||||||
for i in range(len(out_list_s)):
|
for i, n in enumerate(out_list_s):
|
||||||
criterion_transder = TransferLoss(loss_type=self.trans_loss, input_dim=out_list_s[i].shape[2])
|
criterion_transder = TransferLoss(loss_type=self.trans_loss, input_dim=n.shape[2])
|
||||||
for j in range(self.len_seq):
|
for j in range(self.len_seq):
|
||||||
loss_trans = criterion_transder.compute(out_list_s[i][:, j, :], out_list_t[i][:, j, :])
|
loss_trans = criterion_transder.compute(n[:, j, :], out_list_t[i][:, j, :])
|
||||||
loss_transfer = loss_transfer + weight[i, j] * loss_trans
|
loss_transfer = loss_transfer + weight[i, j] * loss_trans
|
||||||
dist_mat[i, j] = loss_trans
|
dist_mat[i, j] = loss_trans
|
||||||
return fc_out, loss_transfer, dist_mat, weight
|
return fc_out, loss_transfer, dist_mat, weight
|
||||||
@@ -546,7 +546,7 @@ class AdaRNN(nn.Module):
|
|||||||
def predict(self, x):
|
def predict(self, x):
|
||||||
out = self.gru_features(x, predict=True)
|
out = self.gru_features(x, predict=True)
|
||||||
fea = out[0]
|
fea = out[0]
|
||||||
if self.use_bottleneck == True:
|
if self.use_bottleneck is True:
|
||||||
fea_bottleneck = self.bottleneck(fea[:, -1, :])
|
fea_bottleneck = self.bottleneck(fea[:, -1, :])
|
||||||
fc_out = self.fc(fea_bottleneck).squeeze()
|
fc_out = self.fc(fea_bottleneck).squeeze()
|
||||||
else:
|
else:
|
||||||
@@ -555,12 +555,13 @@ class AdaRNN(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class TransferLoss:
|
class TransferLoss:
|
||||||
def __init__(self, loss_type="cosine", input_dim=512):
|
def __init__(self, loss_type="cosine", input_dim=512, GPU=0):
|
||||||
"""
|
"""
|
||||||
Supported loss_type: mmd(mmd_lin), mmd_rbf, coral, cosine, kl, js, mine, adv
|
Supported loss_type: mmd(mmd_lin), mmd_rbf, coral, cosine, kl, js, mine, adv
|
||||||
"""
|
"""
|
||||||
self.loss_type = loss_type
|
self.loss_type = loss_type
|
||||||
self.input_dim = input_dim
|
self.input_dim = input_dim
|
||||||
|
self.device = torch.device("cuda:%d" % (GPU) if torch.cuda.is_available() and GPU >= 0 else "cpu")
|
||||||
|
|
||||||
def compute(self, X, Y):
|
def compute(self, X, Y):
|
||||||
"""Compute adaptation loss
|
"""Compute adaptation loss
|
||||||
@@ -572,22 +573,22 @@ class TransferLoss:
|
|||||||
Returns:
|
Returns:
|
||||||
[tensor] -- transfer loss
|
[tensor] -- transfer loss
|
||||||
"""
|
"""
|
||||||
if self.loss_type == "mmd_lin" or self.loss_type == "mmd":
|
if self.loss_type in ("mmd_lin", "mmd"):
|
||||||
mmdloss = MMD_loss(kernel_type="linear")
|
mmdloss = MMD_loss(kernel_type="linear")
|
||||||
loss = mmdloss(X, Y)
|
loss = mmdloss(X, Y)
|
||||||
elif self.loss_type == "coral":
|
elif self.loss_type == "coral":
|
||||||
loss = CORAL(X, Y)
|
loss = CORAL(X, Y, self.device)
|
||||||
elif self.loss_type == "cosine" or self.loss_type == "cos":
|
elif self.loss_type in ("cosine", "cos"):
|
||||||
loss = 1 - cosine(X, Y)
|
loss = 1 - cosine(X, Y)
|
||||||
elif self.loss_type == "kl":
|
elif self.loss_type == "kl":
|
||||||
loss = kl_div(X, Y)
|
loss = kl_div(X, Y)
|
||||||
elif self.loss_type == "js":
|
elif self.loss_type == "js":
|
||||||
loss = js(X, Y)
|
loss = js(X, Y)
|
||||||
elif self.loss_type == "mine":
|
elif self.loss_type == "mine":
|
||||||
mine_model = Mine_estimator(input_dim=self.input_dim, hidden_dim=60).cuda()
|
mine_model = Mine_estimator(input_dim=self.input_dim, hidden_dim=60).to(self.device)
|
||||||
loss = mine_model(X, Y)
|
loss = mine_model(X, Y)
|
||||||
elif self.loss_type == "adv":
|
elif self.loss_type == "adv":
|
||||||
loss = adv(X, Y, input_dim=self.input_dim, hidden_dim=32)
|
loss = adv(X, Y, self.device, input_dim=self.input_dim, hidden_dim=32)
|
||||||
elif self.loss_type == "mmd_rbf":
|
elif self.loss_type == "mmd_rbf":
|
||||||
mmdloss = MMD_loss(kernel_type="rbf")
|
mmdloss = MMD_loss(kernel_type="rbf")
|
||||||
loss = mmdloss(X, Y)
|
loss = mmdloss(X, Y)
|
||||||
@@ -632,12 +633,12 @@ class Discriminator(nn.Module):
|
|||||||
return x
|
return x
|
||||||
|
|
||||||
|
|
||||||
def adv(source, target, input_dim=256, hidden_dim=512):
|
def adv(source, target, device, input_dim=256, hidden_dim=512):
|
||||||
domain_loss = nn.BCELoss()
|
domain_loss = nn.BCELoss()
|
||||||
# !!! Pay attention to .cuda !!!
|
# !!! Pay attention to .cuda !!!
|
||||||
adv_net = Discriminator(input_dim, hidden_dim).cuda()
|
adv_net = Discriminator(input_dim, hidden_dim).to(device)
|
||||||
domain_src = torch.ones(len(source)).cuda()
|
domain_src = torch.ones(len(source)).to(device)
|
||||||
domain_tar = torch.zeros(len(target)).cuda()
|
domain_tar = torch.zeros(len(target)).to(device)
|
||||||
domain_src, domain_tar = domain_src.view(domain_src.shape[0], 1), domain_tar.view(domain_tar.shape[0], 1)
|
domain_src, domain_tar = domain_src.view(domain_src.shape[0], 1), domain_tar.view(domain_tar.shape[0], 1)
|
||||||
reverse_src = ReverseLayerF.apply(source, 1)
|
reverse_src = ReverseLayerF.apply(source, 1)
|
||||||
reverse_tar = ReverseLayerF.apply(target, 1)
|
reverse_tar = ReverseLayerF.apply(target, 1)
|
||||||
@@ -648,16 +649,16 @@ def adv(source, target, input_dim=256, hidden_dim=512):
|
|||||||
return loss
|
return loss
|
||||||
|
|
||||||
|
|
||||||
def CORAL(source, target):
|
def CORAL(source, target, device):
|
||||||
d = source.size(1)
|
d = source.size(1)
|
||||||
ns, nt = source.size(0), target.size(0)
|
ns, nt = source.size(0), target.size(0)
|
||||||
|
|
||||||
# source covariance
|
# source covariance
|
||||||
tmp_s = torch.ones((1, ns)).cuda() @ source
|
tmp_s = torch.ones((1, ns)).to(device) @ source
|
||||||
cs = (source.t() @ source - (tmp_s.t() @ tmp_s) / ns) / (ns - 1)
|
cs = (source.t() @ source - (tmp_s.t() @ tmp_s) / ns) / (ns - 1)
|
||||||
|
|
||||||
# target covariance
|
# target covariance
|
||||||
tmp_t = torch.ones((1, nt)).cuda() @ target
|
tmp_t = torch.ones((1, nt)).to(device) @ target
|
||||||
ct = (target.t() @ target - (tmp_t.t() @ tmp_t) / nt) / (nt - 1)
|
ct = (target.t() @ target - (tmp_t.t() @ tmp_t) / nt) / (nt - 1)
|
||||||
|
|
||||||
# frobenius norm
|
# frobenius norm
|
||||||
@@ -684,9 +685,9 @@ class MMD_loss(nn.Module):
|
|||||||
if fix_sigma:
|
if fix_sigma:
|
||||||
bandwidth = fix_sigma
|
bandwidth = fix_sigma
|
||||||
else:
|
else:
|
||||||
bandwidth = torch.sum(L2_distance.data) / (n_samples ** 2 - n_samples)
|
bandwidth = torch.sum(L2_distance.data) / (n_samples**2 - n_samples)
|
||||||
bandwidth /= kernel_mul ** (kernel_num // 2)
|
bandwidth /= kernel_mul ** (kernel_num // 2)
|
||||||
bandwidth_list = [bandwidth * (kernel_mul ** i) for i in range(kernel_num)]
|
bandwidth_list = [bandwidth * (kernel_mul**i) for i in range(kernel_num)]
|
||||||
kernel_val = [torch.exp(-L2_distance / bandwidth_temp) for bandwidth_temp in bandwidth_list]
|
kernel_val = [torch.exp(-L2_distance / bandwidth_temp) for bandwidth_temp in bandwidth_list]
|
||||||
return sum(kernel_val)
|
return sum(kernel_val)
|
||||||
|
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ from qlib.contrib.model.pytorch_lstm import LSTMModel
|
|||||||
from qlib.contrib.model.pytorch_utils import count_parameters
|
from qlib.contrib.model.pytorch_utils import count_parameters
|
||||||
from qlib.data.dataset import DatasetH
|
from qlib.data.dataset import DatasetH
|
||||||
from qlib.data.dataset.handler import DataHandlerLP
|
from qlib.data.dataset.handler import DataHandlerLP
|
||||||
from qlib.data.dataset.processor import CSRankNorm
|
|
||||||
from qlib.log import get_module_logger
|
from qlib.log import get_module_logger
|
||||||
from qlib.model.base import Model
|
from qlib.model.base import Model
|
||||||
from qlib.utils import get_or_create_path
|
from qlib.utils import get_or_create_path
|
||||||
|
|||||||
@@ -5,7 +5,6 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from typing import Text, Union
|
from typing import Text, Union
|
||||||
@@ -150,7 +149,7 @@ class ALSTM(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
@@ -312,8 +311,8 @@ class ALSTMModel(nn.Module):
|
|||||||
def _build_model(self):
|
def _build_model(self):
|
||||||
try:
|
try:
|
||||||
klass = getattr(nn, self.rnn_type.upper())
|
klass = getattr(nn, self.rnn_type.upper())
|
||||||
except:
|
except Exception as e:
|
||||||
raise ValueError("unknown rnn_type `%s`" % self.rnn_type)
|
raise ValueError("unknown rnn_type `%s`" % self.rnn_type) from e
|
||||||
self.net = nn.Sequential()
|
self.net = nn.Sequential()
|
||||||
self.net.add_module("fc_in", nn.Linear(in_features=self.input_size, out_features=self.hid_size))
|
self.net.add_module("fc_in", nn.Linear(in_features=self.input_size, out_features=self.hid_size))
|
||||||
self.net.add_module("act", nn.Tanh())
|
self.net.add_module("act", nn.Tanh())
|
||||||
|
|||||||
@@ -5,7 +5,6 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from typing import Text, Union
|
from typing import Text, Union
|
||||||
@@ -20,7 +19,7 @@ from torch.utils.data import DataLoader
|
|||||||
|
|
||||||
from .pytorch_utils import count_parameters
|
from .pytorch_utils import count_parameters
|
||||||
from ...model.base import Model
|
from ...model.base import Model
|
||||||
from ...data.dataset import DatasetH, TSDatasetH
|
from ...data.dataset import DatasetH
|
||||||
from ...data.dataset.handler import DataHandlerLP
|
from ...data.dataset.handler import DataHandlerLP
|
||||||
from ...model.utils import ConcatDataset
|
from ...model.utils import ConcatDataset
|
||||||
from ...data.dataset.weight import Reweighter
|
from ...data.dataset.weight import Reweighter
|
||||||
@@ -160,7 +159,7 @@ class ALSTM(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
@@ -320,8 +319,8 @@ class ALSTMModel(nn.Module):
|
|||||||
def _build_model(self):
|
def _build_model(self):
|
||||||
try:
|
try:
|
||||||
klass = getattr(nn, self.rnn_type.upper())
|
klass = getattr(nn, self.rnn_type.upper())
|
||||||
except:
|
except Exception as e:
|
||||||
raise ValueError("unknown rnn_type `%s`" % self.rnn_type)
|
raise ValueError("unknown rnn_type `%s`" % self.rnn_type) from e
|
||||||
self.net = nn.Sequential()
|
self.net = nn.Sequential()
|
||||||
self.net.add_module("fc_in", nn.Linear(in_features=self.input_size, out_features=self.hid_size))
|
self.net.add_module("fc_in", nn.Linear(in_features=self.input_size, out_features=self.hid_size))
|
||||||
self.net.add_module("act", nn.Tanh())
|
self.net.add_module("act", nn.Tanh())
|
||||||
|
|||||||
@@ -5,7 +5,6 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from typing import Text, Union
|
from typing import Text, Union
|
||||||
@@ -158,7 +157,7 @@ class GATs(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
@@ -263,7 +262,9 @@ class GATs(Model):
|
|||||||
pretrained_model.load_state_dict(torch.load(self.model_path, map_location=self.device))
|
pretrained_model.load_state_dict(torch.load(self.model_path, map_location=self.device))
|
||||||
|
|
||||||
model_dict = self.GAT_model.state_dict()
|
model_dict = self.GAT_model.state_dict()
|
||||||
pretrained_dict = {k: v for k, v in pretrained_model.state_dict().items() if k in model_dict}
|
pretrained_dict = {
|
||||||
|
k: v for k, v in pretrained_model.state_dict().items() if k in model_dict # pylint: disable=E1135
|
||||||
|
}
|
||||||
model_dict.update(pretrained_dict)
|
model_dict.update(pretrained_dict)
|
||||||
self.GAT_model.load_state_dict(model_dict)
|
self.GAT_model.load_state_dict(model_dict)
|
||||||
self.logger.info("Loading pretrained model Done...")
|
self.logger.info("Loading pretrained model Done...")
|
||||||
|
|||||||
@@ -5,7 +5,6 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import copy
|
import copy
|
||||||
@@ -19,7 +18,6 @@ from torch.utils.data import Sampler
|
|||||||
|
|
||||||
from .pytorch_utils import count_parameters
|
from .pytorch_utils import count_parameters
|
||||||
from ...model.base import Model
|
from ...model.base import Model
|
||||||
from ...data.dataset import DatasetH
|
|
||||||
from ...data.dataset.handler import DataHandlerLP
|
from ...data.dataset.handler import DataHandlerLP
|
||||||
from ...contrib.model.pytorch_lstm import LSTMModel
|
from ...contrib.model.pytorch_lstm import LSTMModel
|
||||||
from ...contrib.model.pytorch_gru import GRUModel
|
from ...contrib.model.pytorch_gru import GRUModel
|
||||||
@@ -178,7 +176,7 @@ class GATs(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
@@ -279,7 +277,9 @@ class GATs(Model):
|
|||||||
pretrained_model.load_state_dict(torch.load(self.model_path, map_location=self.device))
|
pretrained_model.load_state_dict(torch.load(self.model_path, map_location=self.device))
|
||||||
|
|
||||||
model_dict = self.GAT_model.state_dict()
|
model_dict = self.GAT_model.state_dict()
|
||||||
pretrained_dict = {k: v for k, v in pretrained_model.state_dict().items() if k in model_dict}
|
pretrained_dict = {
|
||||||
|
k: v for k, v in pretrained_model.state_dict().items() if k in model_dict # pylint: disable=E1135
|
||||||
|
}
|
||||||
model_dict.update(pretrained_dict)
|
model_dict.update(pretrained_dict)
|
||||||
self.GAT_model.load_state_dict(model_dict)
|
self.GAT_model.load_state_dict(model_dict)
|
||||||
self.logger.info("Loading pretrained model Done...")
|
self.logger.info("Loading pretrained model Done...")
|
||||||
|
|||||||
@@ -5,7 +5,6 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from typing import Text, Union
|
from typing import Text, Union
|
||||||
@@ -150,7 +149,7 @@ class GRU(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
|
|||||||
@@ -5,7 +5,6 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import copy
|
import copy
|
||||||
@@ -19,7 +18,6 @@ from torch.utils.data import DataLoader
|
|||||||
|
|
||||||
from .pytorch_utils import count_parameters
|
from .pytorch_utils import count_parameters
|
||||||
from ...model.base import Model
|
from ...model.base import Model
|
||||||
from ...data.dataset import DatasetH, TSDatasetH
|
|
||||||
from ...data.dataset.handler import DataHandlerLP
|
from ...data.dataset.handler import DataHandlerLP
|
||||||
from ...model.utils import ConcatDataset
|
from ...model.utils import ConcatDataset
|
||||||
from ...data.dataset.weight import Reweighter
|
from ...data.dataset.weight import Reweighter
|
||||||
@@ -159,7 +157,7 @@ class GRU(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
|
|||||||
501
qlib/contrib/model/pytorch_hist.py
Normal file
501
qlib/contrib/model/pytorch_hist.py
Normal file
@@ -0,0 +1,501 @@
|
|||||||
|
# Copyright (c) Microsoft Corporation.
|
||||||
|
# Licensed under the MIT License.
|
||||||
|
|
||||||
|
|
||||||
|
from __future__ import division
|
||||||
|
from __future__ import print_function
|
||||||
|
|
||||||
|
import os
|
||||||
|
import numpy as np
|
||||||
|
import pandas as pd
|
||||||
|
from typing import Text, Union
|
||||||
|
import urllib.request
|
||||||
|
import copy
|
||||||
|
from ...utils import get_or_create_path
|
||||||
|
from ...log import get_module_logger
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.optim as optim
|
||||||
|
from .pytorch_utils import count_parameters
|
||||||
|
from ...model.base import Model
|
||||||
|
from ...data.dataset import DatasetH
|
||||||
|
from ...data.dataset.handler import DataHandlerLP
|
||||||
|
from ...contrib.model.pytorch_lstm import LSTMModel
|
||||||
|
from ...contrib.model.pytorch_gru import GRUModel
|
||||||
|
|
||||||
|
|
||||||
|
class HIST(Model):
|
||||||
|
"""HIST Model
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
lr : float
|
||||||
|
learning rate
|
||||||
|
d_feat : int
|
||||||
|
input dimensions for each time step
|
||||||
|
metric : str
|
||||||
|
the evaluate metric used in early stop
|
||||||
|
optimizer : str
|
||||||
|
optimizer name
|
||||||
|
GPU : str
|
||||||
|
the GPU ID(s) used for training
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
d_feat=6,
|
||||||
|
hidden_size=64,
|
||||||
|
num_layers=2,
|
||||||
|
dropout=0.0,
|
||||||
|
n_epochs=200,
|
||||||
|
lr=0.001,
|
||||||
|
metric="",
|
||||||
|
early_stop=20,
|
||||||
|
loss="mse",
|
||||||
|
base_model="GRU",
|
||||||
|
model_path=None,
|
||||||
|
stock2concept=None,
|
||||||
|
stock_index=None,
|
||||||
|
optimizer="adam",
|
||||||
|
GPU=0,
|
||||||
|
seed=None,
|
||||||
|
**kwargs
|
||||||
|
):
|
||||||
|
# Set logger.
|
||||||
|
self.logger = get_module_logger("HIST")
|
||||||
|
self.logger.info("HIST pytorch version...")
|
||||||
|
|
||||||
|
# set hyper-parameters.
|
||||||
|
self.d_feat = d_feat
|
||||||
|
self.hidden_size = hidden_size
|
||||||
|
self.num_layers = num_layers
|
||||||
|
self.dropout = dropout
|
||||||
|
self.n_epochs = n_epochs
|
||||||
|
self.lr = lr
|
||||||
|
self.metric = metric
|
||||||
|
self.early_stop = early_stop
|
||||||
|
self.optimizer = optimizer.lower()
|
||||||
|
self.loss = loss
|
||||||
|
self.base_model = base_model
|
||||||
|
self.model_path = model_path
|
||||||
|
self.stock2concept = stock2concept
|
||||||
|
self.stock_index = stock_index
|
||||||
|
self.device = torch.device("cuda:%d" % (GPU) if torch.cuda.is_available() and GPU >= 0 else "cpu")
|
||||||
|
self.seed = seed
|
||||||
|
|
||||||
|
self.logger.info(
|
||||||
|
"HIST parameters setting:"
|
||||||
|
"\nd_feat : {}"
|
||||||
|
"\nhidden_size : {}"
|
||||||
|
"\nnum_layers : {}"
|
||||||
|
"\ndropout : {}"
|
||||||
|
"\nn_epochs : {}"
|
||||||
|
"\nlr : {}"
|
||||||
|
"\nmetric : {}"
|
||||||
|
"\nearly_stop : {}"
|
||||||
|
"\noptimizer : {}"
|
||||||
|
"\nloss_type : {}"
|
||||||
|
"\nbase_model : {}"
|
||||||
|
"\nmodel_path : {}"
|
||||||
|
"\nstock2concept : {}"
|
||||||
|
"\nstock_index : {}"
|
||||||
|
"\nuse_GPU : {}"
|
||||||
|
"\nseed : {}".format(
|
||||||
|
d_feat,
|
||||||
|
hidden_size,
|
||||||
|
num_layers,
|
||||||
|
dropout,
|
||||||
|
n_epochs,
|
||||||
|
lr,
|
||||||
|
metric,
|
||||||
|
early_stop,
|
||||||
|
optimizer.lower(),
|
||||||
|
loss,
|
||||||
|
base_model,
|
||||||
|
model_path,
|
||||||
|
stock2concept,
|
||||||
|
stock_index,
|
||||||
|
GPU,
|
||||||
|
seed,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.seed is not None:
|
||||||
|
np.random.seed(self.seed)
|
||||||
|
torch.manual_seed(self.seed)
|
||||||
|
|
||||||
|
self.HIST_model = HISTModel(
|
||||||
|
d_feat=self.d_feat,
|
||||||
|
hidden_size=self.hidden_size,
|
||||||
|
num_layers=self.num_layers,
|
||||||
|
dropout=self.dropout,
|
||||||
|
base_model=self.base_model,
|
||||||
|
)
|
||||||
|
self.logger.info("model:\n{:}".format(self.HIST_model))
|
||||||
|
self.logger.info("model size: {:.4f} MB".format(count_parameters(self.HIST_model)))
|
||||||
|
if optimizer.lower() == "adam":
|
||||||
|
self.train_optimizer = optim.Adam(self.HIST_model.parameters(), lr=self.lr)
|
||||||
|
elif optimizer.lower() == "gd":
|
||||||
|
self.train_optimizer = optim.SGD(self.HIST_model.parameters(), lr=self.lr)
|
||||||
|
else:
|
||||||
|
raise NotImplementedError("optimizer {} is not supported!".format(optimizer))
|
||||||
|
|
||||||
|
self.fitted = False
|
||||||
|
self.HIST_model.to(self.device)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def use_gpu(self):
|
||||||
|
return self.device != torch.device("cpu")
|
||||||
|
|
||||||
|
def mse(self, pred, label):
|
||||||
|
loss = (pred - label) ** 2
|
||||||
|
return torch.mean(loss)
|
||||||
|
|
||||||
|
def loss_fn(self, pred, label):
|
||||||
|
mask = ~torch.isnan(label)
|
||||||
|
|
||||||
|
if self.loss == "mse":
|
||||||
|
return self.mse(pred[mask], label[mask])
|
||||||
|
|
||||||
|
raise ValueError("unknown loss `%s`" % self.loss)
|
||||||
|
|
||||||
|
def metric_fn(self, pred, label):
|
||||||
|
|
||||||
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
|
if self.metric == "ic":
|
||||||
|
x = pred[mask]
|
||||||
|
y = label[mask]
|
||||||
|
|
||||||
|
vx = x - torch.mean(x)
|
||||||
|
vy = y - torch.mean(y)
|
||||||
|
return torch.sum(vx * vy) / (torch.sqrt(torch.sum(vx**2)) * torch.sqrt(torch.sum(vy**2)))
|
||||||
|
|
||||||
|
if self.metric == ("", "loss"):
|
||||||
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
|
|
||||||
|
def get_daily_inter(self, df, shuffle=False):
|
||||||
|
# organize the train data into daily batches
|
||||||
|
daily_count = df.groupby(level=0).size().values
|
||||||
|
daily_index = np.roll(np.cumsum(daily_count), 1)
|
||||||
|
daily_index[0] = 0
|
||||||
|
if shuffle:
|
||||||
|
# shuffle data
|
||||||
|
daily_shuffle = list(zip(daily_index, daily_count))
|
||||||
|
np.random.shuffle(daily_shuffle)
|
||||||
|
daily_index, daily_count = zip(*daily_shuffle)
|
||||||
|
return daily_index, daily_count
|
||||||
|
|
||||||
|
def train_epoch(self, x_train, y_train, stock_index):
|
||||||
|
|
||||||
|
stock2concept_matrix = np.load(self.stock2concept)
|
||||||
|
x_train_values = x_train.values
|
||||||
|
y_train_values = np.squeeze(y_train.values)
|
||||||
|
stock_index = stock_index.values
|
||||||
|
stock_index[np.isnan(stock_index)] = 733
|
||||||
|
self.HIST_model.train()
|
||||||
|
|
||||||
|
# organize the train data into daily batches
|
||||||
|
daily_index, daily_count = self.get_daily_inter(x_train, shuffle=True)
|
||||||
|
|
||||||
|
for idx, count in zip(daily_index, daily_count):
|
||||||
|
batch = slice(idx, idx + count)
|
||||||
|
feature = torch.from_numpy(x_train_values[batch]).float().to(self.device)
|
||||||
|
concept_matrix = torch.from_numpy(stock2concept_matrix[stock_index[batch]]).float().to(self.device)
|
||||||
|
label = torch.from_numpy(y_train_values[batch]).float().to(self.device)
|
||||||
|
pred = self.HIST_model(feature, concept_matrix)
|
||||||
|
loss = self.loss_fn(pred, label)
|
||||||
|
|
||||||
|
self.train_optimizer.zero_grad()
|
||||||
|
loss.backward()
|
||||||
|
torch.nn.utils.clip_grad_value_(self.HIST_model.parameters(), 3.0)
|
||||||
|
self.train_optimizer.step()
|
||||||
|
|
||||||
|
def test_epoch(self, data_x, data_y, stock_index):
|
||||||
|
|
||||||
|
# prepare training data
|
||||||
|
stock2concept_matrix = np.load(self.stock2concept)
|
||||||
|
x_values = data_x.values
|
||||||
|
y_values = np.squeeze(data_y.values)
|
||||||
|
stock_index = stock_index.values
|
||||||
|
stock_index[np.isnan(stock_index)] = 733
|
||||||
|
self.HIST_model.eval()
|
||||||
|
|
||||||
|
scores = []
|
||||||
|
losses = []
|
||||||
|
|
||||||
|
# organize the test data into daily batches
|
||||||
|
daily_index, daily_count = self.get_daily_inter(data_x, shuffle=False)
|
||||||
|
|
||||||
|
for idx, count in zip(daily_index, daily_count):
|
||||||
|
batch = slice(idx, idx + count)
|
||||||
|
feature = torch.from_numpy(x_values[batch]).float().to(self.device)
|
||||||
|
concept_matrix = torch.from_numpy(stock2concept_matrix[stock_index[batch]]).float().to(self.device)
|
||||||
|
label = torch.from_numpy(y_values[batch]).float().to(self.device)
|
||||||
|
with torch.no_grad():
|
||||||
|
pred = self.HIST_model(feature, concept_matrix)
|
||||||
|
loss = self.loss_fn(pred, label)
|
||||||
|
losses.append(loss.item())
|
||||||
|
|
||||||
|
score = self.metric_fn(pred, label)
|
||||||
|
scores.append(score.item())
|
||||||
|
|
||||||
|
return np.mean(losses), np.mean(scores)
|
||||||
|
|
||||||
|
def fit(
|
||||||
|
self,
|
||||||
|
dataset: DatasetH,
|
||||||
|
evals_result=dict(),
|
||||||
|
save_path=None,
|
||||||
|
):
|
||||||
|
df_train, df_valid, df_test = dataset.prepare(
|
||||||
|
["train", "valid", "test"],
|
||||||
|
col_set=["feature", "label"],
|
||||||
|
data_key=DataHandlerLP.DK_L,
|
||||||
|
)
|
||||||
|
if df_train.empty or df_valid.empty:
|
||||||
|
raise ValueError("Empty data from dataset, please check your dataset config.")
|
||||||
|
|
||||||
|
if not os.path.exists(self.stock2concept):
|
||||||
|
url = "http://fintech.msra.cn/stock_data/downloads/qlib_csi300_stock2concept.npy"
|
||||||
|
urllib.request.urlretrieve(url, self.stock2concept)
|
||||||
|
|
||||||
|
stock_index = np.load(self.stock_index, allow_pickle=True).item()
|
||||||
|
df_train["stock_index"] = 733
|
||||||
|
df_train["stock_index"] = df_train.index.get_level_values("instrument").map(stock_index)
|
||||||
|
df_valid["stock_index"] = 733
|
||||||
|
df_valid["stock_index"] = df_valid.index.get_level_values("instrument").map(stock_index)
|
||||||
|
|
||||||
|
x_train, y_train, stock_index_train = df_train["feature"], df_train["label"], df_train["stock_index"]
|
||||||
|
x_valid, y_valid, stock_index_valid = df_valid["feature"], df_valid["label"], df_valid["stock_index"]
|
||||||
|
|
||||||
|
save_path = get_or_create_path(save_path)
|
||||||
|
|
||||||
|
stop_steps = 0
|
||||||
|
best_score = -np.inf
|
||||||
|
best_epoch = 0
|
||||||
|
evals_result["train"] = []
|
||||||
|
evals_result["valid"] = []
|
||||||
|
|
||||||
|
# load pretrained base_model
|
||||||
|
if self.base_model == "LSTM":
|
||||||
|
pretrained_model = LSTMModel()
|
||||||
|
elif self.base_model == "GRU":
|
||||||
|
pretrained_model = GRUModel()
|
||||||
|
else:
|
||||||
|
raise ValueError("unknown base model name `%s`" % self.base_model)
|
||||||
|
|
||||||
|
if self.model_path is not None:
|
||||||
|
self.logger.info("Loading pretrained model...")
|
||||||
|
pretrained_model.load_state_dict(torch.load(self.model_path))
|
||||||
|
|
||||||
|
model_dict = self.HIST_model.state_dict()
|
||||||
|
pretrained_dict = {k: v for k, v in pretrained_model.state_dict().items() if k in model_dict}
|
||||||
|
model_dict.update(pretrained_dict)
|
||||||
|
self.HIST_model.load_state_dict(model_dict)
|
||||||
|
self.logger.info("Loading pretrained model Done...")
|
||||||
|
|
||||||
|
# train
|
||||||
|
self.logger.info("training...")
|
||||||
|
self.fitted = True
|
||||||
|
|
||||||
|
for step in range(self.n_epochs):
|
||||||
|
self.logger.info("Epoch%d:", step)
|
||||||
|
self.logger.info("training...")
|
||||||
|
self.train_epoch(x_train, y_train, stock_index_train)
|
||||||
|
|
||||||
|
self.logger.info("evaluating...")
|
||||||
|
train_loss, train_score = self.test_epoch(x_train, y_train, stock_index_train)
|
||||||
|
val_loss, val_score = self.test_epoch(x_valid, y_valid, stock_index_valid)
|
||||||
|
self.logger.info("train %.6f, valid %.6f" % (train_score, val_score))
|
||||||
|
evals_result["train"].append(train_score)
|
||||||
|
evals_result["valid"].append(val_score)
|
||||||
|
|
||||||
|
if val_score > best_score:
|
||||||
|
best_score = val_score
|
||||||
|
stop_steps = 0
|
||||||
|
best_epoch = step
|
||||||
|
best_param = copy.deepcopy(self.HIST_model.state_dict())
|
||||||
|
else:
|
||||||
|
stop_steps += 1
|
||||||
|
if stop_steps >= self.early_stop:
|
||||||
|
self.logger.info("early stop")
|
||||||
|
break
|
||||||
|
|
||||||
|
self.logger.info("best score: %.6lf @ %d" % (best_score, best_epoch))
|
||||||
|
self.HIST_model.load_state_dict(best_param)
|
||||||
|
torch.save(best_param, save_path)
|
||||||
|
|
||||||
|
def predict(self, dataset: DatasetH, segment: Union[Text, slice] = "test"):
|
||||||
|
if not self.fitted:
|
||||||
|
raise ValueError("model is not fitted yet!")
|
||||||
|
|
||||||
|
stock2concept_matrix = np.load(self.stock2concept)
|
||||||
|
stock_index = np.load(self.stock_index, allow_pickle=True).item()
|
||||||
|
df_test = dataset.prepare(segment, col_set="feature", data_key=DataHandlerLP.DK_I)
|
||||||
|
df_test["stock_index"] = 733
|
||||||
|
df_test["stock_index"] = df_test.index.get_level_values("instrument").map(stock_index)
|
||||||
|
stock_index_test = df_test["stock_index"].values
|
||||||
|
stock_index_test[np.isnan(stock_index_test)] = 733
|
||||||
|
stock_index_test = stock_index_test.astype("int")
|
||||||
|
df_test = df_test.drop(["stock_index"], axis=1)
|
||||||
|
index = df_test.index
|
||||||
|
|
||||||
|
self.HIST_model.eval()
|
||||||
|
x_values = df_test.values
|
||||||
|
preds = []
|
||||||
|
|
||||||
|
# organize the data into daily batches
|
||||||
|
daily_index, daily_count = self.get_daily_inter(df_test, shuffle=False)
|
||||||
|
|
||||||
|
for idx, count in zip(daily_index, daily_count):
|
||||||
|
batch = slice(idx, idx + count)
|
||||||
|
x_batch = torch.from_numpy(x_values[batch]).float().to(self.device)
|
||||||
|
concept_matrix = torch.from_numpy(stock2concept_matrix[stock_index_test[batch]]).float().to(self.device)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
pred = self.HIST_model(x_batch, concept_matrix).detach().cpu().numpy()
|
||||||
|
|
||||||
|
preds.append(pred)
|
||||||
|
|
||||||
|
return pd.Series(np.concatenate(preds), index=index)
|
||||||
|
|
||||||
|
|
||||||
|
class HISTModel(nn.Module):
|
||||||
|
def __init__(self, d_feat=6, hidden_size=64, num_layers=2, dropout=0.0, base_model="GRU"):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.d_feat = d_feat
|
||||||
|
self.hidden_size = hidden_size
|
||||||
|
|
||||||
|
if base_model == "GRU":
|
||||||
|
self.rnn = nn.GRU(
|
||||||
|
input_size=d_feat,
|
||||||
|
hidden_size=hidden_size,
|
||||||
|
num_layers=num_layers,
|
||||||
|
batch_first=True,
|
||||||
|
dropout=dropout,
|
||||||
|
)
|
||||||
|
elif base_model == "LSTM":
|
||||||
|
self.rnn = nn.LSTM(
|
||||||
|
input_size=d_feat,
|
||||||
|
hidden_size=hidden_size,
|
||||||
|
num_layers=num_layers,
|
||||||
|
batch_first=True,
|
||||||
|
dropout=dropout,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise ValueError("unknown base model name `%s`" % base_model)
|
||||||
|
|
||||||
|
self.fc_es = nn.Linear(hidden_size, hidden_size)
|
||||||
|
torch.nn.init.xavier_uniform_(self.fc_es.weight)
|
||||||
|
self.fc_is = nn.Linear(hidden_size, hidden_size)
|
||||||
|
torch.nn.init.xavier_uniform_(self.fc_is.weight)
|
||||||
|
|
||||||
|
self.fc_es_middle = nn.Linear(hidden_size, hidden_size)
|
||||||
|
torch.nn.init.xavier_uniform_(self.fc_es_middle.weight)
|
||||||
|
self.fc_is_middle = nn.Linear(hidden_size, hidden_size)
|
||||||
|
torch.nn.init.xavier_uniform_(self.fc_is_middle.weight)
|
||||||
|
|
||||||
|
self.fc_es_fore = nn.Linear(hidden_size, hidden_size)
|
||||||
|
torch.nn.init.xavier_uniform_(self.fc_es_fore.weight)
|
||||||
|
self.fc_is_fore = nn.Linear(hidden_size, hidden_size)
|
||||||
|
torch.nn.init.xavier_uniform_(self.fc_is_fore.weight)
|
||||||
|
self.fc_indi_fore = nn.Linear(hidden_size, hidden_size)
|
||||||
|
torch.nn.init.xavier_uniform_(self.fc_indi_fore.weight)
|
||||||
|
|
||||||
|
self.fc_es_back = nn.Linear(hidden_size, hidden_size)
|
||||||
|
torch.nn.init.xavier_uniform_(self.fc_es_back.weight)
|
||||||
|
self.fc_is_back = nn.Linear(hidden_size, hidden_size)
|
||||||
|
torch.nn.init.xavier_uniform_(self.fc_is_back.weight)
|
||||||
|
self.fc_indi = nn.Linear(hidden_size, hidden_size)
|
||||||
|
torch.nn.init.xavier_uniform_(self.fc_indi.weight)
|
||||||
|
|
||||||
|
self.leaky_relu = nn.LeakyReLU()
|
||||||
|
self.softmax_s2t = torch.nn.Softmax(dim=0)
|
||||||
|
self.softmax_t2s = torch.nn.Softmax(dim=1)
|
||||||
|
|
||||||
|
self.fc_out_es = nn.Linear(hidden_size, 1)
|
||||||
|
self.fc_out_is = nn.Linear(hidden_size, 1)
|
||||||
|
self.fc_out_indi = nn.Linear(hidden_size, 1)
|
||||||
|
self.fc_out = nn.Linear(hidden_size, 1)
|
||||||
|
|
||||||
|
def cal_cos_similarity(self, x, y): # the 2nd dimension of x and y are the same
|
||||||
|
xy = x.mm(torch.t(y))
|
||||||
|
x_norm = torch.sqrt(torch.sum(x * x, dim=1)).reshape(-1, 1)
|
||||||
|
y_norm = torch.sqrt(torch.sum(y * y, dim=1)).reshape(-1, 1)
|
||||||
|
cos_similarity = xy / (x_norm.mm(torch.t(y_norm)) + 1e-6)
|
||||||
|
return cos_similarity
|
||||||
|
|
||||||
|
def forward(self, x, concept_matrix):
|
||||||
|
device = torch.device(torch.get_device(x))
|
||||||
|
|
||||||
|
x_hidden = x.reshape(len(x), self.d_feat, -1) # [N, F, T]
|
||||||
|
x_hidden = x_hidden.permute(0, 2, 1) # [N, T, F]
|
||||||
|
x_hidden, _ = self.rnn(x_hidden)
|
||||||
|
x_hidden = x_hidden[:, -1, :]
|
||||||
|
|
||||||
|
# Predefined Concept Module
|
||||||
|
|
||||||
|
stock_to_concept = concept_matrix
|
||||||
|
|
||||||
|
stock_to_concept_sum = torch.sum(stock_to_concept, 0).reshape(1, -1).repeat(stock_to_concept.shape[0], 1)
|
||||||
|
stock_to_concept_sum = stock_to_concept_sum.mul(concept_matrix)
|
||||||
|
|
||||||
|
stock_to_concept_sum = stock_to_concept_sum + (
|
||||||
|
torch.ones(stock_to_concept.shape[0], stock_to_concept.shape[1]).to(device)
|
||||||
|
)
|
||||||
|
stock_to_concept = stock_to_concept / stock_to_concept_sum
|
||||||
|
hidden = torch.t(stock_to_concept).mm(x_hidden)
|
||||||
|
|
||||||
|
hidden = hidden[hidden.sum(1) != 0]
|
||||||
|
|
||||||
|
concept_to_stock = self.cal_cos_similarity(x_hidden, hidden)
|
||||||
|
concept_to_stock = self.softmax_t2s(concept_to_stock)
|
||||||
|
|
||||||
|
e_shared_info = concept_to_stock.mm(hidden)
|
||||||
|
e_shared_info = self.fc_es(e_shared_info)
|
||||||
|
|
||||||
|
e_shared_back = self.fc_es_back(e_shared_info)
|
||||||
|
output_es = self.fc_es_fore(e_shared_info)
|
||||||
|
output_es = self.leaky_relu(output_es)
|
||||||
|
|
||||||
|
# Hidden Concept Module
|
||||||
|
i_shared_info = x_hidden - e_shared_back
|
||||||
|
hidden = i_shared_info
|
||||||
|
i_stock_to_concept = self.cal_cos_similarity(i_shared_info, hidden)
|
||||||
|
dim = i_stock_to_concept.shape[0]
|
||||||
|
diag = i_stock_to_concept.diagonal(0)
|
||||||
|
i_stock_to_concept = i_stock_to_concept * (torch.ones(dim, dim) - torch.eye(dim)).to(device)
|
||||||
|
row = torch.linspace(0, dim - 1, dim).to(device).long()
|
||||||
|
column = i_stock_to_concept.max(1)[1].long()
|
||||||
|
value = i_stock_to_concept.max(1)[0]
|
||||||
|
i_stock_to_concept[row, column] = 10
|
||||||
|
i_stock_to_concept[i_stock_to_concept != 10] = 0
|
||||||
|
i_stock_to_concept[row, column] = value
|
||||||
|
i_stock_to_concept = i_stock_to_concept + torch.diag_embed((i_stock_to_concept.sum(0) != 0).float() * diag)
|
||||||
|
hidden = torch.t(i_shared_info).mm(i_stock_to_concept).t()
|
||||||
|
hidden = hidden[hidden.sum(1) != 0]
|
||||||
|
|
||||||
|
i_concept_to_stock = self.cal_cos_similarity(i_shared_info, hidden)
|
||||||
|
i_concept_to_stock = self.softmax_t2s(i_concept_to_stock)
|
||||||
|
i_shared_info = i_concept_to_stock.mm(hidden)
|
||||||
|
i_shared_info = self.fc_is(i_shared_info)
|
||||||
|
|
||||||
|
i_shared_back = self.fc_is_back(i_shared_info)
|
||||||
|
output_is = self.fc_is_fore(i_shared_info)
|
||||||
|
output_is = self.leaky_relu(output_is)
|
||||||
|
|
||||||
|
# Individual Information Module
|
||||||
|
individual_info = x_hidden - e_shared_back - i_shared_back
|
||||||
|
output_indi = individual_info
|
||||||
|
output_indi = self.fc_indi(output_indi)
|
||||||
|
output_indi = self.leaky_relu(output_indi)
|
||||||
|
|
||||||
|
# Stock Trend Prediction
|
||||||
|
all_info = output_es + output_is + output_indi
|
||||||
|
pred_all = self.fc_out(all_info).squeeze()
|
||||||
|
|
||||||
|
return pred_all
|
||||||
446
qlib/contrib/model/pytorch_igmtf.py
Normal file
446
qlib/contrib/model/pytorch_igmtf.py
Normal file
@@ -0,0 +1,446 @@
|
|||||||
|
# Copyright (c) Microsoft Corporation.
|
||||||
|
# Licensed under the MIT License.
|
||||||
|
|
||||||
|
|
||||||
|
from __future__ import division
|
||||||
|
from __future__ import print_function
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pandas as pd
|
||||||
|
from typing import Text, Union
|
||||||
|
import copy
|
||||||
|
from ...utils import get_or_create_path
|
||||||
|
from ...log import get_module_logger
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.optim as optim
|
||||||
|
|
||||||
|
from .pytorch_utils import count_parameters
|
||||||
|
from ...model.base import Model
|
||||||
|
from ...data.dataset import DatasetH
|
||||||
|
from ...data.dataset.handler import DataHandlerLP
|
||||||
|
from ...contrib.model.pytorch_lstm import LSTMModel
|
||||||
|
from ...contrib.model.pytorch_gru import GRUModel
|
||||||
|
|
||||||
|
|
||||||
|
class IGMTF(Model):
|
||||||
|
"""IGMTF Model
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
d_feat : int
|
||||||
|
input dimension for each time step
|
||||||
|
metric: str
|
||||||
|
the evaluate metric used in early stop
|
||||||
|
optimizer : str
|
||||||
|
optimizer name
|
||||||
|
GPU : str
|
||||||
|
the GPU ID(s) used for training
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
d_feat=6,
|
||||||
|
hidden_size=64,
|
||||||
|
num_layers=2,
|
||||||
|
dropout=0.0,
|
||||||
|
n_epochs=200,
|
||||||
|
lr=0.001,
|
||||||
|
metric="",
|
||||||
|
early_stop=20,
|
||||||
|
loss="mse",
|
||||||
|
base_model="GRU",
|
||||||
|
model_path=None,
|
||||||
|
optimizer="adam",
|
||||||
|
GPU=0,
|
||||||
|
seed=None,
|
||||||
|
**kwargs
|
||||||
|
):
|
||||||
|
# Set logger.
|
||||||
|
self.logger = get_module_logger("IGMTF")
|
||||||
|
self.logger.info("IMGTF pytorch version...")
|
||||||
|
|
||||||
|
# set hyper-parameters.
|
||||||
|
self.d_feat = d_feat
|
||||||
|
self.hidden_size = hidden_size
|
||||||
|
self.num_layers = num_layers
|
||||||
|
self.dropout = dropout
|
||||||
|
self.n_epochs = n_epochs
|
||||||
|
self.lr = lr
|
||||||
|
self.metric = metric
|
||||||
|
self.early_stop = early_stop
|
||||||
|
self.optimizer = optimizer.lower()
|
||||||
|
self.loss = loss
|
||||||
|
self.base_model = base_model
|
||||||
|
self.model_path = model_path
|
||||||
|
self.device = torch.device("cuda:%d" % (GPU) if torch.cuda.is_available() and GPU >= 0 else "cpu")
|
||||||
|
self.seed = seed
|
||||||
|
|
||||||
|
self.logger.info(
|
||||||
|
"IGMTF parameters setting:"
|
||||||
|
"\nd_feat : {}"
|
||||||
|
"\nhidden_size : {}"
|
||||||
|
"\nnum_layers : {}"
|
||||||
|
"\ndropout : {}"
|
||||||
|
"\nn_epochs : {}"
|
||||||
|
"\nlr : {}"
|
||||||
|
"\nmetric : {}"
|
||||||
|
"\nearly_stop : {}"
|
||||||
|
"\noptimizer : {}"
|
||||||
|
"\nloss_type : {}"
|
||||||
|
"\nbase_model : {}"
|
||||||
|
"\nmodel_path : {}"
|
||||||
|
"\nvisible_GPU : {}"
|
||||||
|
"\nuse_GPU : {}"
|
||||||
|
"\nseed : {}".format(
|
||||||
|
d_feat,
|
||||||
|
hidden_size,
|
||||||
|
num_layers,
|
||||||
|
dropout,
|
||||||
|
n_epochs,
|
||||||
|
lr,
|
||||||
|
metric,
|
||||||
|
early_stop,
|
||||||
|
optimizer.lower(),
|
||||||
|
loss,
|
||||||
|
base_model,
|
||||||
|
model_path,
|
||||||
|
GPU,
|
||||||
|
self.use_gpu,
|
||||||
|
seed,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.seed is not None:
|
||||||
|
np.random.seed(self.seed)
|
||||||
|
torch.manual_seed(self.seed)
|
||||||
|
|
||||||
|
self.igmtf_model = IGMTFModel(
|
||||||
|
d_feat=self.d_feat,
|
||||||
|
hidden_size=self.hidden_size,
|
||||||
|
num_layers=self.num_layers,
|
||||||
|
dropout=self.dropout,
|
||||||
|
base_model=self.base_model,
|
||||||
|
)
|
||||||
|
self.logger.info("model:\n{:}".format(self.igmtf_model))
|
||||||
|
self.logger.info("model size: {:.4f} MB".format(count_parameters(self.igmtf_model)))
|
||||||
|
|
||||||
|
if optimizer.lower() == "adam":
|
||||||
|
self.train_optimizer = optim.Adam(self.igmtf_model.parameters(), lr=self.lr)
|
||||||
|
elif optimizer.lower() == "gd":
|
||||||
|
self.train_optimizer = optim.SGD(self.igmtf_model.parameters(), lr=self.lr)
|
||||||
|
else:
|
||||||
|
raise NotImplementedError("optimizer {} is not supported!".format(optimizer))
|
||||||
|
|
||||||
|
self.fitted = False
|
||||||
|
self.igmtf_model.to(self.device)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def use_gpu(self):
|
||||||
|
return self.device != torch.device("cpu")
|
||||||
|
|
||||||
|
def mse(self, pred, label):
|
||||||
|
loss = (pred - label) ** 2
|
||||||
|
return torch.mean(loss)
|
||||||
|
|
||||||
|
def loss_fn(self, pred, label):
|
||||||
|
mask = ~torch.isnan(label)
|
||||||
|
|
||||||
|
if self.loss == "mse":
|
||||||
|
return self.mse(pred[mask], label[mask])
|
||||||
|
|
||||||
|
raise ValueError("unknown loss `%s`" % self.loss)
|
||||||
|
|
||||||
|
def metric_fn(self, pred, label):
|
||||||
|
|
||||||
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
|
if self.metric == "ic":
|
||||||
|
x = pred[mask]
|
||||||
|
y = label[mask]
|
||||||
|
|
||||||
|
vx = x - torch.mean(x)
|
||||||
|
vy = y - torch.mean(y)
|
||||||
|
return torch.sum(vx * vy) / (torch.sqrt(torch.sum(vx**2)) * torch.sqrt(torch.sum(vy**2)))
|
||||||
|
|
||||||
|
if self.metric == ("", "loss"):
|
||||||
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
|
|
||||||
|
def get_daily_inter(self, df, shuffle=False):
|
||||||
|
# organize the train data into daily batches
|
||||||
|
daily_count = df.groupby(level=0).size().values
|
||||||
|
daily_index = np.roll(np.cumsum(daily_count), 1)
|
||||||
|
daily_index[0] = 0
|
||||||
|
if shuffle:
|
||||||
|
# shuffle data
|
||||||
|
daily_shuffle = list(zip(daily_index, daily_count))
|
||||||
|
np.random.shuffle(daily_shuffle)
|
||||||
|
daily_index, daily_count = zip(*daily_shuffle)
|
||||||
|
return daily_index, daily_count
|
||||||
|
|
||||||
|
def get_train_hidden(self, x_train):
|
||||||
|
x_train_values = x_train.values
|
||||||
|
daily_index, daily_count = self.get_daily_inter(x_train, shuffle=True)
|
||||||
|
self.igmtf_model.eval()
|
||||||
|
train_hidden = []
|
||||||
|
train_hidden_day = []
|
||||||
|
|
||||||
|
for idx, count in zip(daily_index, daily_count):
|
||||||
|
batch = slice(idx, idx + count)
|
||||||
|
feature = torch.from_numpy(x_train_values[batch]).float().to(self.device)
|
||||||
|
out = self.igmtf_model(feature, get_hidden=True)
|
||||||
|
train_hidden.append(out.detach().cpu())
|
||||||
|
train_hidden_day.append(out.detach().cpu().mean(dim=0).unsqueeze(dim=0))
|
||||||
|
|
||||||
|
train_hidden = np.asarray(train_hidden, dtype=object)
|
||||||
|
train_hidden_day = torch.cat(train_hidden_day)
|
||||||
|
|
||||||
|
return train_hidden, train_hidden_day
|
||||||
|
|
||||||
|
def train_epoch(self, x_train, y_train, train_hidden, train_hidden_day):
|
||||||
|
|
||||||
|
x_train_values = x_train.values
|
||||||
|
y_train_values = np.squeeze(y_train.values)
|
||||||
|
|
||||||
|
self.igmtf_model.train()
|
||||||
|
|
||||||
|
daily_index, daily_count = self.get_daily_inter(x_train, shuffle=True)
|
||||||
|
|
||||||
|
for idx, count in zip(daily_index, daily_count):
|
||||||
|
batch = slice(idx, idx + count)
|
||||||
|
feature = torch.from_numpy(x_train_values[batch]).float().to(self.device)
|
||||||
|
label = torch.from_numpy(y_train_values[batch]).float().to(self.device)
|
||||||
|
pred = self.igmtf_model(feature, train_hidden=train_hidden, train_hidden_day=train_hidden_day)
|
||||||
|
loss = self.loss_fn(pred, label)
|
||||||
|
|
||||||
|
self.train_optimizer.zero_grad()
|
||||||
|
loss.backward()
|
||||||
|
torch.nn.utils.clip_grad_value_(self.igmtf_model.parameters(), 3.0)
|
||||||
|
self.train_optimizer.step()
|
||||||
|
|
||||||
|
def test_epoch(self, data_x, data_y, train_hidden, train_hidden_day):
|
||||||
|
|
||||||
|
# prepare training data
|
||||||
|
x_values = data_x.values
|
||||||
|
y_values = np.squeeze(data_y.values)
|
||||||
|
|
||||||
|
self.igmtf_model.eval()
|
||||||
|
|
||||||
|
scores = []
|
||||||
|
losses = []
|
||||||
|
|
||||||
|
daily_index, daily_count = self.get_daily_inter(data_x, shuffle=False)
|
||||||
|
|
||||||
|
for idx, count in zip(daily_index, daily_count):
|
||||||
|
batch = slice(idx, idx + count)
|
||||||
|
feature = torch.from_numpy(x_values[batch]).float().to(self.device)
|
||||||
|
label = torch.from_numpy(y_values[batch]).float().to(self.device)
|
||||||
|
|
||||||
|
pred = self.igmtf_model(feature, train_hidden=train_hidden, train_hidden_day=train_hidden_day)
|
||||||
|
loss = self.loss_fn(pred, label)
|
||||||
|
losses.append(loss.item())
|
||||||
|
|
||||||
|
score = self.metric_fn(pred, label)
|
||||||
|
scores.append(score.item())
|
||||||
|
|
||||||
|
return np.mean(losses), np.mean(scores)
|
||||||
|
|
||||||
|
def fit(
|
||||||
|
self,
|
||||||
|
dataset: DatasetH,
|
||||||
|
evals_result=dict(),
|
||||||
|
save_path=None,
|
||||||
|
):
|
||||||
|
|
||||||
|
df_train, df_valid = dataset.prepare(
|
||||||
|
["train", "valid"],
|
||||||
|
col_set=["feature", "label"],
|
||||||
|
data_key=DataHandlerLP.DK_L,
|
||||||
|
)
|
||||||
|
if df_train.empty or df_valid.empty:
|
||||||
|
raise ValueError("Empty data from dataset, please check your dataset config.")
|
||||||
|
|
||||||
|
x_train, y_train = df_train["feature"], df_train["label"]
|
||||||
|
x_valid, y_valid = df_valid["feature"], df_valid["label"]
|
||||||
|
|
||||||
|
save_path = get_or_create_path(save_path)
|
||||||
|
stop_steps = 0
|
||||||
|
train_loss = 0
|
||||||
|
best_score = -np.inf
|
||||||
|
best_epoch = 0
|
||||||
|
evals_result["train"] = []
|
||||||
|
evals_result["valid"] = []
|
||||||
|
|
||||||
|
# load pretrained base_model
|
||||||
|
if self.base_model == "LSTM":
|
||||||
|
pretrained_model = LSTMModel()
|
||||||
|
elif self.base_model == "GRU":
|
||||||
|
pretrained_model = GRUModel()
|
||||||
|
else:
|
||||||
|
raise ValueError("unknown base model name `%s`" % self.base_model)
|
||||||
|
|
||||||
|
if self.model_path is not None:
|
||||||
|
self.logger.info("Loading pretrained model...")
|
||||||
|
pretrained_model.load_state_dict(torch.load(self.model_path, map_location=self.device))
|
||||||
|
|
||||||
|
model_dict = self.igmtf_model.state_dict()
|
||||||
|
pretrained_dict = {
|
||||||
|
k: v for k, v in pretrained_model.state_dict().items() if k in model_dict # pylint: disable=E1135
|
||||||
|
}
|
||||||
|
model_dict.update(pretrained_dict)
|
||||||
|
self.igmtf_model.load_state_dict(model_dict)
|
||||||
|
self.logger.info("Loading pretrained model Done...")
|
||||||
|
|
||||||
|
# train
|
||||||
|
self.logger.info("training...")
|
||||||
|
self.fitted = True
|
||||||
|
|
||||||
|
for step in range(self.n_epochs):
|
||||||
|
self.logger.info("Epoch%d:", step)
|
||||||
|
self.logger.info("training...")
|
||||||
|
train_hidden, train_hidden_day = self.get_train_hidden(x_train)
|
||||||
|
self.train_epoch(x_train, y_train, train_hidden, train_hidden_day)
|
||||||
|
self.logger.info("evaluating...")
|
||||||
|
train_loss, train_score = self.test_epoch(x_train, y_train, train_hidden, train_hidden_day)
|
||||||
|
val_loss, val_score = self.test_epoch(x_valid, y_valid, train_hidden, train_hidden_day)
|
||||||
|
self.logger.info("train %.6f, valid %.6f" % (train_score, val_score))
|
||||||
|
evals_result["train"].append(train_score)
|
||||||
|
evals_result["valid"].append(val_score)
|
||||||
|
|
||||||
|
if val_score > best_score:
|
||||||
|
best_score = val_score
|
||||||
|
stop_steps = 0
|
||||||
|
best_epoch = step
|
||||||
|
best_param = copy.deepcopy(self.igmtf_model.state_dict())
|
||||||
|
else:
|
||||||
|
stop_steps += 1
|
||||||
|
if stop_steps >= self.early_stop:
|
||||||
|
self.logger.info("early stop")
|
||||||
|
break
|
||||||
|
|
||||||
|
self.logger.info("best score: %.6lf @ %d" % (best_score, best_epoch))
|
||||||
|
self.igmtf_model.load_state_dict(best_param)
|
||||||
|
torch.save(best_param, save_path)
|
||||||
|
|
||||||
|
if self.use_gpu:
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
def predict(self, dataset: DatasetH, segment: Union[Text, slice] = "test"):
|
||||||
|
if not self.fitted:
|
||||||
|
raise ValueError("model is not fitted yet!")
|
||||||
|
x_train = dataset.prepare("train", col_set="feature", data_key=DataHandlerLP.DK_L)
|
||||||
|
train_hidden, train_hidden_day = self.get_train_hidden(x_train)
|
||||||
|
x_test = dataset.prepare(segment, col_set="feature", data_key=DataHandlerLP.DK_I)
|
||||||
|
index = x_test.index
|
||||||
|
self.igmtf_model.eval()
|
||||||
|
x_values = x_test.values
|
||||||
|
preds = []
|
||||||
|
|
||||||
|
daily_index, daily_count = self.get_daily_inter(x_test, shuffle=False)
|
||||||
|
|
||||||
|
for idx, count in zip(daily_index, daily_count):
|
||||||
|
batch = slice(idx, idx + count)
|
||||||
|
x_batch = torch.from_numpy(x_values[batch]).float().to(self.device)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
pred = (
|
||||||
|
self.igmtf_model(x_batch, train_hidden=train_hidden, train_hidden_day=train_hidden_day)
|
||||||
|
.detach()
|
||||||
|
.cpu()
|
||||||
|
.numpy()
|
||||||
|
)
|
||||||
|
|
||||||
|
preds.append(pred)
|
||||||
|
|
||||||
|
return pd.Series(np.concatenate(preds), index=index)
|
||||||
|
|
||||||
|
|
||||||
|
class IGMTFModel(nn.Module):
|
||||||
|
def __init__(self, d_feat=6, hidden_size=64, num_layers=2, dropout=0.0, base_model="GRU"):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
if base_model == "GRU":
|
||||||
|
self.rnn = nn.GRU(
|
||||||
|
input_size=d_feat,
|
||||||
|
hidden_size=hidden_size,
|
||||||
|
num_layers=num_layers,
|
||||||
|
batch_first=True,
|
||||||
|
dropout=dropout,
|
||||||
|
)
|
||||||
|
elif base_model == "LSTM":
|
||||||
|
self.rnn = nn.LSTM(
|
||||||
|
input_size=d_feat,
|
||||||
|
hidden_size=hidden_size,
|
||||||
|
num_layers=num_layers,
|
||||||
|
batch_first=True,
|
||||||
|
dropout=dropout,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise ValueError("unknown base model name `%s`" % base_model)
|
||||||
|
self.lins = nn.Sequential()
|
||||||
|
for i in range(2):
|
||||||
|
self.lins.add_module("linear" + str(i), nn.Linear(hidden_size, hidden_size))
|
||||||
|
self.lins.add_module("leakyrelu" + str(i), nn.LeakyReLU())
|
||||||
|
self.fc_output = nn.Linear(hidden_size * 2, hidden_size * 2)
|
||||||
|
self.project1 = nn.Linear(hidden_size, hidden_size, bias=False)
|
||||||
|
self.project2 = nn.Linear(hidden_size, hidden_size, bias=False)
|
||||||
|
self.fc_out_pred = nn.Linear(hidden_size * 2, 1)
|
||||||
|
|
||||||
|
self.leaky_relu = nn.LeakyReLU()
|
||||||
|
self.d_feat = d_feat
|
||||||
|
|
||||||
|
def cal_cos_similarity(self, x, y): # the 2nd dimension of x and y are the same
|
||||||
|
xy = x.mm(torch.t(y))
|
||||||
|
x_norm = torch.sqrt(torch.sum(x * x, dim=1)).reshape(-1, 1)
|
||||||
|
y_norm = torch.sqrt(torch.sum(y * y, dim=1)).reshape(-1, 1)
|
||||||
|
cos_similarity = xy / (x_norm.mm(torch.t(y_norm)) + 1e-6)
|
||||||
|
return cos_similarity
|
||||||
|
|
||||||
|
def sparse_dense_mul(self, s, d):
|
||||||
|
i = s._indices()
|
||||||
|
v = s._values()
|
||||||
|
dv = d[i[0, :], i[1, :]] # get values from relevant entries of dense matrix
|
||||||
|
return torch.sparse.FloatTensor(i, v * dv, s.size())
|
||||||
|
|
||||||
|
def forward(self, x, get_hidden=False, train_hidden=None, train_hidden_day=None, k_day=10, n_neighbor=10):
|
||||||
|
# x: [N, F*T]
|
||||||
|
device = x.device
|
||||||
|
x = x.reshape(len(x), self.d_feat, -1) # [N, F, T]
|
||||||
|
x = x.permute(0, 2, 1) # [N, T, F]
|
||||||
|
out, _ = self.rnn(x)
|
||||||
|
out = out[:, -1, :]
|
||||||
|
out = self.lins(out)
|
||||||
|
mini_batch_out = out
|
||||||
|
if get_hidden is True:
|
||||||
|
return mini_batch_out
|
||||||
|
|
||||||
|
mini_batch_out_day = torch.mean(mini_batch_out, dim=0).unsqueeze(0)
|
||||||
|
day_similarity = self.cal_cos_similarity(mini_batch_out_day, train_hidden_day.to(device))
|
||||||
|
day_index = torch.topk(day_similarity, k_day, dim=1)[1]
|
||||||
|
sample_train_hidden = train_hidden[day_index.long().cpu()].squeeze()
|
||||||
|
sample_train_hidden = torch.cat(list(sample_train_hidden)).to(device)
|
||||||
|
sample_train_hidden = self.lins(sample_train_hidden)
|
||||||
|
cos_similarity = self.cal_cos_similarity(self.project1(mini_batch_out), self.project2(sample_train_hidden))
|
||||||
|
|
||||||
|
row = (
|
||||||
|
torch.linspace(0, x.shape[0] - 1, x.shape[0])
|
||||||
|
.reshape([-1, 1])
|
||||||
|
.repeat(1, n_neighbor)
|
||||||
|
.reshape(1, -1)
|
||||||
|
.to(device)
|
||||||
|
)
|
||||||
|
column = torch.topk(cos_similarity, n_neighbor, dim=1)[1].reshape(1, -1)
|
||||||
|
mask = torch.sparse_coo_tensor(
|
||||||
|
torch.cat([row, column]),
|
||||||
|
torch.ones([row.shape[1]]).to(device) / n_neighbor,
|
||||||
|
(x.shape[0], sample_train_hidden.shape[0]),
|
||||||
|
)
|
||||||
|
cos_similarity = self.sparse_dense_mul(mask, cos_similarity)
|
||||||
|
|
||||||
|
agg_out = torch.sparse.mm(cos_similarity, self.project2(sample_train_hidden))
|
||||||
|
# out = self.fc_out(out).squeeze()
|
||||||
|
out = self.fc_out_pred(torch.cat([mini_batch_out, agg_out], axis=1)).squeeze()
|
||||||
|
return out
|
||||||
@@ -5,7 +5,6 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from typing import Text, Union
|
from typing import Text, Union
|
||||||
@@ -17,11 +16,9 @@ from ...log import get_module_logger
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.optim as optim
|
import torch.optim as optim
|
||||||
from torch.utils.data import DataLoader
|
|
||||||
|
|
||||||
from .pytorch_utils import count_parameters
|
|
||||||
from ...model.base import Model
|
from ...model.base import Model
|
||||||
from ...data.dataset import DatasetH, TSDatasetH
|
from ...data.dataset import DatasetH
|
||||||
from ...data.dataset.handler import DataHandlerLP
|
from ...data.dataset.handler import DataHandlerLP
|
||||||
from torch.nn.modules.container import ModuleList
|
from torch.nn.modules.container import ModuleList
|
||||||
|
|
||||||
@@ -102,7 +99,7 @@ class LocalformerModel(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
|
|||||||
@@ -5,7 +5,6 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import copy
|
import copy
|
||||||
@@ -18,9 +17,8 @@ import torch.nn as nn
|
|||||||
import torch.optim as optim
|
import torch.optim as optim
|
||||||
from torch.utils.data import DataLoader
|
from torch.utils.data import DataLoader
|
||||||
|
|
||||||
from .pytorch_utils import count_parameters
|
|
||||||
from ...model.base import Model
|
from ...model.base import Model
|
||||||
from ...data.dataset import DatasetH, TSDatasetH
|
from ...data.dataset import DatasetH
|
||||||
from ...data.dataset.handler import DataHandlerLP
|
from ...data.dataset.handler import DataHandlerLP
|
||||||
from torch.nn.modules.container import ModuleList
|
from torch.nn.modules.container import ModuleList
|
||||||
|
|
||||||
@@ -101,7 +99,7 @@ class LocalformerModel(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
|
|||||||
@@ -5,7 +5,6 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from typing import Text, Union
|
from typing import Text, Union
|
||||||
@@ -146,7 +145,7 @@ class LSTM(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
|
|||||||
@@ -5,7 +5,6 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import copy
|
import copy
|
||||||
@@ -18,7 +17,6 @@ import torch.optim as optim
|
|||||||
from torch.utils.data import DataLoader
|
from torch.utils.data import DataLoader
|
||||||
|
|
||||||
from ...model.base import Model
|
from ...model.base import Model
|
||||||
from ...data.dataset import DatasetH, TSDatasetH
|
|
||||||
from ...data.dataset.handler import DataHandlerLP
|
from ...data.dataset.handler import DataHandlerLP
|
||||||
from ...model.utils import ConcatDataset
|
from ...model.utils import ConcatDataset
|
||||||
from ...data.dataset.weight import Reweighter
|
from ...data.dataset.weight import Reweighter
|
||||||
@@ -140,7 +138,7 @@ class LSTM(Model):
|
|||||||
loss = weight * (pred - label) ** 2
|
loss = weight * (pred - label) ** 2
|
||||||
return torch.mean(loss)
|
return torch.mean(loss)
|
||||||
|
|
||||||
def loss_fn(self, pred, label):
|
def loss_fn(self, pred, label, weight):
|
||||||
mask = ~torch.isnan(label)
|
mask = ~torch.isnan(label)
|
||||||
|
|
||||||
if weight is None:
|
if weight is None:
|
||||||
@@ -155,8 +153,8 @@ class LSTM(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask], weight=None)
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
|
|
||||||
|
|||||||
@@ -4,11 +4,13 @@
|
|||||||
|
|
||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
from collections import defaultdict
|
||||||
|
|
||||||
import os
|
import os
|
||||||
|
import gc
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from typing import Text, Union
|
from typing import Callable, Optional, Text, Union
|
||||||
from sklearn.metrics import roc_auc_score, mean_squared_error
|
from sklearn.metrics import roc_auc_score, mean_squared_error
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -20,9 +22,17 @@ from ...model.base import Model
|
|||||||
from ...data.dataset import DatasetH
|
from ...data.dataset import DatasetH
|
||||||
from ...data.dataset.handler import DataHandlerLP
|
from ...data.dataset.handler import DataHandlerLP
|
||||||
from ...data.dataset.weight import Reweighter
|
from ...data.dataset.weight import Reweighter
|
||||||
from ...utils import unpack_archive_with_buffer, save_multiple_parts_file, get_or_create_path
|
from ...utils import (
|
||||||
|
auto_filter_kwargs,
|
||||||
|
init_instance_by_config,
|
||||||
|
unpack_archive_with_buffer,
|
||||||
|
save_multiple_parts_file,
|
||||||
|
get_or_create_path,
|
||||||
|
)
|
||||||
from ...log import get_module_logger
|
from ...log import get_module_logger
|
||||||
from ...workflow import R
|
from ...workflow import R
|
||||||
|
from qlib.contrib.meta.data_selection.utils import ICLoss
|
||||||
|
from torch.nn import DataParallel
|
||||||
|
|
||||||
|
|
||||||
class DNNModelPytorch(Model):
|
class DNNModelPytorch(Model):
|
||||||
@@ -49,9 +59,6 @@ class DNNModelPytorch(Model):
|
|||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
input_dim=360,
|
|
||||||
output_dim=1,
|
|
||||||
layers=(256,),
|
|
||||||
lr=0.001,
|
lr=0.001,
|
||||||
max_steps=300,
|
max_steps=300,
|
||||||
batch_size=2000,
|
batch_size=2000,
|
||||||
@@ -64,14 +71,23 @@ class DNNModelPytorch(Model):
|
|||||||
GPU=0,
|
GPU=0,
|
||||||
seed=None,
|
seed=None,
|
||||||
weight_decay=0.0,
|
weight_decay=0.0,
|
||||||
**kwargs
|
data_parall=False,
|
||||||
|
scheduler: Optional[Union[Callable]] = "default", # when it is Callable, it accept one argument named optimizer
|
||||||
|
init_model=None,
|
||||||
|
eval_train_metric=False,
|
||||||
|
pt_model_uri="qlib.contrib.model.pytorch_nn.Net",
|
||||||
|
pt_model_kwargs={
|
||||||
|
"input_dim": 360,
|
||||||
|
"layers": (256,),
|
||||||
|
},
|
||||||
|
valid_key=DataHandlerLP.DK_L,
|
||||||
|
# TODO: Infer Key is a more reasonable key. But it requires more detailed processing on label processing
|
||||||
):
|
):
|
||||||
# Set logger.
|
# Set logger.
|
||||||
self.logger = get_module_logger("DNNModelPytorch")
|
self.logger = get_module_logger("DNNModelPytorch")
|
||||||
self.logger.info("DNN pytorch version...")
|
self.logger.info("DNN pytorch version...")
|
||||||
|
|
||||||
# set hyper-parameters.
|
# set hyper-parameters.
|
||||||
self.layers = layers
|
|
||||||
self.lr = lr
|
self.lr = lr
|
||||||
self.max_steps = max_steps
|
self.max_steps = max_steps
|
||||||
self.batch_size = batch_size
|
self.batch_size = batch_size
|
||||||
@@ -81,41 +97,36 @@ class DNNModelPytorch(Model):
|
|||||||
self.lr_decay_steps = lr_decay_steps
|
self.lr_decay_steps = lr_decay_steps
|
||||||
self.optimizer = optimizer.lower()
|
self.optimizer = optimizer.lower()
|
||||||
self.loss_type = loss
|
self.loss_type = loss
|
||||||
self.device = torch.device("cuda:%d" % (GPU) if torch.cuda.is_available() and GPU >= 0 else "cpu")
|
if isinstance(GPU, str):
|
||||||
|
self.device = torch.device(GPU)
|
||||||
|
else:
|
||||||
|
self.device = torch.device("cuda:%d" % (GPU) if torch.cuda.is_available() and GPU >= 0 else "cpu")
|
||||||
self.seed = seed
|
self.seed = seed
|
||||||
self.weight_decay = weight_decay
|
self.weight_decay = weight_decay
|
||||||
|
self.data_parall = data_parall
|
||||||
|
self.eval_train_metric = eval_train_metric
|
||||||
|
self.valid_key = valid_key
|
||||||
|
|
||||||
|
self.best_step = None
|
||||||
|
|
||||||
self.logger.info(
|
self.logger.info(
|
||||||
"DNN parameters setting:"
|
"DNN parameters setting:"
|
||||||
"\nlayers : {}"
|
f"\nlr : {lr}"
|
||||||
"\nlr : {}"
|
f"\nmax_steps : {max_steps}"
|
||||||
"\nmax_steps : {}"
|
f"\nbatch_size : {batch_size}"
|
||||||
"\nbatch_size : {}"
|
f"\nearly_stop_rounds : {early_stop_rounds}"
|
||||||
"\nearly_stop_rounds : {}"
|
f"\neval_steps : {eval_steps}"
|
||||||
"\neval_steps : {}"
|
f"\nlr_decay : {lr_decay}"
|
||||||
"\nlr_decay : {}"
|
f"\nlr_decay_steps : {lr_decay_steps}"
|
||||||
"\nlr_decay_steps : {}"
|
f"\noptimizer : {optimizer}"
|
||||||
"\noptimizer : {}"
|
f"\nloss_type : {loss}"
|
||||||
"\nloss_type : {}"
|
f"\nseed : {seed}"
|
||||||
"\nseed : {}"
|
f"\ndevice : {self.device}"
|
||||||
"\ndevice : {}"
|
f"\nuse_GPU : {self.use_gpu}"
|
||||||
"\nuse_GPU : {}"
|
f"\nweight_decay : {weight_decay}"
|
||||||
"\nweight_decay : {}".format(
|
f"\nenable data parall : {self.data_parall}"
|
||||||
layers,
|
f"\npt_model_uri: {pt_model_uri}"
|
||||||
lr,
|
f"\npt_model_kwargs: {pt_model_kwargs}"
|
||||||
max_steps,
|
|
||||||
batch_size,
|
|
||||||
early_stop_rounds,
|
|
||||||
eval_steps,
|
|
||||||
lr_decay,
|
|
||||||
lr_decay_steps,
|
|
||||||
optimizer,
|
|
||||||
loss,
|
|
||||||
seed,
|
|
||||||
self.device,
|
|
||||||
self.use_gpu,
|
|
||||||
weight_decay,
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.seed is not None:
|
if self.seed is not None:
|
||||||
@@ -126,7 +137,14 @@ class DNNModelPytorch(Model):
|
|||||||
raise NotImplementedError("loss {} is not supported!".format(loss))
|
raise NotImplementedError("loss {} is not supported!".format(loss))
|
||||||
self._scorer = mean_squared_error if loss == "mse" else roc_auc_score
|
self._scorer = mean_squared_error if loss == "mse" else roc_auc_score
|
||||||
|
|
||||||
self.dnn_model = Net(input_dim, output_dim, layers, loss=self.loss_type)
|
if init_model is None:
|
||||||
|
self.dnn_model = init_instance_by_config({"class": pt_model_uri, "kwargs": pt_model_kwargs})
|
||||||
|
|
||||||
|
if self.data_parall:
|
||||||
|
self.dnn_model = DataParallel(self.dnn_model).to(self.device)
|
||||||
|
else:
|
||||||
|
self.dnn_model = init_model
|
||||||
|
|
||||||
self.logger.info("model:\n{:}".format(self.dnn_model))
|
self.logger.info("model:\n{:}".format(self.dnn_model))
|
||||||
self.logger.info("model size: {:.4f} MB".format(count_parameters(self.dnn_model)))
|
self.logger.info("model size: {:.4f} MB".format(count_parameters(self.dnn_model)))
|
||||||
|
|
||||||
@@ -137,19 +155,24 @@ class DNNModelPytorch(Model):
|
|||||||
else:
|
else:
|
||||||
raise NotImplementedError("optimizer {} is not supported!".format(optimizer))
|
raise NotImplementedError("optimizer {} is not supported!".format(optimizer))
|
||||||
|
|
||||||
# Reduce learning rate when loss has stopped decrease
|
if scheduler == "default":
|
||||||
self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
|
# Reduce learning rate when loss has stopped decrease
|
||||||
self.train_optimizer,
|
self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
|
||||||
mode="min",
|
self.train_optimizer,
|
||||||
factor=0.5,
|
mode="min",
|
||||||
patience=10,
|
factor=0.5,
|
||||||
verbose=True,
|
patience=10,
|
||||||
threshold=0.0001,
|
verbose=True,
|
||||||
threshold_mode="rel",
|
threshold=0.0001,
|
||||||
cooldown=0,
|
threshold_mode="rel",
|
||||||
min_lr=0.00001,
|
cooldown=0,
|
||||||
eps=1e-08,
|
min_lr=0.00001,
|
||||||
)
|
eps=1e-08,
|
||||||
|
)
|
||||||
|
elif scheduler is None:
|
||||||
|
self.scheduler = None
|
||||||
|
else:
|
||||||
|
self.scheduler = scheduler(optimizer=self.train_optimizer)
|
||||||
|
|
||||||
self.fitted = False
|
self.fitted = False
|
||||||
self.dnn_model.to(self.device)
|
self.dnn_model.to(self.device)
|
||||||
@@ -166,40 +189,48 @@ class DNNModelPytorch(Model):
|
|||||||
save_path=None,
|
save_path=None,
|
||||||
reweighter=None,
|
reweighter=None,
|
||||||
):
|
):
|
||||||
df_train, df_valid = dataset.prepare(
|
has_valid = "valid" in dataset.segments
|
||||||
["train", "valid"], col_set=["feature", "label"], data_key=DataHandlerLP.DK_L
|
segments = ["train", "valid"]
|
||||||
)
|
vars = ["x", "y", "w"]
|
||||||
x_train, y_train = df_train["feature"], df_train["label"]
|
all_df = defaultdict(dict) # x_train, x_valid y_train, y_valid w_train, w_valid
|
||||||
x_valid, y_valid = df_valid["feature"], df_valid["label"]
|
all_t = defaultdict(dict) # tensors
|
||||||
|
for seg in segments:
|
||||||
|
if seg in dataset.segments:
|
||||||
|
# df_train df_valid
|
||||||
|
df = dataset.prepare(
|
||||||
|
seg, col_set=["feature", "label"], data_key=self.valid_key if seg == "valid" else DataHandlerLP.DK_L
|
||||||
|
)
|
||||||
|
all_df["x"][seg] = df["feature"]
|
||||||
|
all_df["y"][seg] = df["label"].copy() # We have to use copy to remove the reference to release mem
|
||||||
|
if reweighter is None:
|
||||||
|
all_df["w"][seg] = pd.DataFrame(np.ones_like(all_df["y"][seg].values), index=df.index)
|
||||||
|
elif isinstance(reweighter, Reweighter):
|
||||||
|
all_df["w"][seg] = pd.DataFrame(reweighter.reweight(df))
|
||||||
|
else:
|
||||||
|
raise ValueError("Unsupported reweighter type.")
|
||||||
|
|
||||||
if reweighter is None:
|
# get tensors
|
||||||
w_train = pd.DataFrame(np.ones_like(y_train.values), index=y_train.index)
|
for v in vars:
|
||||||
w_valid = pd.DataFrame(np.ones_like(y_valid.values), index=y_valid.index)
|
all_t[v][seg] = torch.from_numpy(all_df[v][seg].values).float()
|
||||||
elif isinstance(reweighter, Reweighter):
|
# if seg == "valid": # accelerate the eval of validation
|
||||||
w_train = pd.DataFrame(reweighter.reweight(df_train))
|
all_t[v][seg] = all_t[v][seg].to(self.device) # This will consume a lot of memory !!!!
|
||||||
w_valid = pd.DataFrame(reweighter.reweight(df_valid))
|
|
||||||
else:
|
evals_result[seg] = []
|
||||||
raise ValueError("Unsupported reweighter type.")
|
# free memory
|
||||||
|
del df
|
||||||
|
del all_df["x"]
|
||||||
|
gc.collect()
|
||||||
|
|
||||||
save_path = get_or_create_path(save_path)
|
save_path = get_or_create_path(save_path)
|
||||||
stop_steps = 0
|
stop_steps = 0
|
||||||
train_loss = 0
|
train_loss = 0
|
||||||
best_loss = np.inf
|
best_loss = np.inf
|
||||||
evals_result["train"] = []
|
|
||||||
evals_result["valid"] = []
|
|
||||||
# train
|
# train
|
||||||
self.logger.info("training...")
|
self.logger.info("training...")
|
||||||
self.fitted = True
|
self.fitted = True
|
||||||
# return
|
# return
|
||||||
# prepare training data
|
# prepare training data
|
||||||
x_train_values = torch.from_numpy(x_train.values).float()
|
train_num = all_t["y"]["train"].shape[0]
|
||||||
y_train_values = torch.from_numpy(y_train.values).float()
|
|
||||||
w_train_values = torch.from_numpy(w_train.values).float()
|
|
||||||
train_num = y_train_values.shape[0]
|
|
||||||
# prepare validation data
|
|
||||||
x_val_auto = torch.from_numpy(x_valid.values).float().to(self.device)
|
|
||||||
y_val_auto = torch.from_numpy(y_valid.values).float().to(self.device)
|
|
||||||
w_val_auto = torch.from_numpy(w_valid.values).float().to(self.device)
|
|
||||||
|
|
||||||
for step in range(1, self.max_steps + 1):
|
for step in range(1, self.max_steps + 1):
|
||||||
if stop_steps >= self.early_stop_rounds:
|
if stop_steps >= self.early_stop_rounds:
|
||||||
@@ -210,9 +241,9 @@ class DNNModelPytorch(Model):
|
|||||||
self.dnn_model.train()
|
self.dnn_model.train()
|
||||||
self.train_optimizer.zero_grad()
|
self.train_optimizer.zero_grad()
|
||||||
choice = np.random.choice(train_num, self.batch_size)
|
choice = np.random.choice(train_num, self.batch_size)
|
||||||
x_batch_auto = x_train_values[choice].to(self.device)
|
x_batch_auto = all_t["x"]["train"][choice].to(self.device)
|
||||||
y_batch_auto = y_train_values[choice].to(self.device)
|
y_batch_auto = all_t["y"]["train"][choice].to(self.device)
|
||||||
w_batch_auto = w_train_values[choice].to(self.device)
|
w_batch_auto = all_t["w"]["train"][choice].to(self.device)
|
||||||
|
|
||||||
# forward
|
# forward
|
||||||
preds = self.dnn_model(x_batch_auto)
|
preds = self.dnn_model(x_batch_auto)
|
||||||
@@ -226,44 +257,84 @@ class DNNModelPytorch(Model):
|
|||||||
train_loss += loss.val
|
train_loss += loss.val
|
||||||
# for evert `eval_steps` steps or at the last steps, we will evaluate the model.
|
# for evert `eval_steps` steps or at the last steps, we will evaluate the model.
|
||||||
if step % self.eval_steps == 0 or step == self.max_steps:
|
if step % self.eval_steps == 0 or step == self.max_steps:
|
||||||
stop_steps += 1
|
if has_valid:
|
||||||
train_loss /= self.eval_steps
|
stop_steps += 1
|
||||||
|
train_loss /= self.eval_steps
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
self.dnn_model.eval()
|
self.dnn_model.eval()
|
||||||
loss_val = AverageMeter()
|
|
||||||
|
|
||||||
# forward
|
# forward
|
||||||
preds = self.dnn_model(x_val_auto)
|
preds = self._nn_predict(all_t["x"]["valid"], return_cpu=False)
|
||||||
cur_loss_val = self.get_loss(preds, w_val_auto, y_val_auto, self.loss_type)
|
cur_loss_val = self.get_loss(preds, all_t["w"]["valid"], all_t["y"]["valid"], self.loss_type)
|
||||||
loss_val.update(cur_loss_val.item())
|
loss_val = cur_loss_val.item()
|
||||||
R.log_metrics(val_loss=loss_val.val, step=step)
|
metric_val = (
|
||||||
if verbose:
|
self.get_metric(
|
||||||
self.logger.info(
|
preds.reshape(-1), all_t["y"]["valid"].reshape(-1), all_df["y"]["valid"].index
|
||||||
"[Epoch {}]: train_loss {:.6f}, valid_loss {:.6f}".format(step, train_loss, loss_val.val)
|
)
|
||||||
)
|
.detach()
|
||||||
evals_result["train"].append(train_loss)
|
.cpu()
|
||||||
evals_result["valid"].append(loss_val.val)
|
.numpy()
|
||||||
if loss_val.val < best_loss:
|
.item()
|
||||||
|
)
|
||||||
|
R.log_metrics(val_loss=loss_val, step=step)
|
||||||
|
R.log_metrics(val_metric=metric_val, step=step)
|
||||||
|
|
||||||
|
if self.eval_train_metric:
|
||||||
|
metric_train = (
|
||||||
|
self.get_metric(
|
||||||
|
self._nn_predict(all_t["x"]["train"], return_cpu=False),
|
||||||
|
all_t["y"]["train"].reshape(-1),
|
||||||
|
all_df["y"]["train"].index,
|
||||||
|
)
|
||||||
|
.detach()
|
||||||
|
.cpu()
|
||||||
|
.numpy()
|
||||||
|
.item()
|
||||||
|
)
|
||||||
|
R.log_metrics(train_metric=metric_train, step=step)
|
||||||
|
else:
|
||||||
|
metric_train = np.nan
|
||||||
if verbose:
|
if verbose:
|
||||||
self.logger.info(
|
self.logger.info(
|
||||||
"\tvalid loss update from {:.6f} to {:.6f}, save checkpoint.".format(
|
f"[Step {step}]: train_loss {train_loss:.6f}, valid_loss {loss_val:.6f}, train_metric {metric_train:.6f}, valid_metric {metric_val:.6f}"
|
||||||
best_loss, loss_val.val
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
best_loss = loss_val.val
|
evals_result["train"].append(train_loss)
|
||||||
stop_steps = 0
|
evals_result["valid"].append(loss_val)
|
||||||
torch.save(self.dnn_model.state_dict(), save_path)
|
if loss_val < best_loss:
|
||||||
train_loss = 0
|
if verbose:
|
||||||
# update learning rate
|
self.logger.info(
|
||||||
self.scheduler.step(cur_loss_val)
|
"\tvalid loss update from {:.6f} to {:.6f}, save checkpoint.".format(
|
||||||
|
best_loss, loss_val
|
||||||
|
)
|
||||||
|
)
|
||||||
|
best_loss = loss_val
|
||||||
|
self.best_step = step
|
||||||
|
R.log_metrics(best_step=self.best_step, step=step)
|
||||||
|
stop_steps = 0
|
||||||
|
torch.save(self.dnn_model.state_dict(), save_path)
|
||||||
|
train_loss = 0
|
||||||
|
# update learning rate
|
||||||
|
if self.scheduler is not None:
|
||||||
|
auto_filter_kwargs(self.scheduler.step, warning=False)(metrics=cur_loss_val, epoch=step)
|
||||||
|
R.log_metrics(lr=self.get_lr(), step=step)
|
||||||
|
else:
|
||||||
|
# retraining mode
|
||||||
|
if self.scheduler is not None:
|
||||||
|
self.scheduler.step(epoch=step)
|
||||||
|
|
||||||
# restore the optimal parameters after training
|
if has_valid:
|
||||||
self.dnn_model.load_state_dict(torch.load(save_path, map_location=self.device))
|
# restore the optimal parameters after training
|
||||||
|
self.dnn_model.load_state_dict(torch.load(save_path, map_location=self.device))
|
||||||
if self.use_gpu:
|
if self.use_gpu:
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
def get_lr(self):
|
||||||
|
assert len(self.train_optimizer.param_groups) == 1
|
||||||
|
return self.train_optimizer.param_groups[0]["lr"]
|
||||||
|
|
||||||
def get_loss(self, pred, w, target, loss_type):
|
def get_loss(self, pred, w, target, loss_type):
|
||||||
|
pred, w, target = pred.reshape(-1), w.reshape(-1), target.reshape(-1)
|
||||||
if loss_type == "mse":
|
if loss_type == "mse":
|
||||||
sqr_loss = torch.mul(pred - target, pred - target)
|
sqr_loss = torch.mul(pred - target, pred - target)
|
||||||
loss = torch.mul(sqr_loss, w).mean()
|
loss = torch.mul(sqr_loss, w).mean()
|
||||||
@@ -274,15 +345,40 @@ class DNNModelPytorch(Model):
|
|||||||
else:
|
else:
|
||||||
raise NotImplementedError("loss {} is not supported!".format(loss_type))
|
raise NotImplementedError("loss {} is not supported!".format(loss_type))
|
||||||
|
|
||||||
|
def get_metric(self, pred, target, index):
|
||||||
|
# NOTE: the order of the index must follow <datetime, instrument> sorted order
|
||||||
|
return -ICLoss()(pred, target, index) # pylint: disable=E1130
|
||||||
|
|
||||||
|
def _nn_predict(self, data, return_cpu=True):
|
||||||
|
"""Reusing predicting NN.
|
||||||
|
Scenarios
|
||||||
|
1) test inference (data may come from CPU and expect the output data is on CPU)
|
||||||
|
2) evaluation on training (data may come from GPU)
|
||||||
|
"""
|
||||||
|
if not isinstance(data, torch.Tensor):
|
||||||
|
if isinstance(data, pd.DataFrame):
|
||||||
|
data = data.values
|
||||||
|
data = torch.Tensor(data)
|
||||||
|
data = data.to(self.device)
|
||||||
|
preds = []
|
||||||
|
self.dnn_model.eval()
|
||||||
|
with torch.no_grad():
|
||||||
|
batch_size = 8096
|
||||||
|
for i in range(0, len(data), batch_size):
|
||||||
|
x = data[i : i + batch_size]
|
||||||
|
preds.append(self.dnn_model(x.to(self.device)).detach().reshape(-1))
|
||||||
|
if return_cpu:
|
||||||
|
preds = np.concatenate([pr.cpu().numpy() for pr in preds])
|
||||||
|
else:
|
||||||
|
preds = torch.cat(preds, axis=0)
|
||||||
|
return preds
|
||||||
|
|
||||||
def predict(self, dataset: DatasetH, segment: Union[Text, slice] = "test"):
|
def predict(self, dataset: DatasetH, segment: Union[Text, slice] = "test"):
|
||||||
if not self.fitted:
|
if not self.fitted:
|
||||||
raise ValueError("model is not fitted yet!")
|
raise ValueError("model is not fitted yet!")
|
||||||
x_test_pd = dataset.prepare(segment, col_set="feature", data_key=DataHandlerLP.DK_I)
|
x_test_pd = dataset.prepare(segment, col_set="feature", data_key=DataHandlerLP.DK_I)
|
||||||
x_test = torch.from_numpy(x_test_pd.values).float().to(self.device)
|
preds = self._nn_predict(x_test_pd)
|
||||||
self.dnn_model.eval()
|
return pd.Series(preds.reshape(-1), index=x_test_pd.index)
|
||||||
with torch.no_grad():
|
|
||||||
preds = self.dnn_model(x_test).detach().cpu().numpy()
|
|
||||||
return pd.Series(np.squeeze(preds), index=x_test_pd.index)
|
|
||||||
|
|
||||||
def save(self, filename, **kwargs):
|
def save(self, filename, **kwargs):
|
||||||
with save_multiple_parts_file(filename) as model_dir:
|
with save_multiple_parts_file(filename) as model_dir:
|
||||||
@@ -322,15 +418,22 @@ class AverageMeter:
|
|||||||
|
|
||||||
|
|
||||||
class Net(nn.Module):
|
class Net(nn.Module):
|
||||||
def __init__(self, input_dim, output_dim, layers=(256, 512, 768, 512, 256, 128, 64), loss="mse"):
|
def __init__(self, input_dim, output_dim=1, layers=(256,), act="LeakyReLU"):
|
||||||
super(Net, self).__init__()
|
super(Net, self).__init__()
|
||||||
|
|
||||||
layers = [input_dim] + list(layers)
|
layers = [input_dim] + list(layers)
|
||||||
dnn_layers = []
|
dnn_layers = []
|
||||||
drop_input = nn.Dropout(0.05)
|
drop_input = nn.Dropout(0.05)
|
||||||
dnn_layers.append(drop_input)
|
dnn_layers.append(drop_input)
|
||||||
|
hidden_units = input_dim
|
||||||
for i, (_input_dim, hidden_units) in enumerate(zip(layers[:-1], layers[1:])):
|
for i, (_input_dim, hidden_units) in enumerate(zip(layers[:-1], layers[1:])):
|
||||||
fc = nn.Linear(_input_dim, hidden_units)
|
fc = nn.Linear(_input_dim, hidden_units)
|
||||||
activation = nn.LeakyReLU(negative_slope=0.1, inplace=False)
|
if act == "LeakyReLU":
|
||||||
|
activation = nn.LeakyReLU(negative_slope=0.1, inplace=False)
|
||||||
|
elif act == "SiLU":
|
||||||
|
activation = nn.SiLU()
|
||||||
|
else:
|
||||||
|
raise NotImplementedError(f"This type of input is not supported")
|
||||||
bn = nn.BatchNorm1d(hidden_units)
|
bn = nn.BatchNorm1d(hidden_units)
|
||||||
seq = nn.Sequential(fc, bn, activation)
|
seq = nn.Sequential(fc, bn, activation)
|
||||||
dnn_layers.append(seq)
|
dnn_layers.append(seq)
|
||||||
@@ -338,7 +441,7 @@ class Net(nn.Module):
|
|||||||
dnn_layers.append(drop_input)
|
dnn_layers.append(drop_input)
|
||||||
fc = nn.Linear(hidden_units, output_dim)
|
fc = nn.Linear(hidden_units, output_dim)
|
||||||
dnn_layers.append(fc)
|
dnn_layers.append(fc)
|
||||||
# optimizer
|
# optimizer # pylint: disable=W0631
|
||||||
self.dnn_layers = nn.ModuleList(dnn_layers)
|
self.dnn_layers = nn.ModuleList(dnn_layers)
|
||||||
self._weight_init()
|
self._weight_init()
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,6 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from typing import Text, Union
|
from typing import Text, Union
|
||||||
@@ -85,7 +84,7 @@ class SFM_Model(nn.Module):
|
|||||||
if len(self.states) == 0: # hasn't initialized yet
|
if len(self.states) == 0: # hasn't initialized yet
|
||||||
self.init_states(x)
|
self.init_states(x)
|
||||||
self.get_constants(x)
|
self.get_constants(x)
|
||||||
p_tm1 = self.states[0]
|
p_tm1 = self.states[0] # noqa: F841
|
||||||
h_tm1 = self.states[1]
|
h_tm1 = self.states[1]
|
||||||
S_re_tm1 = self.states[2]
|
S_re_tm1 = self.states[2]
|
||||||
S_im_tm1 = self.states[3]
|
S_im_tm1 = self.states[3]
|
||||||
@@ -435,7 +434,7 @@ class SFM(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
|
|||||||
@@ -3,7 +3,6 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from typing import Text, Union
|
from typing import Text, Union
|
||||||
@@ -378,7 +377,7 @@ class TabnetModel(Model):
|
|||||||
|
|
||||||
def metric_fn(self, pred, label):
|
def metric_fn(self, pred, label):
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
|
|
||||||
@@ -478,10 +477,10 @@ class TabNet(nn.Module):
|
|||||||
sparse_loss = []
|
sparse_loss = []
|
||||||
out = torch.zeros(x.size(0), self.n_d).to(x.device)
|
out = torch.zeros(x.size(0), self.n_d).to(x.device)
|
||||||
for step in self.steps:
|
for step in self.steps:
|
||||||
x_te, l = step(x, x_a, priors)
|
x_te, loss = step(x, x_a, priors)
|
||||||
out += F.relu(x_te[:, : self.n_d]) # split the feature from feat_transformer
|
out += F.relu(x_te[:, : self.n_d]) # split the feature from feat_transformer
|
||||||
x_a = x_te[:, self.n_d :]
|
x_a = x_te[:, self.n_d :]
|
||||||
sparse_loss.append(l)
|
sparse_loss.append(loss)
|
||||||
return self.fc(out), sum(sparse_loss)
|
return self.fc(out), sum(sparse_loss)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,6 @@ from ...log import get_module_logger
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.optim as optim
|
import torch.optim as optim
|
||||||
from torch.nn.utils import weight_norm
|
|
||||||
|
|
||||||
from .pytorch_utils import count_parameters
|
from .pytorch_utils import count_parameters
|
||||||
from ...model.base import Model
|
from ...model.base import Model
|
||||||
@@ -158,7 +157,7 @@ class TCN(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
|
|||||||
@@ -158,7 +158,7 @@ class TCN(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
|
|||||||
@@ -5,20 +5,12 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import copy
|
import copy
|
||||||
import random
|
import random
|
||||||
from sklearn.metrics import roc_auc_score, mean_squared_error
|
from ...utils import get_or_create_path
|
||||||
import logging
|
from ...log import get_module_logger
|
||||||
from ...utils import (
|
|
||||||
unpack_archive_with_buffer,
|
|
||||||
save_multiple_parts_file,
|
|
||||||
get_or_create_path,
|
|
||||||
drop_nan_by_y_index,
|
|
||||||
)
|
|
||||||
from ...log import get_module_logger, TimeInspector
|
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
@@ -153,7 +145,7 @@ class TCTS(Model):
|
|||||||
|
|
||||||
init_fore_model = copy.deepcopy(self.fore_model)
|
init_fore_model = copy.deepcopy(self.fore_model)
|
||||||
for p in init_fore_model.parameters():
|
for p in init_fore_model.parameters():
|
||||||
p.init_fore_model = False
|
p.requires_grad = False
|
||||||
|
|
||||||
self.fore_model.train()
|
self.fore_model.train()
|
||||||
self.weight_model.train()
|
self.weight_model.train()
|
||||||
@@ -263,7 +255,7 @@ class TCTS(Model):
|
|||||||
x_valid, y_valid = df_valid["feature"], df_valid["label"]
|
x_valid, y_valid = df_valid["feature"], df_valid["label"]
|
||||||
x_test, y_test = df_test["feature"], df_test["label"]
|
x_test, y_test = df_test["feature"], df_test["label"]
|
||||||
|
|
||||||
if save_path == None:
|
if save_path is None:
|
||||||
save_path = get_or_create_path(save_path)
|
save_path = get_or_create_path(save_path)
|
||||||
best_loss = np.inf
|
best_loss = np.inf
|
||||||
while best_loss > self.lowest_valid_performance:
|
while best_loss > self.lowest_valid_performance:
|
||||||
|
|||||||
@@ -6,10 +6,8 @@ import os
|
|||||||
import copy
|
import copy
|
||||||
import math
|
import math
|
||||||
import json
|
import json
|
||||||
import collections
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import seaborn as sns
|
|
||||||
import matplotlib.pyplot as plt
|
import matplotlib.pyplot as plt
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -24,7 +22,6 @@ except ImportError:
|
|||||||
|
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from qlib.utils import get_or_create_path
|
|
||||||
from qlib.constant import EPS
|
from qlib.constant import EPS
|
||||||
from qlib.log import get_module_logger
|
from qlib.log import get_module_logger
|
||||||
from qlib.model.base import Model
|
from qlib.model.base import Model
|
||||||
@@ -745,7 +742,7 @@ def evaluate(pred):
|
|||||||
score = pred.score
|
score = pred.score
|
||||||
label = pred.label
|
label = pred.label
|
||||||
diff = score - label
|
diff = score - label
|
||||||
MSE = (diff ** 2).mean()
|
MSE = (diff**2).mean()
|
||||||
MAE = (diff.abs()).mean()
|
MAE = (diff.abs()).mean()
|
||||||
IC = score.corr(label, method="spearman")
|
IC = score.corr(label, method="spearman")
|
||||||
return {"MSE": MSE, "MAE": MAE, "IC": IC}
|
return {"MSE": MSE, "MAE": MAE, "IC": IC}
|
||||||
|
|||||||
@@ -5,7 +5,6 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from typing import Text, Union
|
from typing import Text, Union
|
||||||
@@ -17,11 +16,9 @@ from ...log import get_module_logger
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.optim as optim
|
import torch.optim as optim
|
||||||
from torch.utils.data import DataLoader
|
|
||||||
|
|
||||||
from .pytorch_utils import count_parameters
|
|
||||||
from ...model.base import Model
|
from ...model.base import Model
|
||||||
from ...data.dataset import DatasetH, TSDatasetH
|
from ...data.dataset import DatasetH
|
||||||
from ...data.dataset.handler import DataHandlerLP
|
from ...data.dataset.handler import DataHandlerLP
|
||||||
|
|
||||||
# qrun examples/benchmarks/Transformer/workflow_config_transformer_Alpha360.yaml ”
|
# qrun examples/benchmarks/Transformer/workflow_config_transformer_Alpha360.yaml ”
|
||||||
@@ -101,7 +98,7 @@ class TransformerModel(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
|
|||||||
@@ -5,7 +5,6 @@
|
|||||||
from __future__ import division
|
from __future__ import division
|
||||||
from __future__ import print_function
|
from __future__ import print_function
|
||||||
|
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import copy
|
import copy
|
||||||
@@ -18,9 +17,8 @@ import torch.nn as nn
|
|||||||
import torch.optim as optim
|
import torch.optim as optim
|
||||||
from torch.utils.data import DataLoader
|
from torch.utils.data import DataLoader
|
||||||
|
|
||||||
from .pytorch_utils import count_parameters
|
|
||||||
from ...model.base import Model
|
from ...model.base import Model
|
||||||
from ...data.dataset import DatasetH, TSDatasetH
|
from ...data.dataset import DatasetH
|
||||||
from ...data.dataset.handler import DataHandlerLP
|
from ...data.dataset.handler import DataHandlerLP
|
||||||
|
|
||||||
|
|
||||||
@@ -98,7 +96,7 @@ class TransformerModel(Model):
|
|||||||
|
|
||||||
mask = torch.isfinite(label)
|
mask = torch.isfinite(label)
|
||||||
|
|
||||||
if self.metric == "" or self.metric == "loss":
|
if self.metric in ("", "loss"):
|
||||||
return -self.loss_fn(pred[mask], label[mask])
|
return -self.loss_fn(pred[mask], label[mask])
|
||||||
|
|
||||||
raise ValueError("unknown metric `%s`" % self.metric)
|
raise ValueError("unknown metric `%s`" % self.metric)
|
||||||
|
|||||||
@@ -26,12 +26,12 @@ def count_parameters(models_or_parameters, unit="m"):
|
|||||||
else:
|
else:
|
||||||
counts = sum(v.numel() for v in models_or_parameters)
|
counts = sum(v.numel() for v in models_or_parameters)
|
||||||
unit = unit.lower()
|
unit = unit.lower()
|
||||||
if unit == "kb" or unit == "k":
|
if unit in ("kb", "k"):
|
||||||
counts /= 2 ** 10
|
counts /= 2**10
|
||||||
elif unit == "mb" or unit == "m":
|
elif unit in ("mb", "m"):
|
||||||
counts /= 2 ** 20
|
counts /= 2**20
|
||||||
elif unit == "gb" or unit == "g":
|
elif unit in ("gb", "g"):
|
||||||
counts /= 2 ** 30
|
counts /= 2**30
|
||||||
elif unit is not None:
|
elif unit is not None:
|
||||||
raise ValueError("Unknown unit: {:}".format(unit))
|
raise ValueError("Unknown unit: {:}".format(unit))
|
||||||
return counts
|
return counts
|
||||||
|
|||||||
@@ -1,6 +1,5 @@
|
|||||||
# MIT License
|
# MIT License
|
||||||
# Copyright (c) 2018 CMU Locus Lab
|
# Copyright (c) 2018 CMU Locus Lab
|
||||||
import torch
|
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch.nn.utils import weight_norm
|
from torch.nn.utils import weight_norm
|
||||||
|
|
||||||
@@ -56,7 +55,7 @@ class TemporalConvNet(nn.Module):
|
|||||||
layers = []
|
layers = []
|
||||||
num_levels = len(num_channels)
|
num_levels = len(num_channels)
|
||||||
for i in range(num_levels):
|
for i in range(num_levels):
|
||||||
dilation_size = 2 ** i
|
dilation_size = 2**i
|
||||||
in_channels = num_inputs if i == 0 else num_channels[i - 1]
|
in_channels = num_inputs if i == 0 else num_channels[i - 1]
|
||||||
out_channels = num_channels[i]
|
out_channels = num_channels[i]
|
||||||
layers += [
|
layers += [
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user