virtual method

This commit is contained in:
Francesco Gatti
2017-07-03 09:01:40 +00:00
parent 159aa8aae8
commit 33513468e6
+8 -8
View File
@@ -43,7 +43,7 @@ public:
Layer(Network *net, dataDim_t input_dim); Layer(Network *net, dataDim_t input_dim);
virtual ~Layer(); 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"; std::cout<<"No infer action for this layer\n";
return NULL; return NULL;
} }
@@ -86,7 +86,7 @@ public:
const char* fname_weights, const char* fname_bias); const char* fname_weights, const char* fname_bias);
virtual ~Dense(); virtual ~Dense();
value_type* infer(dataDim_t &dim, value_type* srcData); virtual value_type* infer(dataDim_t &dim, value_type* srcData);
protected: protected:
value_type *dstData; //where results will be putted value_type *dstData; //where results will be putted
@@ -111,7 +111,7 @@ public:
Activation(Network *net, dataDim_t input_dim, tkdnnActivationMode_t act_mode); Activation(Network *net, dataDim_t input_dim, tkdnnActivationMode_t act_mode);
virtual ~Activation(); virtual ~Activation();
value_type* infer(dataDim_t &dim, value_type* srcData); virtual value_type* infer(dataDim_t &dim, value_type* srcData);
protected: protected:
tkdnnActivationMode_t act_mode; tkdnnActivationMode_t act_mode;
@@ -130,7 +130,7 @@ public:
const char* fname_weights, const char* fname_bias); const char* fname_weights, const char* fname_bias);
virtual ~Conv2d(); virtual ~Conv2d();
value_type* infer(dataDim_t &dim, value_type* srcData); virtual value_type* infer(dataDim_t &dim, value_type* srcData);
protected: protected:
value_type *dstData; //where results will be putted value_type *dstData; //where results will be putted
@@ -157,7 +157,7 @@ public:
const char* fname_weights, const char* fname_bias); const char* fname_weights, const char* fname_bias);
virtual ~Conv3d(); virtual ~Conv3d();
value_type* infer(dataDim_t &dim, value_type* srcData); virtual value_type* infer(dataDim_t &dim, value_type* srcData);
protected: protected:
value_type *dstData; //where results will be putted value_type *dstData; //where results will be putted
@@ -185,7 +185,7 @@ public:
Flatten(Network *net, dataDim_t input_dim); Flatten(Network *net, dataDim_t input_dim);
virtual ~Flatten(); virtual ~Flatten();
value_type* infer(dataDim_t &dim, value_type* srcData); virtual value_type* infer(dataDim_t &dim, value_type* srcData);
protected: protected:
value_type *dstData; //where results will be putted 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); MulAdd(Network *net, dataDim_t input_dim, value_type mul, value_type add);
virtual ~MulAdd(); virtual ~MulAdd();
value_type* infer(dataDim_t &dim, value_type* srcData); virtual value_type* infer(dataDim_t &dim, value_type* srcData);
protected: protected:
value_type mul, add; value_type mul, add;
@@ -231,7 +231,7 @@ public:
int strideH, int strideW, tkdnnPoolingMode_t pool_mode); int strideH, int strideW, tkdnnPoolingMode_t pool_mode);
virtual ~Pooling(); virtual ~Pooling();
value_type* infer(dataDim_t &dim, value_type* srcData); virtual value_type* infer(dataDim_t &dim, value_type* srcData);
protected: protected: