| 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 | } |