|
| 1 | +""" |
| 2 | +RSI Reversion Strategy - RSI 反转策略 |
| 3 | +Classic mean-reversion strategy using RSI overbought/oversold signals. |
| 4 | +""" |
| 5 | + |
| 6 | +import pandas as pd |
| 7 | +import numpy as np |
| 8 | +from quanteval.core.strategy import Strategy |
| 9 | + |
| 10 | + |
| 11 | +class RSIStrategy(Strategy): |
| 12 | + """ |
| 13 | + RSI 反转策略 - 经典均值回归策略 |
| 14 | +
|
| 15 | + RSI Mean Reversion Strategy. |
| 16 | +
|
| 17 | + Strategy Logic: |
| 18 | + - Buy: When RSI < oversold (market oversold) |
| 19 | + - Sell: When RSI > overbought (market overbought) |
| 20 | +
|
| 21 | + Args: |
| 22 | + window: RSI 计算窗口 (RSI window, default 14) |
| 23 | + oversold: 超卖阈值 (Oversold level, default 30) |
| 24 | + overbought: 超买阈值 (Overbought level, default 70) |
| 25 | +
|
| 26 | + Example: |
| 27 | + >>> strategy = RSIStrategy(window=14, oversold=30, overbought=70) |
| 28 | + >>> bt = Backtester(strategy, data) |
| 29 | + >>> results = bt.run() |
| 30 | + """ |
| 31 | + |
| 32 | + def __init__(self, window: int = 14, oversold: float = 30, overbought: float = 70): |
| 33 | + super().__init__( |
| 34 | + name=f'RSI({window})', |
| 35 | + window=window, |
| 36 | + oversold=oversold, |
| 37 | + overbought=overbought, |
| 38 | + ) |
| 39 | + |
| 40 | + def generate_signals(self, data: pd.DataFrame) -> pd.Series: |
| 41 | + """ |
| 42 | + 生成交易信号 |
| 43 | +
|
| 44 | + Generate trading signals based on RSI levels. |
| 45 | +
|
| 46 | + Returns: |
| 47 | + Series with values: |
| 48 | + 1: Long position (RSI < oversold) |
| 49 | + 0: No position (RSI >= oversold) |
| 50 | + """ |
| 51 | + |
| 52 | + window = self.params['window'] |
| 53 | + oversold = self.params['oversold'] |
| 54 | + overbought = self.params['overbought'] |
| 55 | + |
| 56 | + close = data['Close'] |
| 57 | + |
| 58 | + # Calculate price change |
| 59 | + delta = close.diff() |
| 60 | + |
| 61 | + # Separate gains and losses |
| 62 | + gain = delta.clip(lower=0) |
| 63 | + loss = -delta.clip(upper=0) |
| 64 | + |
| 65 | + # Calculate rolling averages |
| 66 | + avg_gain = gain.rolling(window=window).mean() |
| 67 | + avg_loss = loss.rolling(window=window).mean() |
| 68 | + |
| 69 | + # Calculate RS |
| 70 | + rs = avg_gain / avg_loss |
| 71 | + |
| 72 | + # Calculate RSI |
| 73 | + rsi = 100 - (100 / (1 + rs)) |
| 74 | + |
| 75 | + signal = pd.Series(np.nan, index=data.index, name='Signal') |
| 76 | + signal[rsi < oversold] = 1 |
| 77 | + signal[rsi > overbought] = 0 |
| 78 | + |
| 79 | + return signal.ffill().fillna(0) |
0 commit comments