From 33513468e6f8ad295c5e76174f3c558c84f9ef80 Mon Sep 17 00:00:00 2001 From: Francesco Gatti Date: Mon, 3 Jul 2017 09:01:40 +0000 Subject: [PATCH] virtual method --- include/Layer.h | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/include/Layer.h b/include/Layer.h index 9ec37d4..7059695 100644 --- a/include/Layer.h +++ b/include/Layer.h @@ -43,7 +43,7 @@ public: Layer(Network *net, dataDim_t input_dim); virtual ~Layer(); - value_type* infer(dataDim_t &dim, value_type* srcData) { + virtual value_type* infer(dataDim_t &dim, value_type* srcData) { std::cout<<"No infer action for this layer\n"; return NULL; } @@ -86,7 +86,7 @@ public: const char* fname_weights, const char* fname_bias); virtual ~Dense(); - value_type* infer(dataDim_t &dim, value_type* srcData); + virtual value_type* infer(dataDim_t &dim, value_type* srcData); protected: value_type *dstData; //where results will be putted @@ -111,7 +111,7 @@ public: Activation(Network *net, dataDim_t input_dim, tkdnnActivationMode_t act_mode); virtual ~Activation(); - value_type* infer(dataDim_t &dim, value_type* srcData); + virtual value_type* infer(dataDim_t &dim, value_type* srcData); protected: tkdnnActivationMode_t act_mode; @@ -130,7 +130,7 @@ public: const char* fname_weights, const char* fname_bias); virtual ~Conv2d(); - value_type* infer(dataDim_t &dim, value_type* srcData); + virtual value_type* infer(dataDim_t &dim, value_type* srcData); protected: value_type *dstData; //where results will be putted @@ -157,7 +157,7 @@ public: const char* fname_weights, const char* fname_bias); virtual ~Conv3d(); - value_type* infer(dataDim_t &dim, value_type* srcData); + virtual value_type* infer(dataDim_t &dim, value_type* srcData); protected: value_type *dstData; //where results will be putted @@ -185,7 +185,7 @@ public: Flatten(Network *net, dataDim_t input_dim); virtual ~Flatten(); - value_type* infer(dataDim_t &dim, value_type* srcData); + virtual value_type* infer(dataDim_t &dim, value_type* srcData); protected: value_type *dstData; //where results will be putted @@ -202,7 +202,7 @@ public: MulAdd(Network *net, dataDim_t input_dim, value_type mul, value_type add); virtual ~MulAdd(); - value_type* infer(dataDim_t &dim, value_type* srcData); + virtual value_type* infer(dataDim_t &dim, value_type* srcData); protected: value_type mul, add; @@ -231,7 +231,7 @@ public: int strideH, int strideW, tkdnnPoolingMode_t pool_mode); virtual ~Pooling(); - value_type* infer(dataDim_t &dim, value_type* srcData); + virtual value_type* infer(dataDim_t &dim, value_type* srcData); protected: