diff --git a/src/mlpack/methods/ann/layer/pixel_shuffle_impl.hpp b/src/mlpack/methods/ann/layer/pixel_shuffle_impl.hpp index 9edb09c2eb..89ae8b79d5 100644 --- a/src/mlpack/methods/ann/layer/pixel_shuffle_impl.hpp +++ b/src/mlpack/methods/ann/layer/pixel_shuffle_impl.hpp @@ -61,12 +61,10 @@ void PixelShuffle::Forward( output.zeros(outputHeight * outputWidth * sizeOut, batchSize); for (size_t n = 0; n < batchSize; n++) { - arma::mat inputImage = input.col(n); - arma::mat outputImage = output.col(n); - arma::cube inputTemp(const_cast(inputImage).memptr(), height, - width, size, false, false); - arma::cube outputTemp(const_cast(outputImage).memptr(), - outputHeight, outputWidth, sizeOut, false, false); + arma::cube inputTemp(const_cast(input).memptr(), height, + width, size * batchSize, false, false); + arma::cube outputTemp(const_cast(output).memptr(), + outputHeight, outputWidth, sizeOut * batchSize, false, false); for (size_t c = 0; c < sizeOut ; c++) { @@ -78,13 +76,12 @@ void PixelShuffle::Forward( size_t width_index = w / upscaleFactor; size_t channel_index = (upscaleFactor * (h % upscaleFactor)) + (w % upscaleFactor) + (c * std::pow(upscaleFactor, 2)); - outputTemp(w, h, c) = inputTemp(width_index, height_index, - channel_index); + outputTemp(w, h, c + n * sizeOut) = inputTemp(width_index, height_index, + channel_index + n * size); } } } - output.col(n) = outputImage; } }