feat: add Quant OS A-share baseline
This commit is contained in:
@@ -0,0 +1,363 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user