Refactor Supertrend method for efficiency

Refactor Supertrend calculation to improve clarity and performance by using numpy arrays for calculations.
This commit is contained in:
Kagari
2026-01-04 05:10:33 +08:00
committed by GitHub
parent 42b6030400
commit c07aa68697
+65 -50
View File
@@ -80,30 +80,47 @@ class Supertrend(IStrategy):
sell_p3 = IntParameter(7, 21, default=14)
def populate_indicators(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
new_cols = []
for multiplier in self.buy_m1.range:
for period in self.buy_p1.range:
dataframe[f'supertrend_1_buy_{multiplier}_{period}'] = self.supertrend(dataframe, multiplier, period)['STX']
st = self.supertrend(dataframe, multiplier, period)[['STX']].rename(
columns={'STX': f'supertrend_1_buy_{multiplier}_{period}'})
new_cols.append(st)
for multiplier in self.buy_m2.range:
for period in self.buy_p2.range:
dataframe[f'supertrend_2_buy_{multiplier}_{period}'] = self.supertrend(dataframe, multiplier, period)['STX']
st = self.supertrend(dataframe, multiplier, period)[['STX']].rename(
columns={'STX': f'supertrend_2_buy_{multiplier}_{period}'})
new_cols.append(st)
for multiplier in self.buy_m3.range:
for period in self.buy_p3.range:
dataframe[f'supertrend_3_buy_{multiplier}_{period}'] = self.supertrend(dataframe, multiplier, period)['STX']
st = self.supertrend(dataframe, multiplier, period)[['STX']].rename(
columns={'STX': f'supertrend_3_buy_{multiplier}_{period}'})
new_cols.append(st)
for multiplier in self.sell_m1.range:
for period in self.sell_p1.range:
dataframe[f'supertrend_1_sell_{multiplier}_{period}'] = self.supertrend(dataframe, multiplier, period)['STX']
st = self.supertrend(dataframe, multiplier, period)[['STX']].rename(
columns={'STX': f'supertrend_1_sell_{multiplier}_{period}'})
new_cols.append(st)
for multiplier in self.sell_m2.range:
for period in self.sell_p2.range:
dataframe[f'supertrend_2_sell_{multiplier}_{period}'] = self.supertrend(dataframe, multiplier, period)['STX']
st = self.supertrend(dataframe, multiplier, period)[['STX']].rename(
columns={'STX': f'supertrend_2_sell_{multiplier}_{period}'})
new_cols.append(st)
for multiplier in self.sell_m3.range:
for period in self.sell_p3.range:
dataframe[f'supertrend_3_sell_{multiplier}_{period}'] = self.supertrend(dataframe, multiplier, period)['STX']
st = self.supertrend(dataframe, multiplier, period)[['STX']].rename(
columns={'STX': f'supertrend_3_sell_{multiplier}_{period}'})
new_cols.append(st)
if new_cols:
dataframe = pd.concat([dataframe] + new_cols, axis=1)
return dataframe
def populate_entry_trend(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
@@ -136,42 +153,40 @@ class Supertrend(IStrategy):
Supertrend Indicator; adapted for freqtrade
from: https://github.com/freqtrade/freqtrade-strategies/issues/30
"""
def supertrend(self, dataframe: DataFrame, multiplier, period):
def supertrend(self, dataframe: pd.DataFrame, multiplier, period):
df = dataframe.copy()
df['TR'] = ta.TRANGE(df)
df['ATR'] = ta.SMA(df['TR'], period)
st = 'ST_' + str(period) + '_' + str(multiplier)
stx = 'STX_' + str(period) + '_' + str(multiplier)
# Compute basic upper and lower bands
df['basic_ub'] = (df['high'] + df['low']) / 2 + multiplier * df['ATR']
df['basic_lb'] = (df['high'] + df['low']) / 2 - multiplier * df['ATR']
# Compute final upper and lower bands
df['final_ub'] = 0.00
df['final_lb'] = 0.00
for i in range(period, len(df)):
df['final_ub'].iat[i] = df['basic_ub'].iat[i] if df['basic_ub'].iat[i] < df['final_ub'].iat[i - 1] or df['close'].iat[i - 1] > df['final_ub'].iat[i - 1] else df['final_ub'].iat[i - 1]
df['final_lb'].iat[i] = df['basic_lb'].iat[i] if df['basic_lb'].iat[i] > df['final_lb'].iat[i - 1] or df['close'].iat[i - 1] < df['final_lb'].iat[i - 1] else df['final_lb'].iat[i - 1]
# Set the Supertrend value
df[st] = 0.00
for i in range(period, len(df)):
df[st].iat[i] = df['final_ub'].iat[i] if df[st].iat[i - 1] == df['final_ub'].iat[i - 1] and df['close'].iat[i] <= df['final_ub'].iat[i] else \
df['final_lb'].iat[i] if df[st].iat[i - 1] == df['final_ub'].iat[i - 1] and df['close'].iat[i] > df['final_ub'].iat[i] else \
df['final_lb'].iat[i] if df[st].iat[i - 1] == df['final_lb'].iat[i - 1] and df['close'].iat[i] >= df['final_lb'].iat[i] else \
df['final_ub'].iat[i] if df[st].iat[i - 1] == df['final_lb'].iat[i - 1] and df['close'].iat[i] < df['final_lb'].iat[i] else 0.00
# Mark the trend direction up/down
df[stx] = np.where((df[st] > 0.00), np.where((df['close'] < df[st]), 'down', 'up'), None)
# Remove basic and final bands from the columns
df.drop(['basic_ub', 'basic_lb', 'final_ub', 'final_lb'], inplace=True, axis=1)
df.fillna(0, inplace=True)
return DataFrame(index=df.index, data={
'ST' : df[st],
'STX' : df[stx]
})
high = df['high'].values
low = df['low'].values
close = df['close'].values
length = len(df)
# 1. TR and ATR
tr = ta.TRANGE(df['high'], df['low'], df['close'])
atr = pd.Series(tr).rolling(period).mean().to_numpy()
# 2. basic upper / lower bands
basic_ub = (high + low) / 2 + multiplier * atr
basic_lb = (high + low) / 2 - multiplier * atr
# 3. final upper / lower bands
final_ub = np.zeros(length)
final_lb = np.zeros(length)
final_ub[:period] = basic_ub[:period]
final_lb[:period] = basic_lb[:period]
for i in range(period, length):
final_ub[i] = basic_ub[i] if basic_ub[i] < final_ub[i-1] or close[i-1] > final_ub[i-1] else final_ub[i-1]
final_lb[i] = basic_lb[i] if basic_lb[i] > final_lb[i-1] or close[i-1] < final_lb[i-1] else final_lb[i-1]
# 4. ST calculation
st = np.zeros(length)
for i in range(period, length):
if st[i-1] == final_ub[i-1]:
st[i] = final_ub[i] if close[i] <= final_ub[i] else final_lb[i]
elif st[i-1] == final_lb[i-1]:
st[i] = final_lb[i] if close[i] >= final_lb[i] else final_ub[i]
# 5. STX direction
stx = np.where(st > 0, np.where(close < st, 'down', 'up'), None)
return pd.DataFrame({'ST': st, 'STX': stx}, index=df.index)