97 lines
4.5 KiB
Python
97 lines
4.5 KiB
Python
"""策略通用技术指标计算工具。"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from typing import List, Optional, Sequence, Tuple
|
||
|
||
|
||
def calculate_ma(values: Sequence[float], period: int) -> List[Optional[float]]:
|
||
"""
|
||
功能说明:计算简单移动平均线。
|
||
参数说明:values 为按时间升序排列的数值序列,period 为均线周期。
|
||
返回值说明:返回与输入等长的均线列表,样本不足的位置为 None。
|
||
注意事项:period 必须大于 0。
|
||
"""
|
||
if period <= 0:
|
||
raise ValueError("均线周期必须大于 0")
|
||
result: List[Optional[float]] = []
|
||
window_sum = 0.0
|
||
for index, value in enumerate(values):
|
||
window_sum += value
|
||
if index >= period:
|
||
window_sum -= values[index - period]
|
||
result.append(window_sum / period if index + 1 >= period else None)
|
||
return result
|
||
|
||
|
||
def calculate_ema(values: Sequence[float], period: int) -> List[Optional[float]]:
|
||
"""
|
||
功能说明:计算指数移动平均线。
|
||
参数说明:values 为按时间升序排列的数值序列,period 为 EMA 周期。
|
||
返回值说明:返回与输入等长的 EMA 列表,空输入返回空列表。
|
||
注意事项:首个 EMA 使用首个输入值初始化,period 必须大于 0。
|
||
"""
|
||
if period <= 0:
|
||
raise ValueError("EMA 周期必须大于 0")
|
||
if not values:
|
||
return []
|
||
factor = 2 / (period + 1)
|
||
result: List[Optional[float]] = [float(values[0])]
|
||
for value in values[1:]:
|
||
result.append(float(value) * factor + result[-1] * (1 - factor))
|
||
return result
|
||
|
||
|
||
def calculate_macd(values: Sequence[float], fast_period: int = 12, slow_period: int = 26, signal_period: int = 9) -> Tuple[List[Optional[float]], List[Optional[float]], List[Optional[float]]]:
|
||
"""
|
||
功能说明:计算 MACD 指标的 DIF、DEA 和柱值。
|
||
参数说明:values 为按时间升序排列的收盘价序列,fast_period 为快线周期,slow_period 为慢线周期,signal_period 为信号线周期。
|
||
返回值说明:返回 DIF、DEA、MACD 柱值三个等长列表。
|
||
注意事项:周期必须大于 0,MACD 柱值按 A 股常用口径计算为 2 * (DIF - DEA)。
|
||
"""
|
||
if fast_period <= 0 or slow_period <= 0 or signal_period <= 0:
|
||
raise ValueError("MACD 周期必须大于 0")
|
||
fast_ema = calculate_ema(values, fast_period)
|
||
slow_ema = calculate_ema(values, slow_period)
|
||
dif = [fast_value - slow_value for fast_value, slow_value in zip(fast_ema, slow_ema)]
|
||
dea = calculate_ema(dif, signal_period)
|
||
macd = [2 * (dif_value - dea_value) for dif_value, dea_value in zip(dif, dea)]
|
||
return dif, dea, macd
|
||
|
||
|
||
def calculate_kdj(highs: Sequence[float], lows: Sequence[float], closes: Sequence[float], period: int = 9, k_period: int = 3, d_period: int = 3) -> Tuple[List[Optional[float]], List[Optional[float]], List[Optional[float]]]:
|
||
"""
|
||
功能说明:计算 KDJ 指标的 K、D、J 值。
|
||
参数说明:highs、lows、closes 分别为按时间升序排列的最高价、最低价、收盘价序列,period 为 RSV 周期,k_period 和 d_period 为平滑周期。
|
||
返回值说明:返回 K、D、J 三个等长列表,样本不足的位置为 None。
|
||
注意事项:三个价格序列长度必须一致,周期必须大于 0,首个有效 K 和 D 从 50 开始平滑。
|
||
"""
|
||
if len(highs) != len(lows) or len(highs) != len(closes):
|
||
raise ValueError("KDJ 输入序列长度必须一致")
|
||
if period <= 0 or k_period <= 0 or d_period <= 0:
|
||
raise ValueError("KDJ 周期必须大于 0")
|
||
k_values: List[Optional[float]] = []
|
||
d_values: List[Optional[float]] = []
|
||
j_values: List[Optional[float]] = []
|
||
previous_k = 50.0
|
||
previous_d = 50.0
|
||
for index, close in enumerate(closes):
|
||
if index + 1 < period:
|
||
k_values.append(None)
|
||
d_values.append(None)
|
||
j_values.append(None)
|
||
continue
|
||
start_index = index + 1 - period
|
||
highest = max(highs[start_index : index + 1])
|
||
lowest = min(lows[start_index : index + 1])
|
||
rsv = 50.0 if highest == lowest else (close - lowest) / (highest - lowest) * 100
|
||
current_k = (previous_k * (k_period - 1) + rsv) / k_period
|
||
current_d = (previous_d * (d_period - 1) + current_k) / d_period
|
||
current_j = 3 * current_k - 2 * current_d
|
||
k_values.append(current_k)
|
||
d_values.append(current_d)
|
||
j_values.append(current_j)
|
||
previous_k = current_k
|
||
previous_d = current_d
|
||
return k_values, d_values, j_values
|