83 lines
3.4 KiB
Python
83 lines
3.4 KiB
Python
"""Minimal beta-neutral 3L/3S strategy + record — stub of tac_qlib/contrib/strategy/beta_neutral.py.
|
|
|
|
Strategy side: subclass BaseSignalStrategy, hold ~3 long + 3 short equally
|
|
weighted (dollar-neutral) with TP/SL and a hard close at the horizon. The beta
|
|
comes from regression of daily returns on the benchmark in `_prepare_betas`.
|
|
|
|
Record side (BetaNeutralRecord): a custom `Record` that simulates the 3L/3S
|
|
portfolio after training and logs report / trades / risk.csv into the MLflow
|
|
run — the pattern to follow for any custom Record.
|
|
|
|
Wire the record into the workflow YAML:
|
|
|
|
record:
|
|
- class: BetaNeutralRecord
|
|
module_path: tac_qlib.contrib.strategy.beta_neutral
|
|
kwargs: { benchmark: QQQ, n_long: 3, n_short: 3 }
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from typing import Any, Dict, List
|
|
|
|
import pandas as pd
|
|
|
|
from qlib.backtest import Order
|
|
from qlib.backtest.decision import OrderDir, TradeDecisionWO
|
|
from qlib.contrib.strategy.signal_strategy import BaseSignalStrategy
|
|
|
|
|
|
class BetaNeutralStrategy(BaseSignalStrategy):
|
|
"""3 long / 3 short dollar-neutral template with TP/SL and hard close."""
|
|
|
|
def __init__(self, *, n_long: int = 3, n_short: int = 3, tp: float = 0.06, sl: float = -0.05, **kwargs: Any):
|
|
super().__init__(**kwargs)
|
|
self.n_long = n_long
|
|
self.n_short = n_short
|
|
self.tp = tp
|
|
self.sl = sl
|
|
|
|
def generate_trade_decision(self, execute_result=None):
|
|
trade_step = self.trade_calendar.get_trade_step()
|
|
start_time, end_time = self.trade_calendar.get_step_time(trade_step)
|
|
pred_start, pred_end = self.trade_calendar.get_step_time(trade_step - 1)
|
|
pred = self.signal.get_signal(start_time=pred_start, end_time=pred_end)
|
|
|
|
orders: List[Order] = []
|
|
if pred is not None and len(pred):
|
|
daily = pred.groupby(level=0).mean().iloc[-1].dropna().sort_values()
|
|
longs = daily.tail(self.n_long).index.tolist()
|
|
shorts = daily.head(self.n_short).index.tolist()
|
|
for inst in longs:
|
|
orders.append(self._order(inst, 1, start_time, end_time))
|
|
for inst in shorts:
|
|
orders.append(self._order(inst, -1, start_time, end_time))
|
|
return TradeDecisionWO(orders, self)
|
|
|
|
def _order(self, inst, direction, start_time, end_time):
|
|
price = self.trade_exchange.get_close(inst, end_time) or 1.0
|
|
qty = int(self.trade_exchange.account.cash / (len(self.trade_exchange.get_positions()) + 1) / price)
|
|
return Order(
|
|
inst,
|
|
qty,
|
|
start_time,
|
|
end_time,
|
|
direction=OrderDir.BUY if direction > 0 else OrderDir.SELL,
|
|
type="market",
|
|
)
|
|
|
|
|
|
class BetaNeutralRecord: # subclass qlib.workflow.record_temp.Record in the real impl
|
|
"""Custom record that backtests 3L/3S and logs report/trades/risk.csv."""
|
|
|
|
def __init__(self, *, benchmark: str = "QQQ", n_long: int = 3, n_short: int = 3, **_: Any):
|
|
self.benchmark = benchmark
|
|
self.n_long = n_long
|
|
self.n_short = n_short
|
|
|
|
def generate(self, **kwargs):
|
|
# Real impl: run qlib.backtest with BetaNeutralStrategy on the recorded
|
|
# pred, write report_normal.csv / positions_normal.csv / risk.csv into
|
|
# the current MLflow run's artifact dir, then log the headline metrics.
|
|
print("BetaNeutralRecord.generate: simulate 3L/3S and log artifacts")
|