pulsatrix
Loading...
Searching...
No Matches
batch_norm_fold.hpp
Go to the documentation of this file.
1
6#pragma once
7
8#include <vector>
9
12
13namespace pulsatrix {
14
29public:
38
39 BatchNormFold(const BatchNormFold&) = delete;
43
44private:
45 Conv2DModule& conv_;
46 BatchNormModule& bn_;
47 std::vector<float> original_kernel_;
48 std::vector<float> original_bias_;
49};
50
51} // namespace pulsatrix
Batch normalization (Ioffe & Szegedy, 2015) – per-channel statistics computed across the batch and sp...
While alive, merges bn's affine map into conv's weights and makes bn an exact identity; on destructio...
Definition batch_norm_fold.hpp:28
BatchNormFold & operator=(const BatchNormFold &)=delete
BatchNormFold(const BatchNormFold &)=delete
BatchNormFold & operator=(BatchNormFold &&)=delete
BatchNormFold(Conv2DModule &conv, BatchNormModule &bn)
BatchNormFold(BatchNormFold &&)=delete
y_{n,c,h,w} = gamma_c * (x_{n,c,h,w} - mu_c)/std_c + beta_c, mu_c/std_c computed per channel c over e...
Definition batch_norm_module.hpp:42
2D convolution, batched (input/output are rank-4: N x channels x H x W) – migrated from the original ...
Definition conv2d_module.hpp:31
2D convolution – implemented via im2col + DeviceBackend::gemm (no new backend primitive).
Definition acquisition_functions.hpp:16