Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
27 changes: 20 additions & 7 deletions docs/component/rl/quickstart.rst
Original file line number Diff line number Diff line change
Expand Up @@ -50,13 +50,13 @@ QlibRL provides an example of an implementation of a single asset order executio
data:
source:
order_dir: ./data/training_order_split
data_dir: ./data/pickle_dataframe/backtest
feature_root_dir: ./data/pickle_dataframe/backtest
# number of time indexes
total_time: 240
# start time index
default_start_time: 0
# end time index
default_end_time: 240
# start time index (must be divisible by simulator.data_granularity)
default_start_time_index: 0
# end time index (must be divisible by simulator.data_granularity)
default_end_time_index: 240
proc_data_dim: 6
num_workers: 0
queue_size: 20
Expand Down Expand Up @@ -84,6 +84,19 @@ QlibRL provides an example of an implementation of a single asset order executio
checkpoint_path: ./checkpoints
checkpoint_every_n_iters: 1

.. warning::

``data_config["source"]["default_start_time_index"]`` and
``data_config["source"]["default_end_time_index"]`` **must both be
divisible by** ``simulator.data_granularity``. This is because the RL
environment slices each trading day into uniform windows of
``data_granularity`` ticks, so misaligned indices would corrupt the step
boundaries and raise a ``ValueError`` at startup.

For example, if ``simulator.data_granularity: 5``, valid values are
``0, 5, 10, 15, …``. The safest starting point is always
``default_start_time_index: 0``.


And the config file for backtesting:

Expand Down Expand Up @@ -160,13 +173,13 @@ With the above config files, you can start training the agent by the following c

.. code-block:: console

$ python -m qlib.rl.contrib.train_onpolicy.py --config_path train_config.yml
$ python -m qlib.rl.contrib.train_onpolicy --config_path train_config.yml

After the training, you can backtest with the following command:

.. code-block:: console

$ python -m qlib.rl.contrib.backtest.py --config_path backtest_config.yml
$ python -m qlib.rl.contrib.backtest --config_path backtest_config.yml

In that case, :class:`~qlib.rl.order_execution.simulator_qlib.SingleAssetOrderExecution` and :class:`~qlib.rl.order_execution.simulator_simple.SingleAssetOrderExecutionSimple` as examples for simulator, :class:`qlib.rl.order_execution.interpreter.FullHistoryStateInterpreter` and :class:`qlib.rl.order_execution.interpreter.CategoricalActionInterpreter` as examples for interpreter, :class:`qlib.rl.order_execution.policy.PPO` as an example for policy, and :class:`qlib.rl.order_execution.reward.PAPenaltyReward` as an example for reward.
For the single asset order execution task, if developers have already defined their simulator/interpreters/reward function/policy, they could launch the training and backtest pipeline by simply modifying the corresponding settings in the config files.
Expand Down
11 changes: 9 additions & 2 deletions qlib/backtest/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

import numpy as np

from qlib.utils.time import epsilon_change
from qlib.utils.time import epsilon_change, Freq

