Merge pull request #1491 from ShikharJ/Conv

Improve Convolution Gradient() Method.
This commit is contained in:
Marcus Edel
2018-08-15 18:58:22 +02:00
committed by GitHub
3 changed files with 27 additions and 37 deletions
@@ -246,17 +246,8 @@ void AtrousConvolution<
arma::Mat<eT>&& error,
arma::Mat<eT>&& gradient)
{
arma::cube mappedError;
if (padW != 0 && padH != 0)
{
mappedError = arma::cube(error.memptr(), outputWidth / padW,
outputHeight / padH, outSize * batchSize, false, false);
}
else
{
mappedError = arma::cube(error.memptr(), outputWidth,
outputHeight, outSize * batchSize, false, false);
}
arma::cube mappedError(error.memptr(), outputWidth, outputHeight,
outSize * batchSize, false, false);
gradient.set_size(weights.n_elem, 1);
gradientTemp = arma::Cube<eT>(gradient.memptr(), weight.n_rows,
@@ -303,13 +294,17 @@ void AtrousConvolution<
}
}
if ((padW != 0 || padH != 0) &&
(gradientTemp.n_rows < output.n_rows &&
gradientTemp.n_cols < output.n_cols))
if (gradientTemp.n_rows < output.n_rows ||
gradientTemp.n_cols < output.n_cols)
{
gradientTemp.slice(outMapIdx) += output.submat(padW, padH,
padW + gradientTemp.n_rows - 1,
padH + gradientTemp.n_cols - 1);
gradientTemp.slice(outMapIdx) += output.submat(0, 0,
gradientTemp.n_rows - 1, gradientTemp.n_cols - 1);
}
else if (gradientTemp.n_rows > output.n_rows ||
gradientTemp.n_cols > output.n_cols)
{
gradientTemp.slice(outMapIdx).submat(0, 0, output.n_rows - 1,
output.n_cols - 1) += output;
}
else
{
@@ -239,17 +239,8 @@ void Convolution<
arma::Mat<eT>&& error,
arma::Mat<eT>&& gradient)
{
arma::cube mappedError;
if (padW != 0 && padH != 0)
{
mappedError = arma::cube(error.memptr(), outputWidth / padW,
outputHeight / padH, outSize * batchSize, false, false);
}
else
{
mappedError = arma::cube(error.memptr(), outputWidth,
outputHeight, outSize * batchSize, false, false);
}
arma::cube mappedError(error.memptr(), outputWidth,
outputHeight, outSize * batchSize, false, false);
gradient.set_size(weights.n_elem, 1);
gradientTemp = arma::Cube<eT>(gradient.memptr(), weight.n_rows,
@@ -283,13 +274,17 @@ void Convolution<
GradientConvolutionRule::Convolution(inputSlice, deltaSlice,
output, dW, dH);
if ((padW != 0 || padH != 0) &&
(gradientTemp.n_rows < output.n_rows &&
gradientTemp.n_cols < output.n_cols))
if (gradientTemp.n_rows < output.n_rows ||
gradientTemp.n_cols < output.n_cols)
{
gradientTemp.slice(outMapIdx) += output.submat(padW, padH,
padW + gradientTemp.n_rows - 1,
padH + gradientTemp.n_cols - 1);
gradientTemp.slice(outMapIdx) += output.submat(0, 0,
gradientTemp.n_rows - 1, gradientTemp.n_cols - 1);
}
else if (gradientTemp.n_rows > output.n_rows ||
gradientTemp.n_cols > output.n_cols)
{
gradientTemp.slice(outMapIdx).submat(0, 0, output.n_rows - 1,
output.n_cols - 1) += output;
}
else
{
@@ -130,15 +130,15 @@ class Reparametrization
//! Locally-stored number of output units.
size_t latentSize;
//! The beta hyperparameter for constrained variational frameworks.
double beta;
//! If false, sample will be constant.
bool stochastic;
//! If false, KL error will not be included in Backward function.
bool includeKl;
//! The beta hyperparameter for constrained variational frameworks.
double beta;
//! Locally-stored delta object.
OutputDataType delta;