| 776 | } |
| 777 | |
| 778 | void batch_normalize ( |
| 779 | const double eps, |
| 780 | resizable_tensor& dest, |
| 781 | resizable_tensor& means, |
| 782 | resizable_tensor& invstds, |
| 783 | const double averaging_factor, |
| 784 | resizable_tensor& running_means, |
| 785 | resizable_tensor& running_variances, |
| 786 | const tensor& src, |
| 787 | const tensor& gamma, |
| 788 | const tensor& beta |
| 789 | ) |
| 790 | { |
| 791 | DLIB_CASSERT(0 <= averaging_factor && averaging_factor <= 1, "averaging_factor: " << averaging_factor); |
| 792 | DLIB_CASSERT(averaging_factor==1 || have_same_dimensions(running_means,means)); |
| 793 | DLIB_CASSERT(averaging_factor==1 || have_same_dimensions(running_variances,invstds)); |
| 794 | DLIB_CASSERT( |
| 795 | src.num_samples() > 1 && |
| 796 | gamma.num_samples() == 1 && |
| 797 | beta.num_samples() == 1 && |
| 798 | gamma.nr() == beta.nr() && beta.nr() == src.nr() && |
| 799 | gamma.nc() == beta.nc() && beta.nc() == src.nc() && |
| 800 | gamma.k() == beta.k() && beta.k() == src.k() && |
| 801 | eps > 0, |
| 802 | "\ngamma.num_samples(): " << gamma.num_samples() << |
| 803 | "\ngamma.k(): " << gamma.k() << |
| 804 | "\ngamma.nr(): " << gamma.nr() << |
| 805 | "\ngamma.nc(): " << gamma.nc() << |
| 806 | "\nbeta.num_samples(): " << beta.num_samples() << |
| 807 | "\nbeta.k(): " << beta.k() << |
| 808 | "\nbeta.nr(): " << beta.nr() << |
| 809 | "\nbeta.nc(): " << beta.nc() << |
| 810 | "\nsrc.k(): " << src.k() << |
| 811 | "\nsrc.nr(): " << src.nr() << |
| 812 | "\nsrc.nc(): " << src.nc() << |
| 813 | "\neps: " << eps |
| 814 | ); |
| 815 | |
| 816 | dest.copy_size(src); |
| 817 | means.set_size(1, src.k(), src.nr(), src.nc()); |
| 818 | invstds.set_size(1, src.k(), src.nr(), src.nc()); |
| 819 | running_means.set_size(1, src.k(), src.nr(), src.nc()); |
| 820 | running_variances.set_size(1, src.k(), src.nr(), src.nc()); |
| 821 | |
| 822 | // first compute means and invstds |
| 823 | const auto p_invstds = invstds.host(); |
| 824 | const auto p_means = means.host(); |
| 825 | auto p_src = src.host(); |
| 826 | const auto rvar = running_variances.host(); |
| 827 | const long num = src.k()*src.nr()*src.nc(); |
| 828 | |
| 829 | // This scale makes the running variances unbiased. |
| 830 | const double scale = (src.num_samples())/(src.num_samples()-1.0); |
| 831 | |
| 832 | // Apply Welford's algorithm to improve numerical stability |
| 833 | for (long i = 0; i < num; ++i) |
| 834 | { |
| 835 | double mean = 0.0; |