""" ATR 平均真实波幅因子。 """ import pandas as pd from factors.base import BaseFactor class ATRFactor(BaseFactor): """Average True Range,衡量波动性。""" category = "technical" def __init__(self, period: int = 14): self.period = period self.name = f"atr_{period}" def calculate(self, df: pd.DataFrame) -> pd.Series: high, low, close = df["high"], df["low"], df["close"] prev_close = close.shift(1) tr = pd.concat([ (high - low).abs(), (high - prev_close).abs(), (low - prev_close).abs(), ], axis=1).max(axis=1) return tr.ewm(span=self.period, min_periods=self.period).mean() def get_required_columns(self) -> list[str]: return ["high", "low", "close"] class ATRRatioFactor(BaseFactor): """ATR / close 归一化,便于跨股票比较。""" category = "technical" def __init__(self, period: int = 14): self.period = period self.name = f"atr_ratio_{period}" def calculate(self, df: pd.DataFrame) -> pd.Series: atr = ATRFactor(period=self.period).calculate(df) return atr / df["close"].replace(0, float("nan")) * 100 def get_required_columns(self) -> list[str]: return ["high", "low", "close"]