1. 我们会用到什么模型呢?# m7 n# D. G1 i3 F: N0 X
在之前的投票分类器中,我们用到了一个分类模型和一个回归模型。 在回归模型中,我们在用于计算分类时,用预测价格替代预测价格走势(下跌、上涨、不变)。 然而,在这种情况下,我们不能依据分类得到概率分布,而对于所谓的“软投票”这样是不允许的。+ h% k! o: K! D
我们已准备了 3 个分类模型。 在“如何在 MQL5 中集成 ONNX 模型的示例”一文中已用到两个模型。 第一个模型(回归)被转换为分类模型。 基于 10 个 OHLC 价格序列进行了培训。 第二个模型是分类模型。 基于 63 个收盘价序列进) i% A& V7 u4 F+ d) ~4 a9 |
//| https://www.mql5.com |" z. x6 }" @- j# ~1 {( C$ P" V
//+------------------------------------------------------------------+
6 J2 f% r6 }& f3 \//--- price movement prediction" E% i) I5 d- D# k& s# t* d( D# m
#define PRICE_UP 0. U3 u8 z k+ F
#define PRICE_SAME 1
+ B, S# K; v* ^& a. X% t: K#define PRICE_DOWN 2, p6 ?4 ]8 l, U
//+------------------------------------------------------------------+/ s4 R# R% N0 L5 @& A: {$ W
//| Base class for models based on trained symbol and period |# ]8 m7 u' G6 C: o8 V$ F
//+------------------------------------------------------------------+
6 i: V2 a2 l) w7 A: Xclass CModelSymbolPeriod7 H$ ]5 ]% S/ y* X6 n J
{ p3 j% i- k/ L# Y$ v
protected:
% @; Y6 J8 i/ g2 f( P" plong m_handle; // created model session handle
8 l$ N4 F; c/ \5 ]string m_symbol; // symbol of trained data
' d. g) |7 X+ r7 D( hENUM_TIMEFRAMES m_period; // timeframe of trained data
7 }+ j5 b" Z5 udatetime m_next_bar; // time of next bar (we work at bar begin only)
8 K1 ^: R% \, I% Y. h1 ndouble m_class_delta; // delta to recognize "price the same" in regression models
/ n& K8 o0 _" p1 i9 _3 ?public:
~! A3 X7 o1 a//+------------------------------------------------------------------+. C# Q9 d E) D/ l- k
//| Constructor |
: }& b( W- j j2 M* Z//+------------------------------------------------------------------+. w J9 u1 f# O$ ?' e" _
CModelSymbolPeriod(const string symbol,const ENUM_TIMEFRAMES period,const double class_delta=0.0001)" z: E' |, v' T; g
{
. V1 N$ U5 F- ]+ n0 U8 F. c" }m_handle=INVALID_HANDLE;* P1 N V9 i# i2 y/ _+ ]) `3 T6 \
m_symbol=symbol;
! B% l, M: v* N @6 G) C. wm_period=period;/ ]# W. W; e5 f5 d
m_next_bar=0;* ^; e2 J7 N6 @+ s0 |0 b9 X, b
m_class_delta=class_delta;
& S$ r& e0 n% @# V6 E}
. [( L8 }4 v7 j8 Z9 s//+------------------------------------------------------------------+. h4 F: W" p q) V& q" A2 S
//| Destructor |5 `0 u6 S) I) P7 r3 f
//| Check for initialization, create model |* ~$ P% e n: Y9 i1 Z
//+------------------------------------------------------------------+8 Q- c1 c3 v" m" U- c: w
bool CheckInit(const string symbol,const ENUM_TIMEFRAMES period,const uchar& model[])
# O" p/ b" g K% s- M{- l8 \- ^9 T0 k
//--- check symbol, period
9 m2 Z0 b! c5 p- u( _0 gif(symbol!=m_symbol || period!=m_period)
. D& @4 }5 W$ |5 m% _' g* }8 {5 e{1 \; ^5 H) n8 {4 d. ?7 v7 A
PrintFormat("Model must work with %s,%s",m_symbol,EnumToString(m_period)); o) P+ R2 W* L/ d7 W( S1 Q0 d. }
return(false);# s- H+ |& [, O7 O; v+ i
}+ c% ]- t8 ~* M, v' |* k
//--- create a model from static buffer! M) K1 s. P* \/ t w3 B9 `5 T1 l
m_handle=OnnxCreateFromBuffer(model,ONNX_DEFAULT);/ N f5 ?0 C2 K7 w1 d) H
if(m_handle==INVALID_HANDLE)
+ Y& @2 T. L' d2 i' P( u$ k8 l4 S{
- G7 I2 u7 T" J) O/ h* }% a/ O# YPrint("OnnxCreateFromBuffer error ",GetLastError());
* c( b! R5 h/ ~* W% ~return(false);
& m; o4 Y; h, O}
& y/ l+ l# `2 `+ r4 Z//--- ok g$ e5 ^* f7 a3 y! }
return(true);/ p/ Y" C6 v1 l% @# X2 }0 A
}* z! n, O* E2 m/ R, V
//+------------------------------------------------------------------+
8 K J5 X9 y( N2 x+ q0 am_next_bar=TimeCurrent();
M) E5 r0 b- `, h" qm_next_bar-=m_next_bar%PeriodSeconds(m_period);
9 v+ b ~: B% W2 Y4 z3 r" }2 Km_next_bar+=PeriodSeconds(m_period);
; s. ^' M( K% V% x//--- work on new day bar$ e# q7 ^2 z. W( j. i+ B
return(true);# W# Z9 h# k9 D
}
: h8 w. g4 V) ]//+------------------------------------------------------------------+$ I% o2 D' A( \( {8 ?; h( W
//| virtual stub for PredictPrice (regression model) |. k. I! {1 S6 l
//+------------------------------------------------------------------+0 F+ }1 ~- p6 t9 Z$ x6 v
virtual double PredictPrice(void)
6 `/ M1 E$ B y4 q% C; g{1 ^0 ^. t# G3 M* U5 l9 ~8 |
return(DBL_MAX);
! J+ B3 @3 ]& a" o& p}
$ A4 M, x4 u v" a4 c T//+------------------------------------------------------------------+7 v0 A% Y7 d0 I P/ t
//| Predict class (regression -> classification) |0 N1 K0 ]+ a9 Y! I7 |
//+------------------------------------------------------------------+
) I5 h' P1 Z& O+ q6 O" `9 X% M4 Mvirtual int PredictClass(void)8 ~: U- y. s" L2 X) i, O; C, {) x
{
' B; E. U: t# k" A$ ~double predicted_price=PredictPrice();6 f+ l |4 T, J+ c+ z
if(predicted_price==DBL_MAX)
9 Q$ R0 a9 @" O) \8 ~return(-1);
+ T( x! W# j; t, G, nint predicted_class=-1;6 y+ I# Y, }- k$ v% r( ^( o
double last_close=iClose(m_symbol,m_period,1);
. B, w8 }" X- c5 n4 S//--- classify predicted price movement
% r' J O5 [9 Q- p u- Xdouble delta=last_close-predicted_price;
" J2 k, l3 Q& D. a% t& l- ~if(fabs(delta)<=m_class_delta)$ u( w& ~$ X* ]! X
predicted_class=PRICE_SAME;
! y N* u S9 e! Kelse' e+ Q- ?6 J7 A5 v* o
private:
* ]" I* k) K) _, B, ?) {# \" X Rint m_sample_size;
4 I( o* B0 D# c9 k o1 Z) [+ }//+------------------------------------------------------------------+
! L8 |& n, h2 Z5 F( J Hvirtual bool Init(const string symbol, const ENUM_TIMEFRAMES period)
' M" b% b3 B/ C4 M{: P7 P# U7 Y5 p3 P$ I: T" w4 r
//--- check symbol, period, create model
. [0 V( k# X9 i* bif(!CModelSymbolPeriod::CheckInit(symbol,period,model_eurusd_D1_10_class))
% [ y T5 T' J4 r9 k{
' E9 M2 Z; t# Y/ F7 V2 nPrint("model_eurusd_D1_10_class : initialization error");# f% Q7 k5 O, O% A8 M2 ?
return(false);
& Z( W- J/ }5 A2 ]}9 P5 k6 u# y& V4 V* g$ L
//--- since not all sizes defined in the input tensor we must set them explicitly
1 ~1 W! T5 G% N' T( W2 ~/ P//--- first index - batch size, second index - series size, third index - number of series (OHLC)( E5 S* {' ^1 D/ h. }
const long input_shape[] = {1,m_sample_size,4};" i% M5 j5 I$ `) q a$ T% H
if(!OnnxSetInputShape(m_handle,0,input_shape))! i- [: }" e4 z" M6 [
{
, N/ M, }, s I6 {' y5 v3 k8 i* ^Print("model_eurusd_D1_10_class : OnnxSetInputShape error ",GetLastError());7 U. g) S+ v2 Z- R* y7 Y8 G+ i
return(false);. D7 H# v4 T$ S9 j
}/ ?1 T6 W/ Y% B- ^9 t4 ?
//--- since not all sizes defined in the output tensor we must set them explicitly7 K2 V, H- q3 W5 C
//--- first index - batch size, must match the batch size of the input tensor
- J9 P' e6 v1 Y+ z! x R//--- second index - number of classes (up, same or down)! @1 D: @0 W+ a ^& r0 g' O
const long output_shape[] = {1,3};
3 ]5 I" ]+ e) X! e6 aif(!OnnxSetOutputShape(m_handle,0,output_shape))
4 h9 r( s- V) T" J. S% W" ?* U6 ^{
1 R& X" R1 N9 xPrint("model_eurusd_D1_10_class : OnnxSetOutputShape error ",GetLastError());
1 B# @2 W: L: s2 {1 Ireturn(false);- ^, Z2 [* q! C
}# _! X- q( C j; w
//--- ok( N; I' B, f* b$ \, H& `9 ]
return(true);
; \! N, ~* E, C3 F8 q}- ?2 X& c8 b7 p0 {% R! r: b
//+------------------------------------------------------------------+ C, S [9 X$ X( q8 u! U
//| Predict class |
. i! c p8 e3 F( b. E5 Z+ E//+------------------------------------------------------------------+# g- r) C& N* K; o& ]
virtual int PredictClass(void)
1 b5 v0 v3 q# V- W# y: @9 {{ |