本页源码来自 TradingView 公开发布的开源脚本,版权归原作者所有, 请遵循其原始许可(Pine 脚本常见 CC BY-NC-SA / MPL-2.0 / MIT)。 本项目仅用于研究检索与许可范围内的移植。
// This Pine Script® code is subject to the terms of the Mozilla Public License 2.0 at https://mozilla.org/MPL/2.0/
// © SatohK
//@version=5
indicator("⚡AI-SuperTrend (KNN Machine Learning)", max_labels_count = 200, overlay=true, max_bars_back=2000)
// ==========================================
// --- CONSTANTS & STYLING ---
// ==========================================
color_bull = color.new(#289eff, 0)
color_bear = color.new(#ce3f6c, 0)
color_bull_dim = color.new(#00ffbb, 50)
color_bear_dim = color.new(#ff3355, 50)
color_neutral = color.new(#64748b, 20)
color_gold = color.new(#ffd700, 0)
// ==========================================
// --- MA SELECTOR & HELPER FUNCTIONS ---
// ==========================================
f_zlsma(s, l) =>
lsma = ta.linreg(s, l, 0)
lsma + (lsma - ta.sma(s, l))
f_dema(s, l) =>
e1 = ta.ema(s, l)
2 * e1 - ta.ema(e1, l)
f_tema(s, l) =>
e1 = ta.ema(s, l)
e2 = ta.ema(e1, l)
3 * (e1 - e2) + ta.ema(e2, l)
f_thma(s, l) =>
l_3 = math.max(1, math.round(l / 3))
l_2 = math.max(1, math.round(l / 2))
ta.wma(ta.wma(s, l_3) * 3 - ta.wma(s, l_2) - ta.wma(s, l), l)
calcMA(type, s, l) =>
len = math.max(1, l)
switch type
"SMA" => ta.sma(s, len)
"EMA" => ta.ema(s, len)
"DEMA" => f_dema(s, len)
"TEMA" => f_tema(s, len)
"LSMA" => ta.linreg(s, len, 0)
"WMA" => ta.wma(s, len)
"HMA" => ta.hma(s, len)
"ZLSMA" => f_zlsma(s, len)
"SMMA" => ta.rma(s, len)
"THMA" => f_thma(s, len)
=> ta.sma(s, len)
// ==========================================
// --- INPUT PARAMETERS ---
// ==========================================
group_st = "SuperTrend"
atrPeriod = input.int(10, "ATR Length", minval = 1,group = group_st)
factor = input.float(2.0, "Factor", minval = 0.01, step = 0.01,group = group_st)
group_ml = "🧠 Machine Learning Engine"
k_neighbors = input.int(10, "K-Neighbors (K)", minval=1, group=group_ml, tooltip="Number of nearest neighbors to consider.")
sampling_window_size = input.int(1000, "Learning Window Size", minval=10, group=group_ml, tooltip="Historical data lookback for training.")
momentum_window = input.int(15, "Stride", minval=1, group=group_ml, tooltip="Stride period for Sampling Data")
prob_threshold = input.float(0.9, "Prediction Threshold", minval=0.1, maxval=1.0, step=0.01, group=group_ml, tooltip="Confidence level required for a signal.")
group_feat = "📊 Feature Engineering"
feat_ma_type = input.string("SMA", "Feature MA Type", options=["SMA", "EMA", "DEMA", "TEMA", "LSMA", "WMA", "HMA", "ZLSMA", "SMMA", "THMA"], group=group_feat, tooltip="MA type used for feature calculation.")
rsi_len = input.int(20, "Short RSI Period", group=group_feat)
ma_len = input.int(20, "Short MA Period", group=group_feat)
signal_len = input.int(10, "Signal Line Period", group=group_feat)
window_size = input.int(1000, "Normalizing Window Size", minval=10, group=group_feat, tooltip="Historical data lookback for normalizing.")
p_param = input.float(2.0, "Minkowski Parameter (p)", group=group_feat, tooltip="Distance metric exponent. 2=Euclidean, 1=Manhattan.")
w_param = input.float(2.0, "Shape Parameter", group=group_feat, tooltip="Gausian Weighting exponent.")
group_pca = "⚡ Dimensionality Reduction"
use_pca = input.bool(true, "Enable PCA Compression", group=group_pca, tooltip="Compresses features into 3 Principal Components to reduce noise.")
group_vis = "🎨 Visual Analytics"
use_bar_color = input.bool(true, "Dynamic Bar Coloring", group=group_vis, tooltip="Colors bars based on KNN prediction confidence.")
// ==========================================
// --- LABELING (Supervised Learning) ---
// ==========================================
[supertrend, direction] = ta.supertrend(factor, atrPeriod)
supertrend := barstate.isfirst ? na : supertrend
target = direction*-1
for i = 1 to 5
if direction[i] != direction
target := 0
// ==========================================
// --- FEATURE CALCULATION & NORMALIZATION ---
// ==========================================
normalize(src, len) =>
float _mean = ta.sma(src[1], len)
float _std = ta.stdev(src[1], len)
(src - _mean) / math.max(_std, 0.00001)
f_rsi_s = ta.rsi(close, rsi_len)
f_rsi_m = ta.rsi(close[momentum_window], rsi_len)
f_rsi_l = ta.rsi(close[momentum_window*2], rsi_len)
f_ma_s_dev = (close - calcMA(feat_ma_type, close[1], ma_len)) / calcMA(feat_ma_type, close[1], ma_len)
f_ma_m_dev = (close[momentum_window] - calcMA(feat_ma_type, close[momentum_window+1], ma_len)) / calcMA(feat_ma_type, close[momentum_window+1], ma_len)
f_ma_l_dev = (close[momentum_window*2] - calcMA(feat_ma_type, close[momentum_window*2+1], ma_len)) / calcMA(feat_ma_type, close[momentum_window*2+1], ma_len)
f_rsi_s_sig_dist = (f_rsi_s - ta.sma(f_rsi_s[1], signal_len)) /ta.sma(f_rsi_s[1], signal_len)
f_rsi_m_sig_dist = (f_rsi_m - ta.sma(f_rsi_m[1], signal_len)) /ta.sma(f_rsi_m[1], signal_len)
f_rsi_l_sig_dist = (f_rsi_l - ta.sma(f_rsi_l[1], signal_len)) /ta.sma(f_rsi_l[1], signal_len)
f_rsi_s_z = normalize(f_rsi_s, window_size)
f_rsi_m_z = normalize(f_rsi_m, window_size)
f_rsi_l_z = normalize(f_rsi_l, window_size)
f_ma_s_dev_z = normalize(f_ma_s_dev, window_size)
f_ma_m_dev_z = normalize(f_ma_m_dev, window_size)
f_ma_l_dev_z = normalize(f_ma_l_dev, window_size)
f_rsi_s_sd_z = normalize(f_rsi_s_sig_dist, window_size)
f_rsi_m_sd_z = normalize(f_rsi_m_sig_dist, window_size)
f_rsi_l_sd_z = normalize(f_rsi_l_sig_dist, window_size)
// ==========================================
// --- DIMENSIONALITY REDUCTION ---
// ==========================================
float pc1 = 0.0, float pc2 = 0.0, float pc3 = 0.0
if use_pca
pc1 := (f_rsi_s_z + f_ma_s_dev_z + f_rsi_s_sd_z*0.5)
pc2 := (f_rsi_m_z + f_ma_m_dev_z + f_rsi_m_sd_z*0.5) * 0.9
pc3 := (f_rsi_l_z + f_ma_l_dev_z + f_rsi_l_sd_z*0.5) * 0.8
else
pc1 := f_rsi_m_z
pc2 := f_ma_m_dev_z
pc3 := f_rsi_m_sd_z
// ==========================================
// --- KNN CORE ENGINE ---
// ==========================================
float prob_up = 0.0, float prob_down = 0.0
var float[] distances = array.new_float(0)
var float[] labels = array.new_float(0)
stride = momentum_window
if bar_index > window_size + momentum_window
array.clear(distances)
array.clear(labels)
for i = momentum_window to sampling_window_size + momentum_window by stride
if target[i] != 0
float d1 = math.abs(pc1 - pc1[i])
float d2 = math.abs(pc2 - pc2[i])
float d3 = math.abs(pc3 - pc3[i])
float dist_knn = math.pow(math.pow(d1, p_param) + math.pow(d2, p_param) + math.pow(d3, p_param), 1/p_param)
array.push(distances, dist_knn)
array.push(labels, target[i])
if array.size(distances) >= k_neighbors
int[] sorted_indices = array.sort_indices(distances, order.ascending)
float sum_weight_up = 0.0, float sum_weight_down = 0.0, float total_weight = 0.0
float[] dist_sorted = array.copy(distances)
array.sort(dist_sorted)
float sigma = array.get(dist_sorted, math.min(int(k_neighbors/2), array.size(dist_sorted)-1))
sigma := math.max(sigma, 0.0001)
for j = 0 to k_neighbors - 1
int idx = array.get(sorted_indices, j)
float d = array.get(distances, idx)
float lbl = array.get(labels, idx)
float weight = math.exp(-math.pow(d, w_param) / (2 * math.pow(sigma, 2)))
if lbl == 1
sum_weight_up += weight
else if lbl == -1
sum_weight_down += weight
total_weight += weight
prob_up := total_weight > 0 ? sum_weight_up / total_weight : 0.0
prob_down := total_weight > 0 ? sum_weight_down / total_weight : 0.0
// ==========================================
// --- SIGNAL GENERATION ---
// ==========================================
bool raw_long_signal = ta.crossover(prob_up, prob_threshold)
bool raw_short_signal = ta.crossover(prob_down, prob_threshold)
var int last_dir = 0
bool long_signal = false
bool short_signal = false
if raw_long_signal and last_dir <= 0
long_signal := true
last_dir := 1
if raw_short_signal and last_dir >= 0
short_signal := true
last_dir := -1
var last_stdir = 0.
if direction != nz(direction[1]) or (raw_long_signal) or (raw_short_signal)
if direction < 0 and last_dir==1 and prob_up > 0.5
last_stdir := direction*-1
if direction > 0 and last_dir==-1 and prob_down > 0.5
last_stdir := direction*-1
// ==========================================
// --- VISUALIZATION & PLOTTING ---
// ==========================================
color g_color = prob_up > prob_down ? color.from_gradient(prob_up, 0.5, 0.95, color_neutral, color_bull) : color.from_gradient(prob_down, 0.5, 0.95, color_neutral, color_bear)
barcolor(use_bar_color ? g_color : na)
plotshape(ta.crossover(last_stdir,0), "Major Long",shape.labelup,location.belowbar,color.new(color_bull,30), size=size.tiny, text="▲", textcolor=color.white)
plotshape(ta.crossunder(last_stdir,0), "Major Short",shape.labeldown,location.abovebar,color.new(color_bear,30), size=size.tiny, text="▼", textcolor=color.white)
plotshape(ta.crossunder(direction,0)?supertrend:na, "ST Long",shape.circle,location.absolute,last_stdir>0?color.new(color_bull,40):color.new(color_bear,40), size=size.tiny)
plotshape(ta.crossover(direction,0)?supertrend:na, "ST Short",shape.circle,location.absolute,last_stdir>0?color.new(color_bull,40):color.new(color_bear,40), size=size.tiny)
if ta.crossunder(direction,0)
label.new(
bar_index,
low,
text="Pred\n" + str.tostring(prob_up*100, "#.#") + "%",
color=color.new(color_bull, 90),
textcolor=color_bull,
style=label.style_label_up,
yloc=yloc.belowbar,
size=size.small
)
// 2. Major Short (Bear)
if ta.crossover(direction,0)
label.new(
bar_index,
high,
text="Pred\n" + str.tostring(prob_down*100, "#.#") + "%",
color=color.new(color_bear, 90),
textcolor=color_bear,
style=label.style_label_down,
yloc=yloc.abovebar,
size=size.small
)
upTrend = plot(last_stdir==1? supertrend : na, "Up Trend", color = color_bull, style = plot.style_linebr)
downTrend = plot(last_stdir==-1? supertrend : na, "Down Trend", color = color_bear, style = plot.style_linebr)
bodyMiddle = plot(barstate.isfirst ? na : (open + close) / 2, "Body Middle",display = display.none)
low_p = plot(low,display = display.none)
high_p = plot(high,display = display.none)
fill(low_p, upTrend, (open + close) / 2, supertrend,direction<0?color_bull:na,na)
fill(high_p, downTrend, supertrend, (open + close) / 2, na,direction>0?color_bear:na)
// ==========================================
// --- ALERTS ---
// ==========================================
alert_long = ta.crossover(last_stdir,0)
alert_short = ta.crossunder(last_stdir,0)
alert_combo = alert_long or alert_short
alertcondition(alert_long, "Long", "Bullish Signal with AI Confirmation")
alertcondition(alert_short, "Short", "Bearish Signal with AI Confirmation")
alertcondition(alert_combo, "Signal", "Siganl with AI Confirmation")