pulsatrix
Loading...
Searching...
No Matches
pulsatrix::BatchNormModule Class Reference

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 every (n, h, w) element jointly – BatchNorm's defining statistic, and the reason this module didn't exist before campaign_exai_dl_library_batch_dimension_support: it has nothing to compute over without a real batch dimension. Input/output are rank-4 (N, channels, H, W), the same convention Conv2DModule/GroupNormModule already establish. More...

#include <batch_norm_module.hpp>

Inheritance diagram for pulsatrix::BatchNormModule:
Collaboration diagram for pulsatrix::BatchNormModule:

Public Member Functions

 BatchNormModule (int64_t num_channels, DeviceBackend *backend, DeviceType device, float eps=1e-6f, float momentum=0.1f)
 Constructs a BatchNorm layer with zero-initialized gamma and beta.
 
 BatchNormModule (int64_t num_channels, DeviceBackend *backend)
 On backend's own device (backend->device()), default eps. Previously the device defaulted to Cpu regardless of backend (GPU-native-kernels Mission 0).
 
Tensor backward (const Tensor &grad_output) override
 Computes the gradient w.r.t. this module's input, and accumulates gamma's/ beta's gradients internally.
 
OpType op_type () const override
 Normalization per charter's closed OpType set.
 
void set_gamma (std::initializer_list< float > values)
 Overwrites the per-channel gamma buffer – test/initialization use only.
 
void set_beta (std::initializer_list< float > values)
 Overwrites the per-channel beta buffer – test/initialization use only.
 
void set_gamma (const std::vector< float > &values)
 Vector overload for runtime-sized sources – see Tensor's own vector ctor.
 
void set_beta (const std::vector< float > &values)
 Vector overload for runtime-sized sources – see Tensor's own vector ctor.
 
const Tensor & running_mean () const
 Per-channel running mean, used in eval mode. Shape (num_channels).
 
const Tensor & running_var () const
 Per-channel running variance, used in eval mode. Shape (num_channels).
 
void set_running_mean (const std::vector< float > &values)
 Overwrites the running mean, e.g. when loading a pretrained model.
 
void set_running_var (const std::vector< float > &values)
 Overwrites the running variance, e.g. when loading a pretrained model.
 
const Tensor & gamma () const
 
const Tensor & beta () const
 
const Tensor & gamma_grad () const
 
const Tensor & beta_grad () const
 
Tensor propagate_relevance (const Tensor &relevance_out, const LRPRuleConfig &config) override
 Identity-rule LRP relevance propagation (AttnLRP, Achtibat et al. 2024).
 
std::vector< NamedParamRef > named_parameters () override
 This module's trainable parameters, each with its hierarchical name – the one place a module declares its parameters (roadmap FND-1).
 
std::optional< DeviceType > compute_device () const override
 Where this layer computes, so forward() rejects an input on another device (FND-8).
 
- Public Member Functions inherited from pulsatrix::Module
virtual ~Module ()=default
 
Tensor forward (const Tensor &input)
 Runs this module's forward computation.
 
std::pair< Tensor, NodeId > forward_traced (const Tensor &input, NodeId input_node, ComputationGraph &graph, Autograd &autograd)
 Runs forward() while also registering a ComputationGraph node (tagged with this module's op_type(), parented to input_node) and wiring an Autograd backward function that reuses this module's own backward() – the opt-in traced/explainable path, per Phase 2 Mission 0.
 
virtual bool supports_lrp_rule (LRPRule rule) const
 Whether propagate_relevance() implements rule (no silent fallback: callers such as ExplainerContext::relevance_pass() throw rather than run a module on a rule it does not implement).
 
virtual std::vector< ParamRef > parameters ()
 This module's trainable parameters and their gradients, for an optimizer to update uniformly across module types.
 
void set_requires_grad (bool requires_grad, const std::string &prefix="")
 Freezes (false) or unfreezes (true) parameters by name (roadmap FND-2).
 
virtual void set_training (bool training)
 Sets this module's training/eval mode. Defaults to training (matches every mainstream framework's Module default).
 
