feat(selection): 迁移知行B1选股策略
This commit is contained in:
@@ -0,0 +1,230 @@
|
||||
"""TDX-style indicator primitives used by the Zhixing B1 formula.
|
||||
|
||||
All inputs are ascending by trading date. Rolling functions intentionally
|
||||
use available observations for the early rows, while ``EVERY`` keeps its
|
||||
full-window requirement. This matches the legacy formula's warm-up behavior
|
||||
without using future rows.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
|
||||
def _check_window(window: int) -> None:
|
||||
"""Validate a positive TDX lookback window."""
|
||||
|
||||
if window < 1:
|
||||
raise ValueError("window must be positive")
|
||||
|
||||
|
||||
def MA(series: pd.Series, window: int) -> pd.Series:
|
||||
"""Return a simple moving average with available-row warm-up."""
|
||||
|
||||
_check_window(window)
|
||||
return series.rolling(window=window, min_periods=1).mean()
|
||||
|
||||
|
||||
def EMA(series: pd.Series, window: int) -> pd.Series:
|
||||
"""Return an adjust-false exponential moving average."""
|
||||
|
||||
_check_window(window)
|
||||
return series.ewm(span=window, adjust=False, min_periods=1).mean()
|
||||
|
||||
|
||||
def LLV(series: pd.Series, window: int) -> pd.Series:
|
||||
"""Return the lowest value in the trailing window."""
|
||||
|
||||
_check_window(window)
|
||||
return series.rolling(window=window, min_periods=1).min()
|
||||
|
||||
|
||||
def HHV(series: pd.Series, window: int) -> pd.Series:
|
||||
"""Return the highest value in the trailing window."""
|
||||
|
||||
_check_window(window)
|
||||
return series.rolling(window=window, min_periods=1).max()
|
||||
|
||||
|
||||
def SMA(series: pd.Series, window: int, weight: int = 1) -> pd.Series:
|
||||
"""Return TDX ``SMA(X,N,M)`` using its recursive weighted average."""
|
||||
|
||||
_check_window(window)
|
||||
if weight < 0 or weight > window:
|
||||
raise ValueError("weight must be between zero and window")
|
||||
return series.ewm(alpha=weight / window, adjust=False, min_periods=1).mean()
|
||||
|
||||
|
||||
def REF(series: pd.Series, periods: int) -> pd.Series:
|
||||
"""Return the value ``periods`` trading rows ago."""
|
||||
|
||||
if periods < 0:
|
||||
raise ValueError("periods must not be negative")
|
||||
return series.shift(periods)
|
||||
|
||||
|
||||
def EXIST(condition: pd.Series, window: int) -> pd.Series:
|
||||
"""Return whether a condition occurred at least once in the window."""
|
||||
|
||||
_check_window(window)
|
||||
values = condition.fillna(False).astype(bool).astype(float)
|
||||
return values.rolling(window=window, min_periods=1).max().astype(bool)
|
||||
|
||||
|
||||
def EVERY(condition: pd.Series, window: int) -> pd.Series:
|
||||
"""Return whether every row in a complete trailing window is true."""
|
||||
|
||||
_check_window(window)
|
||||
values = condition.fillna(False).astype(bool).astype(float)
|
||||
return values.rolling(window=window, min_periods=window).min().fillna(0).astype(bool)
|
||||
|
||||
|
||||
def COUNT(condition: pd.Series, window: int) -> pd.Series:
|
||||
"""Count true rows in the trailing window."""
|
||||
|
||||
_check_window(window)
|
||||
values = condition.fillna(False).astype(bool).astype(float)
|
||||
return values.rolling(window=window, min_periods=1).sum()
|
||||
|
||||
|
||||
def HHVBARS(series: pd.Series, window: int) -> pd.Series:
|
||||
"""Return periods since the most recent trailing maximum."""
|
||||
|
||||
_check_window(window)
|
||||
values = series.to_numpy(dtype=float)
|
||||
result = np.full(len(values), np.nan, dtype=float)
|
||||
for index in range(len(values)):
|
||||
start = max(0, index - window + 1)
|
||||
trailing = values[start : index + 1]
|
||||
finite = np.isfinite(trailing)
|
||||
if not finite.any():
|
||||
continue
|
||||
maximum = np.nanmax(trailing)
|
||||
latest = np.flatnonzero(finite & (trailing == maximum))[-1]
|
||||
result[index] = len(trailing) - 1 - int(latest)
|
||||
return pd.Series(result, index=series.index, dtype=float)
|
||||
|
||||
|
||||
def BARSLAST(condition: pd.Series) -> pd.Series:
|
||||
"""Return periods since the most recent true row, or NaN before one."""
|
||||
|
||||
values = condition.fillna(False).astype(bool).to_numpy()
|
||||
result = np.full(len(values), np.nan, dtype=float)
|
||||
last_true = -1
|
||||
for index, matched in enumerate(values):
|
||||
if matched:
|
||||
last_true = index
|
||||
if last_true >= 0:
|
||||
result[index] = index - last_true
|
||||
return pd.Series(result, index=condition.index, dtype=float)
|
||||
|
||||
|
||||
def CROSS(left: pd.Series, right: pd.Series) -> pd.Series:
|
||||
"""Return rows where ``left`` crosses from below to at-or-above right."""
|
||||
|
||||
previous_left = REF(left, 1)
|
||||
previous_right = REF(right, 1)
|
||||
return (
|
||||
previous_left.notna()
|
||||
& previous_right.notna()
|
||||
& left.notna()
|
||||
& right.notna()
|
||||
& (previous_left < previous_right)
|
||||
& (left >= right)
|
||||
)
|
||||
|
||||
|
||||
def compute_kdj(frame: pd.DataFrame, window: int = 9) -> pd.DataFrame:
|
||||
"""Compute ascending-data K, D and J values.
|
||||
|
||||
A zero high-low range is represented as NaN. K and D carry their prior
|
||||
state across such a row, while J remains NaN there, preventing a flat or
|
||||
incomplete bar from becoming an oversold signal.
|
||||
"""
|
||||
|
||||
_check_window(window)
|
||||
if frame.empty:
|
||||
return pd.DataFrame(index=frame.index, data={"K": [], "D": [], "J": []})
|
||||
low = LLV(frame["low"], window)
|
||||
high = HHV(frame["high"], window)
|
||||
denominator = high - low
|
||||
rsv = ((frame["close"] - low) / denominator.replace(0, np.nan) * 100).to_numpy(float)
|
||||
k = np.full(len(rsv), np.nan, dtype=float)
|
||||
d = np.full(len(rsv), np.nan, dtype=float)
|
||||
previous_k = 50.0
|
||||
previous_d = 50.0
|
||||
for index, value in enumerate(rsv):
|
||||
if np.isfinite(value):
|
||||
previous_k = (2.0 * previous_k + value) / 3.0
|
||||
previous_d = (2.0 * previous_d + previous_k) / 3.0
|
||||
k[index] = previous_k
|
||||
d[index] = previous_d
|
||||
j = 3.0 * k - 2.0 * d
|
||||
return pd.DataFrame(index=frame.index, data={"K": k, "D": d, "J": j})
|
||||
|
||||
|
||||
def compute_rsi(close: pd.Series, window: int = 3) -> pd.Series:
|
||||
"""Compute TDX RSI from close prices, preserving zero-denominator NaN."""
|
||||
|
||||
_check_window(window)
|
||||
previous = REF(close, 1)
|
||||
change = close - previous
|
||||
gain = change.clip(lower=0)
|
||||
absolute_change = change.abs()
|
||||
denominator = SMA(absolute_change, window, 1)
|
||||
return SMA(gain, window, 1).div(denominator.replace(0, np.nan)).mul(100)
|
||||
|
||||
|
||||
def compute_zhixing_lines(close: pd.Series) -> tuple[pd.Series, pd.Series]:
|
||||
"""Return the formula's trend white line and 4-MA yellow line."""
|
||||
|
||||
white = EMA(EMA(close, 10), 10)
|
||||
yellow = (MA(close, 14) + MA(close, 28) + MA(close, 57) + MA(close, 114)) / 4
|
||||
return white, yellow
|
||||
|
||||
|
||||
def is_wide_limit(code: str) -> bool:
|
||||
"""Return whether a code belongs to the 20-percent-limit prefixes."""
|
||||
|
||||
return code.startswith(("68", "30", "4", "8", "9"))
|
||||
|
||||
|
||||
def compute_amplitude_params(code: str, close: pd.Series | pd.DataFrame) -> tuple[float, float]:
|
||||
"""Return ``(daily_range_limit, change_relaxation)`` for one history.
|
||||
|
||||
Ordinary stocks are widened when a more-than-15-percent historical move
|
||||
appears in the available trailing 200 trading rows. The function accepts
|
||||
either a close series or a frame containing ``close`` for test and caller
|
||||
convenience.
|
||||
"""
|
||||
|
||||
values = close["close"] if isinstance(close, pd.DataFrame) else close
|
||||
wide = is_wide_limit(code)
|
||||
if not wide and not values.empty:
|
||||
ratio = values / REF(values, 1)
|
||||
wide = bool(EXIST(ratio > 1.15, min(200, len(values))).iloc[-1])
|
||||
return (8.0, 0.9) if wide else (5.0, 1.0)
|
||||
|
||||
|
||||
def finite_or_none(value: object) -> float | None:
|
||||
"""Convert one numeric scalar to a JSON-safe float or ``None``."""
|
||||
|
||||
if value is None:
|
||||
return None
|
||||
number = float(str(value))
|
||||
return number if np.isfinite(number) else None
|
||||
|
||||
|
||||
def serializable_metrics(values: Sequence[tuple[str, object]]) -> dict[str, float | str | None]:
|
||||
"""Convert target-row metrics into a JSON-safe details mapping."""
|
||||
|
||||
result: dict[str, float | str | None] = {}
|
||||
for key, value in values:
|
||||
if isinstance(value, str) or value is None:
|
||||
result[key] = value
|
||||
else:
|
||||
result[key] = finite_or_none(value)
|
||||
return result
|
||||
Reference in New Issue
Block a user