Files
zhixing-system/zhixing-server/src/zhixing_server/modules/selection/domain/indicators.py
T

231 lines
7.7 KiB
Python
Raw Normal View History

2026-08-08 22:41:45 +08:00
"""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