bool is_training () const
 Whether this module is currently in training mode.
 

Protected Member Functions

Tensor forward_impl (const Tensor &input) override
 The actual forward computation. Called by forward() after precondition checks.
 

Friends

class BatchNormFold
 

Detailed Description

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 every (n, h, w) element jointly – BatchNorm's defining statistic, and the reason this module didn't exist before campaign_exai_dl_library_batch_dimension_support: it has nothing to compute over without a real batch dimension. Input/output are rank-4 (N, channels, H, W), the same convention Conv2DModule/GroupNormModule already establish.

Note
propagate_relevance is an identity pass-through, cited to AttnLRP's normalization-layer treatment (Achtibat et al. 2024) – the same rule, same citation, RMSNormModule/LayerNormModule/GroupNormModule already use; BatchNorm is architecturally the same normalization category, just a different statistic grouping (channel-over-batch-and-spatial instead of group-over-spatial-per-row). backward() is the real, undetached training gradient.
Training mode (the Module default) normalizes with the current batch's statistics and folds them into running statistics with PyTorch's rule: running = (1 - momentum) * running + momentum * batch, using the unbiased batch variance. Eval mode (set_training(false)) normalizes with the running statistics instead, so each sample's output depends only on that sample (roadmap FND-5, lrp_issues #8). Running statistics start at mean 0, variance 1. Put the model in eval mode before explaining it.
For LRP, fold an eval-mode BatchNorm into the Conv2D before it with BatchNormFold: the convolution's rule then distributes relevance through the combined affine map, and this module becomes an exact identity.

Constructor & Destructor Documentation

◆ BatchNormModule() [1/2]

pulsatrix::BatchNormModule::BatchNormModule ( int64_t  num_channels,
DeviceBackend *  backend,
DeviceType  device,
float  eps = 1e-6f,
float  momentum = 0.1f 
)

Constructs a BatchNorm layer with zero-initialized gamma and beta.

