⚡AI-SuperTrend (KNN Machine Learning)

SatohK · study · 259 行 · 点赞 3,367 · TradingView 原页

本页源码来自 TradingView 公开发布的开源脚本,版权归原作者所有, 请遵循其原始许可(Pine 脚本常见 CC BY-NC-SA / MPL-2.0 / MIT)。 本项目仅用于研究检索与许可范围内的移植。

Pine Script

// 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")

← 返回列表