1. 我们会用到什么模型呢?
( F7 l/ a( e7 Z0 ~; h在之前的投票分类器中,我们用到了一个分类模型和一个回归模型。 在回归模型中,我们在用于计算分类时,用预测价格替代预测价格走势(下跌、上涨、不变)。 然而,在这种情况下,我们不能依据分类得到概率分布,而对于所谓的“软投票”这样是不允许的。9 J/ q$ E; @$ h! O: b# r Z
我们已准备了 3 个分类模型。 在“如何在 MQL5 中集成 ONNX 模型的示例”一文中已用到两个模型。 第一个模型(回归)被转换为分类模型。 基于 10 个 OHLC 价格序列进行了培训。 第二个模型是分类模型。 基于 63 个收盘价序列进2 Q3 v/ J/ w0 Q# e1 n; u/ A% Q- T
//| https://www.mql5.com |
8 {# d' X, D0 ~8 Y& Q3 ], G//+------------------------------------------------------------------+
: S/ @2 \( i9 Z//--- price movement prediction
! k4 s2 M( b+ C$ h) q+ c {#define PRICE_UP 05 Q, h5 ^( L: H# c* `
#define PRICE_SAME 12 S: l8 T4 D% }& D8 n: Y$ J6 z
#define PRICE_DOWN 2+ Y% _( [5 w0 x
//+------------------------------------------------------------------+! p4 V# h) n1 q0 ~, l# b, b
//| Base class for models based on trained symbol and period |# e5 F/ } X7 _5 ?
//+------------------------------------------------------------------+8 U# \- a. C# j8 i
class CModelSymbolPeriod" ?2 ?5 _6 H( C: h# o Q
{! B; h% K. ^' _' s; @/ ^; v
protected:3 o4 a# z4 ?7 D$ y
long m_handle; // created model session handle) i6 B; c8 u0 `) J$ k; |
string m_symbol; // symbol of trained data2 a z8 Y% g; B; ^. ]# X; S( O
ENUM_TIMEFRAMES m_period; // timeframe of trained data
. h3 z2 e- S! f* k- _datetime m_next_bar; // time of next bar (we work at bar begin only)& c4 G* f: `4 O9 J, `; J/ e& ~. d% T
double m_class_delta; // delta to recognize "price the same" in regression models
! U7 ~& V3 {; }$ [% U4 Y& C. \& Spublic:
" E2 s W4 Y/ I$ F* e//+------------------------------------------------------------------+
: o1 Y2 q6 _ V' e5 X//| Constructor |! a3 b+ Y `4 e8 g( F& D
//+------------------------------------------------------------------+2 W( t9 T* O/ W5 \! J7 o( d! o
CModelSymbolPeriod(const string symbol,const ENUM_TIMEFRAMES period,const double class_delta=0.0001); x4 f+ e5 a/ P- V7 a. j
{
4 ]* U2 ]' \. _! a+ C7 |% Im_handle=INVALID_HANDLE;
K: v S7 |2 z# @m_symbol=symbol;
[) r' r) O% lm_period=period;
0 @: _& M" a5 h1 Q- P2 N: A& P% {m_next_bar=0;
* j$ l* O* O, Sm_class_delta=class_delta;
; o0 T( D7 j; G. ]/ ?# W8 e. Z; D B/ `}
: l1 |! m0 b( y# i6 n7 _- Y//+------------------------------------------------------------------+
! C! e. a0 a% y0 i3 V//| Destructor |4 `0 E- N1 w) E0 I+ R' i( C. g6 x
//| Check for initialization, create model |
& _9 _$ O, q; e" [/ y//+------------------------------------------------------------------+, [1 k- T$ Q+ G2 D; r% T
bool CheckInit(const string symbol,const ENUM_TIMEFRAMES period,const uchar& model[])
, y- x' Y5 g/ y" ~ D8 Q( z{
8 K: }' K& a7 I8 k$ H {//--- check symbol, period L! e; X5 K8 g' R+ m4 ^
if(symbol!=m_symbol || period!=m_period)
/ r5 d' D3 U+ J+ a3 K6 v2 @{& [6 j, j5 W& G$ F* o
PrintFormat("Model must work with %s,%s",m_symbol,EnumToString(m_period));( k P$ O7 F+ d
return(false);
$ G7 L1 N2 j6 d9 `) s# F}+ W0 M' b+ I3 R7 i0 B- I4 E
//--- create a model from static buffer
8 j+ H: h1 ?8 k i& x0 ~# E5 Y$ a) am_handle=OnnxCreateFromBuffer(model,ONNX_DEFAULT);
* N1 n) E" q: p# K: L7 e4 W9 dif(m_handle==INVALID_HANDLE)
( b0 r* W( d& U, N, S6 }{" }$ z( V% r& c3 D* A# a3 q
Print("OnnxCreateFromBuffer error ",GetLastError());
( ] p6 u# Z' p$ @7 sreturn(false);; q# @( o8 S: d& ]. x
}
4 f' e( @: M0 v `8 v7 N0 ~- J; v//--- ok
' l# N/ z& e! V# t) k3 ~; o6 U' D5 oreturn(true);
/ p% l( h. X. b* k2 t3 p}
3 g& Z9 E, K/ h' X. n* k//+------------------------------------------------------------------+( d) y) O( ?, H- C# c
m_next_bar=TimeCurrent();3 |5 l+ z+ X/ ` h% l1 P' q
m_next_bar-=m_next_bar%PeriodSeconds(m_period);
0 l" w/ }% z h/ ?7 x; [$ l4 g/ ^m_next_bar+=PeriodSeconds(m_period);
' ^0 S6 Y( s, M- b: P r* N//--- work on new day bar
9 H N( n' J. Z( X' C4 d Ereturn(true);4 T( S9 S% [1 \ K4 ^
}
* i, k9 e& i/ i9 a//+------------------------------------------------------------------+
v4 q# m" d7 S2 k2 Y//| virtual stub for PredictPrice (regression model) |
) }) P. t6 S. g2 e/ O2 \1 Z( {//+------------------------------------------------------------------+
0 P0 q# b+ C+ n/ [) p0 Pvirtual double PredictPrice(void)( C9 C2 c" S! t* f9 ~! y
{
- i; U# h# f. H6 G! ?! greturn(DBL_MAX);
/ [3 e1 |; d W) @8 l# R) j}
2 G, n8 H, V2 K, z7 P//+------------------------------------------------------------------+
/ f) a/ Q6 Y; \( a//| Predict class (regression -> classification) |$ b' t# C$ m2 K6 o/ C- }
//+------------------------------------------------------------------+
2 t( p4 W! ]1 p6 X$ G& cvirtual int PredictClass(void)
) P7 S( b/ r6 o6 p' x{
4 V1 J7 T0 F* l/ u& K e; |# fdouble predicted_price=PredictPrice();2 t& N. u& M% q2 X* ~1 V% L
if(predicted_price==DBL_MAX)( G# q0 n! C( p( \ {) s
return(-1);* @- Z/ \( t2 ]) f/ C
int predicted_class=-1;8 M! {' Z0 w6 Y
double last_close=iClose(m_symbol,m_period,1);
4 a; c# P, L5 w% y* r: w//--- classify predicted price movement" f% @( x; Z+ t$ p2 _0 c
double delta=last_close-predicted_price;7 o9 J h) P. |
if(fabs(delta)<=m_class_delta)
8 ?( s4 M+ s1 G* s; jpredicted_class=PRICE_SAME;& K* N& f( E. I& o( S: }
else
& d& G$ m8 W# h4 ]% lprivate:
+ I2 ]% m% ]/ K5 ^& V4 Gint m_sample_size;7 F( H, G7 g" {) G9 Q: J
//+------------------------------------------------------------------+, y7 M5 w* I* u- L: V
virtual bool Init(const string symbol, const ENUM_TIMEFRAMES period)* @% J8 Q+ E. B0 }
{& x) }5 V0 m) \. ]+ S
//--- check symbol, period, create model
: c+ C7 p) T# k, ~: Oif(!CModelSymbolPeriod::CheckInit(symbol,period,model_eurusd_D1_10_class))% o5 a( F. \; ]! M
{3 a' m7 ~3 B! X: L* v, A# \
Print("model_eurusd_D1_10_class : initialization error");- o9 R( v4 S' N) n( t& G
return(false);/ K' j$ h3 t3 D6 |
}% F" e* v2 Q* o, m8 S9 F; s
//--- since not all sizes defined in the input tensor we must set them explicitly. S9 V/ n& c' E* ?( J
//--- first index - batch size, second index - series size, third index - number of series (OHLC)) M, O; |3 Z$ R! P
const long input_shape[] = {1,m_sample_size,4};
. K) _ U4 D5 Zif(!OnnxSetInputShape(m_handle,0,input_shape))& Q5 b5 P4 p# V0 r0 R( _8 ?: x
{
' B" ? G5 q( N; XPrint("model_eurusd_D1_10_class : OnnxSetInputShape error ",GetLastError());
0 X2 B& x+ q! v0 Ereturn(false);
1 I- z$ ]. L6 y" x* ~; z' @}, ?- K, V! U7 ?" C$ O/ ^" {
//--- since not all sizes defined in the output tensor we must set them explicitly
; y' p! A- q& Y5 i7 i//--- first index - batch size, must match the batch size of the input tensor9 R0 q( h' p" g8 n5 ?9 S
//--- second index - number of classes (up, same or down)# C# e5 |0 x; e$ w$ L6 L+ q+ P
const long output_shape[] = {1,3};
2 C* }0 W; @9 ]5 {% S: Hif(!OnnxSetOutputShape(m_handle,0,output_shape))
' K2 `. P: }" s% @/ c/ S{
! s! Q- s- I' ]! |* OPrint("model_eurusd_D1_10_class : OnnxSetOutputShape error ",GetLastError());
4 h8 E2 c4 [7 Mreturn(false);
7 f; w2 o% s, F8 V$ \}( ~' { W# p' a$ J i0 ?
//--- ok. J# ~" y) B2 U2 {# s+ }
return(true);
2 j. @% j6 o& f% Y( m: f}; V% c) ?7 F2 m, X/ ]0 \# u; n
//+------------------------------------------------------------------+' e. Y) G2 C* y6 u; o
//| Predict class |
/ y8 o5 I, c0 D+ x//+------------------------------------------------------------------+
' U! k5 m, J1 |5 s! m7 j6 Wvirtual int PredictClass(void)
* F2 {: E. p# R; G! Y{ |