399 lines
14 KiB
Python
399 lines
14 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,
|
|
_validate_managed_provider_request,
|
|
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",
|
|
)
|
|
|
|
def test_managed_provider_request_fails_before_terminal_session(self):
|
|
evidence = {
|
|
"managed": True,
|
|
"calendar": {
|
|
"start": "2024-01-02",
|
|
"end": "2024-01-11",
|
|
},
|
|
"market": {"name": "tushare_a"},
|
|
"benchmark": {"symbol": "SH999999"},
|
|
}
|
|
with self.assertRaisesRegex(
|
|
ValueError,
|
|
"simulator consumes the next provider session",
|
|
):
|
|
_validate_managed_provider_request(
|
|
evidence,
|
|
market="tushare_a",
|
|
benchmark="SH999999",
|
|
start_time="2024-01-02",
|
|
end_time="2024-01-11",
|
|
)
|
|
_validate_managed_provider_request(
|
|
evidence,
|
|
market="tushare_a",
|
|
benchmark="SH999999",
|
|
start_time="2024-01-02",
|
|
end_time="2024-01-10",
|
|
)
|
|
|
|
@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",
|
|
"run_parameters": {"start_time": "2024-01-02"},
|
|
}
|
|
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.assertEqual(
|
|
payload["run_parameters"]["start_time"],
|
|
"2024-01-02",
|
|
)
|
|
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()
|