| 580 | } |
| 581 | |
| 582 | void batch_normalize_conv ( |
| 583 | const double eps, |
| 584 | resizable_tensor& dest, |
| 585 | resizable_tensor& means, |
| 586 | resizable_tensor& invstds, |
| 587 | const double averaging_factor, |
| 588 | resizable_tensor& running_means, |
| 589 | resizable_tensor& running_variances, |
| 590 | const tensor& src, |
| 591 | const tensor& gamma, |
| 592 | const tensor& beta |
| 593 | ) |
| 594 | { |
| 595 | DLIB_CASSERT(0 <= averaging_factor && averaging_factor <= 1, "averaging_factor: " << averaging_factor); |
| 596 | DLIB_CASSERT(averaging_factor==1 || have_same_dimensions(running_means,means)); |
| 597 | DLIB_CASSERT(averaging_factor==1 || have_same_dimensions(running_variances,invstds)); |
| 598 | DLIB_CASSERT( |
| 599 | src.num_samples() > 1 && |
| 600 | gamma.num_samples() == 1 && |
| 601 | beta.num_samples() == 1 && |
| 602 | gamma.nr() == 1 && |
| 603 | beta.nr() == 1 && |
| 604 | gamma.nc() == 1 && |
| 605 | beta.nc() == 1 && |
| 606 | gamma.k() == beta.k() && beta.k() == src.k() && |
| 607 | eps > 0, |
| 608 | "\ngamma.num_samples(): " << gamma.num_samples() << |
| 609 | "\ngamma.k(): " << gamma.k() << |
| 610 | "\ngamma.nr(): " << gamma.nr() << |
| 611 | "\ngamma.nc(): " << gamma.nc() << |
| 612 | "\nbeta.num_samples(): " << beta.num_samples() << |
| 613 | "\nbeta.k(): " << beta.k() << |
| 614 | "\nbeta.nr(): " << beta.nr() << |
| 615 | "\nbeta.nc(): " << beta.nc() << |
| 616 | "\nsrc.k(): " << src.k() << |
| 617 | "\nsrc.nr(): " << src.nr() << |
| 618 | "\nsrc.nc(): " << src.nc() << |
| 619 | "\neps: " << eps |
| 620 | ); |
| 621 | const float in_scale = 1; |
| 622 | const float out_scale = 0; |
| 623 | |
| 624 | dest.copy_size(src); |
| 625 | means.set_size(1, src.k()); |
| 626 | invstds.copy_size(means); |
| 627 | running_means.copy_size(means); |
| 628 | running_variances.copy_size(means); |
| 629 | // cuDNN requires that running_means and running_variances be initialized to |
| 630 | // some valid float values even if the averaging factor would have ignored |
| 631 | // them. |
| 632 | if (averaging_factor == 1) |
| 633 | { |
| 634 | running_means = 0; |
| 635 | running_variances = 1; |
| 636 | } |
| 637 | |
| 638 | CHECK_CUDNN(cudnnBatchNormalizationForwardTraining( |
| 639 | context(), |
nothing calls this directly
no test coverage detected