ClikaRT::nn::BatchNorm
class
Header: ClikaRT/nn/batch_norm.h
Inherits: ClikaRT::nn::Module
Batch normalization, inference form.
Every operand is per channel ([C]), the channel axis is the last one, and the stored statistics ARE the statistics (no batch statistics, no momentum). The module form of ops::batch_norm.
auto bn = ClikaRT::nn::BatchNorm::make(64);
bn->load_state_dict(checkpoint); // weight, bias, running_mean, ...
auto y = (*bn)(x); // [N, H, W, 64] -> [N, H, W, 64]
Static member functions
make()
static std::shared_ptr<BatchNorm> make(std::int64_t num_features, BatchNormOptions options = {})
Construct from config with storage-free slots: registers the [num_features] parameter slots (weight, bias, when options.affine) and buffer slots (running_mean, running_var, plus the num_batches_tracked counter) on options.dtype / options.device; no weight memory is allocated. named_parameters() and named_buffers() report the slots, so a strict load_state_dict both expects and satisfies them; the first forward builds. Raises ClikaRT::Error on a channel count or eps that is not positive.
Declared in ClikaRT/nn/batch_norm.h, line 92
Member functions
set_weights()
void set_weights(Tensor weight, OptionalTensor bias = {})
Bind the affine parameters positionally: weight (scale, [C]) and bias (shift, [C]). The declared placement wins; shapes must match the declaration. Re-binding drops the built primitive; the next forward rebuilds. Raises ClikaRT::Error on a geometry mismatch or on a module made with affine = false.
Declared in ClikaRT/nn/batch_norm.h, line 103
set_running_stats()
Bind the statistics positionally: running_mean and running_var, both [C]. Same placement and shape rules as set_weights; the next forward rebuilds. Raises ClikaRT::Error on a geometry mismatch.
Declared in ClikaRT/nn/batch_norm.h, line 112
~BatchNorm()
~BatchNorm() override
Declared in ClikaRT/nn/batch_norm.h, line 116
to_impl(StreamOrDevice)
virtual Result<void> to_impl(StreamOrDevice where) override
Move to a placement (Device / Stream) or cast to a dtype. Rebuilds over the moved slots on the next forward.
Declared in ClikaRT/nn/batch_norm.h, line 120
to_impl(DataType)
Declared in ClikaRT/nn/batch_norm.h, line 121
forward()
Normalizes input [N, D1..Dn, C] with the stored statistics and the affine parameters, then the fused activation when set; output shape and dtype follow input. The first call builds (once, thread-safe). Raises ClikaRT::Error when a slot is still unbound or the channel count does not match.
Declared in ClikaRT/nn/batch_norm.h, line 129
operator()()
Same as forward.
Declared in ClikaRT/nn/batch_norm.h, line 133
num_features()
std::int64_t num_features() const noexcept
The channel count the module was made with.
Declared in ClikaRT/nn/batch_norm.h, line 138
options()
const BatchNormOptions& options() const noexcept
The settings the module was made with.
Declared in ClikaRT/nn/batch_norm.h, line 140
Protected member functions
initialize_impl()
virtual Result<void> initialize_impl() override
The pack hook (Module::initialize_impl): builds the normalization over the bound slots (idempotent, thread-safe). Every declared slot must hold a real (loaded) tensor; a still-unbound slot is a clean error naming it.
Declared in ClikaRT/nn/batch_norm.h, line 151