Parameters
num_channelsNumber of channels (the statistic-bearing dimension).
backendBackend to allocate/compute through. Not owned; must outlive this module.
deviceWhich device every internal Tensor member is tagged as. Defaults to Cpu.
epsStabilizer added inside the sqrt. Defaults to 1e-6, matching RMSNormModule/LayerNormModule/GroupNormModule's default for consistency within this codebase's normalization family.
momentumWeight of each new batch in the running statistics, in (0, 1]. Defaults to 0.1, PyTorch's default.
Exceptions
std::invalid_argumentif num_channels <= 0 or momentum is outside (0, 1] – external boundary (construction arguments can originate from Phase 5's Python bindings with no upstream validation), per cpp_tdd/context_tdd_adversarial_boundary_testing.md.

◆ BatchNormModule() [2/2]

pulsatrix::BatchNormModule::BatchNormModule ( int64_t  num_channels,
DeviceBackend *  backend 
)

On backend's own device (backend->device()), default eps. Previously the device defaulted to Cpu regardless of backend (GPU-native-kernels Mission 0).

Member Function Documentation

◆ backward()

Tensor pulsatrix::BatchNormModule::backward ( const Tensor &  grad_output)
overridevirtual

Computes the gradient w.r.t. this module's input, and accumulates gamma's/ beta's gradients internally.

Parameters
grad_outputGradient w.r.t. this module's output. Must match the shape of the most recent forward() call's output.
Returns
Gradient w.r.t. this module's input.
Exceptions
std::logic_errorif forward() has never been called.
Note
Device-generic: runs on Cpu, Cuda or Hip tensors (GPU-native-kernels Mission 4).

Implements pulsatrix::Module.

◆ beta()

const Tensor & pulsatrix::BatchNormModule::beta ( ) const
inline

◆ beta_grad()

const Tensor & pulsatrix::BatchNormModule::beta_grad ( ) const
inline

◆ compute_device()

std::optional< DeviceType > pulsatrix::BatchNormModule::compute_device ( ) const
inlineoverridevirtual

Where this layer computes, so forward() rejects an input on another device (FND-8).

Reimplemented from pulsatrix::Module.

◆ forward_impl()

Tensor pulsatrix::BatchNormModule::forward_impl ( const Tensor &  input)
overrideprotectedvirtual

The actual forward computation. Called by forward() after precondition checks.

Implements pulsatrix::Module.

◆ gamma()

const Tensor & pulsatrix::BatchNormModule::gamma ( ) const
inline

◆ gamma_grad()

const Tensor & pulsatrix::BatchNormModule::gamma_grad ( ) const
inline

◆ named_parameters()

std::vector< NamedParamRef > pulsatrix::BatchNormModule::named_parameters ( )
inlineoverridevirtual

This module's trainable parameters, each with its hierarchical name – the one place a module declares its parameters (roadmap FND-1).

Returns
{name, {value, grad}} entries pointing directly at this module's own members, in a fixed order. Names are unique within the module tree. Default: empty (a parameterless module like ReluModule needs no override).
Note
Override this, not parameters(): saving, loading, freezing by name and optimizer parameter groups all key on these names.

Reimplemented from pulsatrix::Module.

◆ op_type()

OpType pulsatrix::BatchNormModule::op_type ( ) const
inlineoverridevirtual

Normalization per charter's closed OpType set.

Implements pulsatrix::Module.

◆ propagate_relevance()

Tensor pulsatrix::BatchNormModule::propagate_relevance ( const Tensor &  relevance_out,
const LRPRuleConfig &  config 
)
overridevirtual

Identity-rule LRP relevance propagation (AttnLRP, Achtibat et al. 2024).

Parameters
relevance_outRelevance at this module's output. Must match the shape of the most recent forward() call's output.
configUnused – the identity rule has no tunable parameter.
Returns
relevance_out, unchanged – conservation holds trivially by construction.
Exceptions
std::logic_errorif forward() has never been called.
std::invalid_argumentif relevance_out's element count doesn't match the cached forward output's element count.

Implements pulsatrix::Module.

◆ running_mean()

const Tensor & pulsatrix::BatchNormModule::running_mean ( ) const
inline

Per-channel running mean, used in eval mode. Shape (num_channels).

◆ running_var()

const Tensor & pulsatrix::BatchNormModule::running_var ( ) const
inline

Per-channel running variance, used in eval mode. Shape (num_channels).

◆ set_beta() [1/2]

void pulsatrix::BatchNormModule::set_beta ( const std::vector< float > &  values)

Vector overload for runtime-sized sources – see Tensor's own vector ctor.

◆ set_beta() [2/2]

void pulsatrix::BatchNormModule::set_beta ( std::initializer_list< float >  values)

Overwrites the per-channel beta buffer – test/initialization use only.

◆ set_gamma() [1/2]

void pulsatrix::BatchNormModule::set_gamma ( const std::vector< float > &  values)

Vector overload for runtime-sized sources – see Tensor's own vector ctor.

◆ set_gamma() [2/2]

void pulsatrix::BatchNormModule::set_gamma ( std::initializer_list< float >  values)

Overwrites the per-channel gamma buffer – test/initialization use only.

◆ set_running_mean()

void pulsatrix::BatchNormModule::set_running_mean ( const std::vector< float > &  values)

Overwrites the running mean, e.g. when loading a pretrained model.

Exceptions
std::invalid_argumenton a size other than num_channels or a non-finite value.

◆ set_running_var()

void pulsatrix::BatchNormModule::set_running_var ( const std::vector< float > &  values)

Overwrites the running variance, e.g. when loading a pretrained model.

Exceptions
std::invalid_argumenton a size other than num_channels, or a negative or non-finite value.

Friends And Related Symbol Documentation

◆ BatchNormFold

friend class BatchNormFold
friend

The documentation for this class was generated from the following file: