3 分钟快速生成代码
输入想法,AI 即刻生成可运行代码
在 PTrade 量化交易平台中,除了直接使用内置的 get_RSI 函数外,许多开发者希望基于历史价格数据(通过 get_history 接口)自定义实现 RSI(Relative Strength Index,相对强弱指标)算法,以便调整平滑方式(如 SMA 或 EMA)或与其他逻辑深度结合。
RSI 指标通过比较一段时间内的平均涨幅和平均跌幅,来衡量价格变化的强弱。计算步骤如下:
利用 numpy 或 pandas 库,我们可以轻松实现上述逻辑:
import numpy as np
import pandas as pd
def calculate_custom_rsi(close_prices, n=6):
"""
自定义 RSI 计算函数
:param close_prices: 历史收盘价序列 (pd.Series 或 np.ndarray)
:param n: RSI 周期,默认 6
:return: RSI 计算结果数组 (np.ndarray)
"""
# 计算每日价格变动
deltas = np.diff(close_prices)
# 分离上涨和下跌幅度
gains = np.where(deltas > 0, deltas, 0.0)
losses = np.where(deltas < 0, -deltas, 0.0)
# 使用移动平均计算 N 日平均涨跌幅
# 此处采用简单移动平均 (SMA)
gains_series = pd.Series(gains)
losses_series = pd.Series(losses)
avg_gains = gains_series.rolling(window=n).mean()
avg_losses = losses_series.rolling(window=n).mean()
# 计算 RSI
rs = avg_gains / (avg_losses + 1e-10) # 防止除以 0
rsi = 100.0 - (100.0 / (1.0 + rs))
return rsi.values
以下示例展示了如何在 PTrade 的 handle_data 编写逻辑,调用 get_history 获取历史数据并计算自定义 RSI:
import numpy as np
import pandas as pd
def initialize(context):
# 设置目标标的
g.security = '600570.SS'
set_universe(g.security)
# RSI 参数设置
g.rsi_period = 6
def calculate_custom_rsi(close_prices, n=6):
deltas = np.diff(close_prices)
gains = np.where(deltas > 0, deltas, 0.0)
losses = np.where(deltas < 0, -deltas, 0.0)
gains_series = pd.Series(gains)
losses_series = pd.Series(losses)
avg_gains = gains_series.rolling(window=n).mean()
avg_losses = losses_series.rolling(window=n).mean()
rsi = 100.0 * avg_gains / (avg_gains + avg_losses + 1e-10)
return rsi.values
def handle_data(context, data):
security = g.security
n = g.rsi_period
# 获取包含计算所需足够的历史数据 (周期 N + 缓冲数据)
df = get_history(n + 15, '1d', 'close', security_list=security, fq=None, include=False)
if df is None or df.empty:
return
close_data = df.query('code in [@security]')['close'].values
if len(close_data) < n + 1:
return
# 计算 RSI 序列
rsi_array = calculate_custom_rsi(close_data, n=n)
current_rsi = rsi_array[-1]
log.info("%s 当前 %d 日 RSI 值为: %.2f" % (security, n, current_rsi))
# 简易策略逻辑:RSI 超卖买入,超买卖出
cash = context.portfolio.cash
position = get_position(security).amount
if current_rsi < 20 and cash > 0:
order_value(security, cash)
log.info("RSI 低于 20,触发超卖买入")
elif current_rsi > 80 and position > 0:
order_target(security, 0)
log.info("RSI 高于 80,触发超买卖出")
np.diff 会使数组长度减少 1,因此调用 get_history 获取历史数据时,指定的 count 至少需要大于 N + 1,建议预留 10~20 根 K 线的缓冲以确保滑动窗口计算稳定。get_history 的 fq 参数设置为 'pre'(前复权),以消除除权除息对价格跳空产生的影响。rolling().mean() 替换为 ewm()。