Files
quant-os/tests/platform_qlib_test.py
T

364 lines
13 KiB
Python

import importlib.util
import os
import struct
import sys
import tempfile
import unittest
from contextlib import redirect_stderr
from io import StringIO
from pathlib import Path
from unittest import mock
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
sys.path.insert(0, str(ROOT / "src"))
from platforms.qlib_runner import (
QlibUnavailableError,
_main,
build_alpha158_lightgbm_task,
generate_portable_momentum_signal,
run_native_momentum_backtest,
)
from tools.build_qlib_tiny_fixture import build_default_fixture
class QlibFixtureTest(unittest.TestCase):
def test_standard_library_fixture_has_expected_bin_header(self):
with tempfile.TemporaryDirectory() as temp:
report = build_default_fixture(Path(temp), days=40)
feature = Path(temp) / "features/sh600000/close.day.bin"
payload = feature.read_bytes()
values = struct.unpack("<" + "f" * (len(payload) // 4), payload)
self.assertEqual(values[0], 0.0)
self.assertEqual(len(values), 41)
self.assertEqual(report["calendar_count"], 40)
self.assertTrue((Path(temp) / "calendars/day.txt").is_file())
self.assertTrue((Path(temp) / "instruments/csi300.txt").is_file())
def test_alpha158_lightgbm_config_uses_stable_097_paths(self):
config = build_alpha158_lightgbm_task(
market="csi300",
benchmark="SH000300",
train=("2018-01-01", "2020-12-31"),
valid=("2021-01-01", "2021-12-31"),
test=("2022-01-01", "2022-12-31"),
)
self.assertEqual(
config["task"]["model"]["module_path"],
"qlib.contrib.model.gbdt",
)
self.assertEqual(
config["task"]["dataset"]["kwargs"]["handler"]["module_path"],
"qlib.contrib.data.handler",
)
self.assertEqual(
config["portfolio_analysis"]["executor"]["class"],
"SimulatorExecutor",
)
self.assertEqual(
config["portfolio_analysis"]["strategy"]["class"],
"TopkDropoutStrategy",
)
@unittest.skipUnless(
importlib.util.find_spec("qlib") is not None
and os.environ.get("QUANT60_RUN_QLIB_NATIVE_SMOKE") == "1",
"set QUANT60_RUN_QLIB_NATIVE_SMOKE=1 in a pyqlib 0.9.7 env",
)
def test_optional_native_momentum_smoke(self):
with tempfile.TemporaryDirectory() as temp:
build_default_fixture(Path(temp), days=80)
result = run_native_momentum_backtest(
provider_uri=temp,
market="csi300",
benchmark="SH000300",
start_time="2024-02-01",
end_time="2024-04-19",
topk=1,
n_drop=1,
)
self.assertGreater(len(result["signal"]), 0)
def test_missing_provider_fails_before_optional_import(self):
with self.assertRaises(FileNotFoundError):
run_native_momentum_backtest(
provider_uri="/definitely/missing/quant60-provider",
market="csi300",
benchmark="SH000300",
start_time="2024-01-01",
end_time="2024-02-01",
)
@unittest.skipUnless(
importlib.util.find_spec("pandas") is not None,
"pandas is optional outside the Qlib environment",
)
def test_weekly_signal_uses_true_first_provider_session(self):
import pandas
dates = pandas.bdate_range("2024-01-02", periods=20)
index = pandas.MultiIndex.from_product(
[["SH600000"], dates],
names=["instrument", "datetime"],
)
frame = pandas.DataFrame(
{"$close": [10.0 + index * 0.1 for index in range(len(dates))]},
index=index,
)
class FakeD:
@staticmethod
def features(*args, **kwargs):
del args, kwargs
return frame
signal = generate_portable_momentum_signal(
FakeD,
["SH600000"],
start_time="2024-01-02",
end_time="2024-01-31",
feature_start_time="2024-01-02",
lookback=2,
rebalance="weekly",
)
selected_dates = list(
signal.index.get_level_values("datetime").unique()
)
# Warm-up finishes mid-week, so that partial week is skipped. Every
# emitted date thereafter is the actual first provider session.
self.assertNotIn(pandas.Timestamp("2024-01-04"), selected_dates)
for selected in selected_dates:
week_dates = [
value
for value in dates
if value.to_period("W") == selected.to_period("W")
]
self.assertEqual(selected, min(week_dates))
class QlibCliTest(unittest.TestCase):
def test_legacy_momentum_cli_remains_default_and_writes_json(self):
fake_result = {
"qlib_version": "0.9.7",
"signal": [0.1, 0.2],
"portfolio_metrics": {"1day": object()},
"clock_contract": "test-clock",
"strategy_fidelity": "signal-only",
"a_share_rule_fidelity": "approximation",
}
with tempfile.TemporaryDirectory() as temp:
output = Path(temp) / "nested" / "summary.json"
with mock.patch(
"platforms.qlib_runner.run_native_momentum_backtest",
return_value=fake_result,
) as runner, mock.patch("builtins.print") as printer:
status = _main(
[
"--provider-uri",
temp,
"--start",
"2024-01-02",
"--end",
"2024-03-29",
"--feature-start",
"2023-12-01",
"--lookback",
"10",
"--skip",
"2",
"--topk",
"20",
"--n-drop",
"3",
"--result-json",
str(output),
]
)
self.assertEqual(status, 0)
runner.assert_called_once_with(
provider_uri=temp,
market="csi300",
benchmark="SH000300",
start_time="2024-01-02",
end_time="2024-03-29",
feature_start_time="2023-12-01",
lookback=10,
skip=2,
topk=20,
n_drop=3,
rebalance="weekly",
)
payload = __import__("json").loads(output.read_text(encoding="utf-8"))
self.assertEqual(payload["workflow"], "momentum")
self.assertEqual(payload["signal_rows"], 2)
self.assertEqual(payload["portfolio_frequencies"], ["1day"])
self.assertIn('"workflow": "momentum"', printer.call_args.args[0])
def test_alpha158_cli_dispatches_with_validated_segments(self):
fake_result = {
"qlib_version": "0.9.7",
"experiment_name": "quant-os-test",
"recorder_id": "rec-123",
}
with tempfile.TemporaryDirectory() as temp:
with mock.patch(
"platforms.qlib_runner.run_alpha158_lightgbm_workflow",
return_value=fake_result,
) as runner, mock.patch("builtins.print"):
status = _main(
[
"--workflow",
"alpha158",
"--provider-uri",
temp,
"--market",
"csi500",
"--benchmark",
"SH000905",
"--train-start",
"2018-01-01",
"--train-end",
"2020-12-31",
"--valid-start",
"2021-01-01",
"--valid-end",
"2021-12-31",
"--test-start",
"2022-01-01",
"--test-end",
"2022-12-31",
"--experiment-name",
"quant-os-test",
]
)
self.assertEqual(status, 0)
runner.assert_called_once_with(
provider_uri=temp,
market="csi500",
benchmark="SH000905",
train=("2018-01-01", "2020-12-31"),
valid=("2021-01-01", "2021-12-31"),
test=("2022-01-01", "2022-12-31"),
experiment_name="quant-os-test",
topk=50,
n_drop=5,
)
def test_alpha158_cli_rejects_missing_or_overlapping_segments(self):
base = [
"--workflow",
"alpha158",
"--provider-uri",
"/provider",
"--train-start",
"2018-01-01",
"--train-end",
"2020-12-31",
"--valid-start",
"2020-12-31",
"--valid-end",
"2021-12-31",
"--test-start",
"2022-01-01",
"--test-end",
"2022-12-31",
]
with mock.patch(
"platforms.qlib_runner.run_alpha158_lightgbm_workflow"
) as runner, redirect_stderr(StringIO()):
with self.assertRaises(SystemExit) as overlap:
_main(base)
with self.assertRaises(SystemExit) as missing:
_main(base[:-2])
self.assertEqual(overlap.exception.code, 2)
self.assertEqual(missing.exception.code, 2)
runner.assert_not_called()
def test_momentum_cli_rejects_invalid_dates_and_position_counts(self):
with mock.patch(
"platforms.qlib_runner.run_native_momentum_backtest"
) as runner, redirect_stderr(StringIO()):
with self.assertRaises(SystemExit) as dates:
_main(
[
"--provider-uri",
"/provider",
"--start",
"2024-02-01",
"--end",
"2024-01-01",
]
)
with self.assertRaises(SystemExit) as counts:
_main(
[
"--provider-uri",
"/provider",
"--start",
"2024-01-01",
"--end",
"2024-02-01",
"--topk",
"2",
"--n-drop",
"3",
]
)
self.assertEqual(dates.exception.code, 2)
self.assertEqual(counts.exception.code, 2)
runner.assert_not_called()
def test_cli_rejects_date_arguments_from_the_other_workflow(self):
with mock.patch(
"platforms.qlib_runner.run_native_momentum_backtest"
) as momentum_runner, redirect_stderr(StringIO()):
with self.assertRaises(SystemExit) as momentum:
_main(
[
"--provider-uri",
"/provider",
"--start",
"2024-01-01",
"--end",
"2024-02-01",
"--train-start",
"2020-01-01",
]
)
with mock.patch(
"platforms.qlib_runner.run_alpha158_lightgbm_workflow"
) as alpha_runner, redirect_stderr(StringIO()):
with self.assertRaises(SystemExit) as alpha:
_main(
[
"--workflow",
"alpha158",
"--provider-uri",
"/provider",
"--start",
"2022-01-01",
"--train-start",
"2018-01-01",
"--train-end",
"2020-12-31",
"--valid-start",
"2021-01-01",
"--valid-end",
"2021-12-31",
"--test-start",
"2022-01-01",
"--test-end",
"2022-12-31",
]
)
self.assertEqual(momentum.exception.code, 2)
self.assertEqual(alpha.exception.code, 2)
momentum_runner.assert_not_called()
alpha_runner.assert_not_called()
if __name__ == "__main__":
unittest.main()