upsample ok
This commit is contained in:
+7
-5
@@ -8,14 +8,14 @@ namespace tk { namespace dnn {
|
||||
Upsample::Upsample(Network *net, int stride) : Layer(net) {
|
||||
|
||||
this->stride = stride;
|
||||
|
||||
|
||||
output_dim.n = input_dim.n;
|
||||
output_dim.c = input_dim.c*stride*stride;
|
||||
output_dim.h = input_dim.h/stride;
|
||||
output_dim.w = input_dim.w/stride;
|
||||
output_dim.c = input_dim.c;
|
||||
output_dim.h = input_dim.h*stride;
|
||||
output_dim.w = input_dim.w*stride;
|
||||
output_dim.l = input_dim.l;
|
||||
|
||||
checkCuda( cudaMalloc(&dstData, input_dim.tot()*sizeof(dnnType)) );
|
||||
checkCuda( cudaMalloc(&dstData, output_dim.tot()*sizeof(dnnType)) );
|
||||
}
|
||||
|
||||
Upsample::~Upsample() {
|
||||
@@ -25,6 +25,8 @@ Upsample::~Upsample() {
|
||||
|
||||
dnnType* Upsample::infer(dataDim_t &dim, dnnType* srcData) {
|
||||
|
||||
fill(dstData, output_dim.tot(), 0.0);
|
||||
upsampleForward(srcData, dstData, input_dim.n, input_dim.c, input_dim.h, input_dim.w, stride, 1, 1);
|
||||
|
||||
dim = output_dim;
|
||||
return dstData;
|
||||
|
||||
Reference in New Issue
Block a user