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

While alive, merges bn's affine map into conv's weights and makes bn an exact identity; on destruction, restores both bit for bit. More...

#include <batch_norm_fold.hpp>

Public Member Functions

 BatchNormFold (Conv2DModule &conv, BatchNormModule &bn)
 
 ~BatchNormFold ()
 
 BatchNormFold (const BatchNormFold &)=delete
 
BatchNormFold & operator= (const BatchNormFold &)=delete
 
 BatchNormFold (BatchNormFold &&)=delete
 
BatchNormFold & operator= (BatchNormFold &&)=delete
 

Detailed Description

While alive, merges bn's affine map into conv's weights and makes bn an exact identity; on destruction, restores both bit for bit.

Note
With s_c = gamma_c / sqrt(running_var_c + eps), the merged convolution has kernel w'[c] = s_c * w[c] and bias b'[c] = s_c * (b[c] - running_mean_c) + beta_c. The pair computes the same function before and during the fold (up to float rounding), but LRP now distributes relevance through one affine layer with the convolution's own rule, instead of passing it unchanged through BatchNorm's identity rule.
Use it as a scope around explaining: { BatchNormFold fold(conv, bn); auto r = lrp.explain(...); }. Don't train while a fold is active: gradients would update the merged weights, and the restore would overwrite those updates.

Constructor & Destructor Documentation

◆ BatchNormFold() [1/3]

pulsatrix::BatchNormFold::BatchNormFold ( Conv2DModule &  conv,
BatchNormModule &  bn 
)
Exceptions
std::invalid_argumentif bn is in training mode (its statistics are not fixed, so there is no single affine map to fold), or if bn's channel count differs from conv's output channels.
std::logic_errorif bn is already folded.

◆ ~BatchNormFold()

pulsatrix::BatchNormFold::~BatchNormFold ( )

◆ BatchNormFold() [2/3]

pulsatrix::BatchNormFold::BatchNormFold ( const BatchNormFold &  )
delete

◆ BatchNormFold() [3/3]

pulsatrix::BatchNormFold::BatchNormFold ( BatchNormFold &&  )
delete

Member Function Documentation

◆ operator=() [1/2]

BatchNormFold & pulsatrix::BatchNormFold::operator= ( BatchNormFold &&  )
delete

◆ operator=() [2/2]

BatchNormFold & pulsatrix::BatchNormFold::operator= ( const BatchNormFold &  )
delete

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