回测基本一致
This commit is contained in:
@@ -0,0 +1,78 @@
|
||||
"""Gate primitives — boolean masks over bars (doc 04 Rule 6).
|
||||
|
||||
A gate is a callable that takes the bars DataFrame and returns a boolean
|
||||
``numpy.ndarray`` (``True`` = entry allowed on that bar). Gates never open or
|
||||
close trades; they only mask the signal array the caller hands to the engine.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol, runtime_checkable
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class Gate(Protocol):
|
||||
"""A callable producing a boolean mask over bars (``True`` = allow entry)."""
|
||||
|
||||
def __call__(self, bars: pd.DataFrame) -> np.ndarray: ...
|
||||
|
||||
|
||||
def time_of_day_gate(
|
||||
bars: pd.DataFrame,
|
||||
*,
|
||||
start_hour: int,
|
||||
end_hour: int,
|
||||
timezone: str | None = None,
|
||||
) -> np.ndarray:
|
||||
"""Allow entries only within ``[start_hour, end_hour)`` (hours, 0–24).
|
||||
|
||||
Useful for sessions that only trade London/NY open. ``timezone`` is
|
||||
applied to ``bars['timestamp']`` if given; otherwise the timestamp's
|
||||
existing tz is used (or naive local time).
|
||||
"""
|
||||
ts = pd.to_datetime(bars["timestamp"])
|
||||
if timezone is not None:
|
||||
ts = ts.dt.tz_localize(None).dt.tz_localize(timezone) if ts.dt.tz is None else ts.dt.tz_convert(timezone)
|
||||
hours = ts.dt.hour
|
||||
if start_hour <= end_hour:
|
||||
mask = (hours >= start_hour) & (hours < end_hour)
|
||||
else:
|
||||
# Wrap past midnight, e.g. 22 → 6.
|
||||
mask = (hours >= start_hour) | (hours < end_hour)
|
||||
return mask.to_numpy()
|
||||
|
||||
|
||||
def regime_gate(
|
||||
bars: pd.DataFrame,
|
||||
*,
|
||||
trend_filter: np.ndarray,
|
||||
direction: int,
|
||||
) -> np.ndarray:
|
||||
"""Allow entries only when ``trend_filter`` agrees with ``direction``.
|
||||
|
||||
``trend_filter`` is a +1/-1 array (e.g. from an EMA slope or ADX sign).
|
||||
``direction=+1`` keeps bars where the trend is up; ``-1`` keeps downtrend.
|
||||
"""
|
||||
tf = np.asarray(trend_filter)
|
||||
mask = tf == direction
|
||||
return mask
|
||||
|
||||
|
||||
def exhaustion_gate(
|
||||
rsi_arr: np.ndarray,
|
||||
*,
|
||||
overbought: float = 70.0,
|
||||
oversold: float = 30.0,
|
||||
) -> np.ndarray:
|
||||
"""Block entries when RSI is in the exhaustion zone for the direction.
|
||||
|
||||
Returns ``True`` where entry is *allowed* (i.e. not exhausted). Block longs
|
||||
when ``rsi >= overbought`` and shorts when ``rsi <= oversold`` — combine
|
||||
with the directional signal in the caller.
|
||||
"""
|
||||
rsi_arr = np.asarray(rsi_arr, dtype=float)
|
||||
allow = (rsi_arr < overbought) & (rsi_arr > oversold)
|
||||
allow = np.where(np.isnan(rsi_arr), False, allow)
|
||||
return allow
|
||||
Reference in New Issue
Block a user