if TYPE_CHECKING:
from qlib.backtest.decision import BaseTradeDecision
Expand Down Expand Up @@ -128,7 +128,14 @@ def get_step_time(self, trade_step: int | None = None, shift: int = 0) -> Tuple[
if trade_step is None:
trade_step = self.get_trade_step()
calendar_index = self.start_index + trade_step - shift
return self._calendar[calendar_index], epsilon_change(self._calendar[calendar_index + 1])
left = self._calendar[calendar_index]
# if we're at the very last bar, there's no next entry to peek at, so
# derive the right endpoint from the current bar's own frequency instead
if calendar_index + 1 < len(self._calendar):
right = self._calendar[calendar_index + 1]
else:
right = left + Freq.get_timedelta(*Freq.parse(self.freq))
return left, epsilon_change(right)

def get_data_cal_range(self, rtype: str = "full") -> Tuple[int, int]:
"""
Expand Down
78 changes: 78 additions & 0 deletions qlib/rl/contrib/_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT License.
"""Lightweight validation utilities for RL training configuration.

This module intentionally has **zero** heavy dependencies (no torch, tianshou,
gym, etc.) so that the helpers here can be imported and unit-tested in any
environment, including lightweight CI images.
"""

from __future__ import annotations


def validate_time_alignment(
default_start_time_index: int,
default_end_time_index: int,
data_granularity: int,
) -> None:
"""Validate that start/end time indices are aligned to ``data_granularity``.

Both ``default_start_time_index`` and ``default_end_time_index`` must be
exact multiples of ``data_granularity``. This is required because the RL
environment slices each trading day into uniform windows of
``data_granularity`` ticks; a misaligned start/end index would silently
corrupt step boundaries.

Parameters
----------
default_start_time_index : int
The starting tick index, taken from
``data_config["source"]["default_start_time_index"]``.
default_end_time_index : int
The ending tick index (exclusive), taken from
``data_config["source"]["default_end_time_index"]``.
data_granularity : int
Number of raw ticks grouped into a single data point (from
``simulator_config["data_granularity"]``).

Raises
------
ValueError
If either index is not divisible by ``data_granularity``.

Examples
--------
Valid – 0 and 240 are both multiples of 5:

>>> validate_time_alignment(0, 240, 5)

Invalid – 3 is not a multiple of 5:

>>> validate_time_alignment(3, 240, 5) # doctest: +ELLIPSIS
Traceback (most recent call last):
...
ValueError: data_config['source']['default_start_time_index'] (=3) ...
"""
if default_start_time_index % data_granularity != 0:
_suggested = (default_start_time_index // data_granularity) * data_granularity
raise ValueError(
f"data_config['source']['default_start_time_index'] "
f"(={default_start_time_index}) must be divisible by "
f"data_granularity (={data_granularity}), but "
f"{default_start_time_index} % {data_granularity} = "
f"{default_start_time_index % data_granularity}. "
f"Hint: set default_start_time_index to a multiple of "
f"{data_granularity} (e.g. {_suggested})."
)

if default_end_time_index % data_granularity != 0:
_suggested_end = (default_end_time_index // data_granularity) * data_granularity
raise ValueError(
f"data_config['source']['default_end_time_index'] "
f"(={default_end_time_index}) must be divisible by "
f"data_granularity (={data_granularity}), but "
f"{default_end_time_index} % {data_granularity} = "
f"{default_end_time_index % data_granularity}. "
f"Hint: set default_end_time_index to a multiple of "
f"{data_granularity} (e.g. {_suggested_end})."
)
65 changes: 63 additions & 2 deletions qlib/rl/contrib/train_onpolicy.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,8 @@ def __getitem__(self, index: int) -> Order:

return order

from qlib.rl.contrib._utils import validate_time_alignment as _validate_time_alignment


def train_and_test(
env_config: dict,
Expand All @@ -109,6 +111,58 @@ def train_and_test(
run_training: bool,
run_backtest: bool,
) -> None:
"""Run the training and/or backtest pipeline for a single-asset order execution RL task.

Parameters
----------
env_config : dict
Environment configuration (``concurrency``, ``parallel_mode``, etc.).
simulator_config : dict
Simulator configuration including ``time_per_step``, ``vol_limit``, and
optionally ``data_granularity`` (default ``1``). ``data_granularity``
controls how many raw ticks are grouped into a single data point fed to
the agent.
trainer_config : dict
Trainer configuration (``max_epoch``, ``batch_size``, etc.).
data_config : dict
Data source configuration. The sub-key ``data_config["source"]`` must
contain at minimum:

- ``order_dir`` – path to the directory holding order pickle files.
- ``feature_root_dir`` – path to the feature data directory.
- ``default_start_time_index`` – the starting tick index of the trading
window. **Must be divisible by** ``data_granularity``.
- ``default_end_time_index`` – the ending tick index of the trading
window (exclusive). **Must be divisible by** ``data_granularity``.

.. note::
Both ``default_start_time_index`` and ``default_end_time_index``
must be exact multiples of ``data_granularity``. This is required
because the RL environment slices each trading day into uniform
windows of ``data_granularity`` ticks, so any misaligned start/end
index would silently corrupt the step boundaries. A safe default
is ``default_start_time_index: 0``. For example, if
``data_granularity`` is ``5``, valid values are ``0, 5, 10, …``.

state_interpreter : StateInterpreter
Interpreter that converts raw simulator state to an RL observation.
action_interpreter : ActionInterpreter
Interpreter that converts RL action to a simulator action.
policy : BasePolicy
The tianshou policy to train or evaluate.
reward : Reward
Reward function.
run_training : bool
Whether to run the training loop.
run_backtest : bool
Whether to run the backtest loop.

Raises
------
ValueError
If ``default_start_time_index`` or ``default_end_time_index`` is not
divisible by ``data_granularity``.
"""
order_root_path = Path(data_config["source"]["order_dir"])

data_granularity = simulator_config.get("data_granularity", 1)
Expand All @@ -124,8 +178,15 @@ def _simulator_factory_simple(order: Order) -> SingleAssetOrderExecutionSimple:
vol_threshold=simulator_config["vol_limit"],
)

assert data_config["source"]["default_start_time_index"] % data_granularity == 0
assert data_config["source"]["default_end_time_index"] % data_granularity == 0
_start = data_config["source"]["default_start_time_index"]
_end = data_config["source"]["default_end_time_index"]

_validate_time_alignment(
default_start_time_index=_start,
default_end_time_index=_end,
data_granularity=data_granularity,
)


if run_training:
train_dataset, valid_dataset = [
Expand Down
Loading