MCPcopy Create free account
hub / github.com/davisking/dlib / batch_normalize_inference

Function batch_normalize_inference

dlib/cuda/cpu_dlib.cpp:717–776  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

715 // -----------------------------------------------------------------------------------
716
717 void batch_normalize_inference (
718 const double eps,
719 resizable_tensor& dest,
720 const tensor& src,
721 const tensor& gamma,
722 const tensor& beta,
723 const tensor& running_means,
724 const tensor& running_variances
725 )
726 {
727 DLIB_CASSERT(
728 gamma.num_samples() == 1 &&
729 gamma.nr() == src.nr() &&
730 gamma.nc() == src.nc() &&
731 gamma.k() == src.k() &&
732 have_same_dimensions(gamma, beta) &&
733 have_same_dimensions(gamma, running_means) &&
734 have_same_dimensions(gamma, running_variances) &&
735 eps > 0,
736 "\ngamma.num_samples(): " << gamma.num_samples() <<
737 "\ngamma.k(): " << gamma.k() <<
738 "\ngamma.nr(): " << gamma.nr() <<
739 "\ngamma.nc(): " << gamma.nc() <<
740 "\nbeta.num_samples(): " << beta.num_samples() <<
741 "\nbeta.k(): " << beta.k() <<
742 "\nbeta.nr(): " << beta.nr() <<
743 "\nbeta.nc(): " << beta.nc() <<
744 "\nrunning_means.num_samples(): " << running_means.num_samples() <<
745 "\nrunning_means.k(): " << running_means.k() <<
746 "\nrunning_means.nr(): " << running_means.nr() <<
747 "\nrunning_means.nc(): " << running_means.nc() <<
748 "\nrunning_variances.num_samples(): " << running_variances.num_samples() <<
749 "\nrunning_variances.k(): " << running_variances.k() <<
750 "\nrunning_variances.nr(): " << running_variances.nr() <<
751 "\nrunning_variances.nc(): " << running_variances.nc() <<
752 "\nsrc.k(): " << src.k() <<
753 "\nsrc.nr(): " << src.nr() <<
754 "\nsrc.nc(): " << src.nc() <<
755 "\neps: " << eps
756 );
757 dest.copy_size(src);
758
759 auto d = dest.host();
760 auto s = src.host();
761 auto g = gamma.host();
762 auto b = beta.host();
763 auto m = running_means.host();
764 auto v = running_variances.host();
765
766 const long num = src.k()*src.nr()*src.nc();
767 for (long n = 0; n < src.num_samples(); ++n)
768 {
769 for (long k = 0; k < num; ++k)
770 {
771 *d = g[k]*(*s - m[k])/std::sqrt(v[k]+eps) + b[k];
772 ++d;
773 ++s;
774 }

Callers 2

test_batch_normalizeFunction · 0.50
forwardMethod · 0.50

Calls 8

copy_sizeMethod · 0.80
have_same_dimensionsFunction · 0.70
sqrtFunction · 0.50
num_samplesMethod · 0.45
nrMethod · 0.45
ncMethod · 0.45
kMethod · 0.45
hostMethod · 0.45

Tested by 1

test_batch_normalizeFunction · 0.40