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

Max pooling, rank-4 (N, channels, H, W), matching Conv2DModule's convention. Stride fixed equal to kernel size (non-overlapping windows), no padding, no dilation – deferred until a real use case needs them, same minimal-cut discipline as Conv2DModule's original stride-1/no-padding scope cut. More...

#include <max_pool2d_module.hpp>

Inheritance diagram for pulsatrix::MaxPool2DModule:
Collaboration diagram for pulsatrix::MaxPool2DModule:

Public Member Functions

 MaxPool2DModule (int64_t kernel_h, int64_t kernel_w, DeviceBackend *backend)
 Constructs a max-pool layer.
 
Tensor backward (const Tensor &grad_output) override
 Computes the gradient w.r.t. this module's input – only the cached argmax position within each window receives grad_output; every other position is 0.
 
OpType op_type () const override
 Pooling per charter's closed OpType set.
 
Tensor propagate_relevance (const Tensor &relevance_out, const LRPRuleConfig &config) override
 Winner-take-all LRP relevance propagation (Bach et al. 2015).
 
bool supports_lrp_rule (LRPRule) const override
 Winner-take-all ignores the config: the same under every rule, so supports all of them.
 
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 std::vector< NamedParamRef > named_parameters ()
 This module's trainable parameters, each with its hierarchical name – the one place a module declares its parameters (roadmap FND-1).
 
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 – per-window max, argmax cached per output element for backward()/propagate_relevance() to reuse.
 

Detailed Description

Max pooling, rank-4 (N, channels, H, W), matching Conv2DModule's convention. Stride fixed equal to kernel size (non-overlapping windows), no padding, no dilation – deferred until a real use case needs them, same minimal-cut discipline as Conv2DModule's original stride-1/no-padding scope cut.

Note
propagate_relevance is winner-take-all (Bach et al. 2015's supplementary treatment of max-pooling, the same rule iNNvestigate/zennit ship): all relevance at an output position flows to the single input position that was the argmax in forward(); every other position in that window gets zero. Conserves exactly by construction. backward() routes gradient the same way (only the argmax position receives grad_output; standard max-pool gradient semantics).

Constructor & Destructor Documentation

◆ MaxPool2DModule()

pulsatrix::MaxPool2DModule::MaxPool2DModule ( int64_t  kernel_h,
int64_t  kernel_w,
DeviceBackend *  backend 
)

Constructs a max-pool layer.

Parameters
kernel_hWindow height. Also the vertical stride (non-overlapping).
kernel_wWindow width. Also the horizontal stride (non-overlapping).
backendBackend to allocate/compute through. Not owned; must outlive this module.
Exceptions
std::invalid_argumentif kernel_h <= 0 or kernel_w <= 0 – 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.

Member Function Documentation

◆ backward()

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

Computes the gradient w.r.t. this module's input – only the cached argmax position within each window receives grad_output; every other position is 0.

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.
std::invalid_argumentif grad_output's shape doesn't match the cached forward output shape.
Note
Device-generic: runs on Cpu, Cuda or Hip tensors (GPU-native-kernels Mission 4).

Implements pulsatrix::Module.

◆ compute_device()

std::optional< DeviceType > pulsatrix::MaxPool2DModule::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::MaxPool2DModule::forward_impl ( const Tensor &  input)
overrideprotectedvirtual

The actual forward computation – per-window max, argmax cached per output element for backward()/propagate_relevance() to reuse.

Exceptions
std::invalid_argumentif input isn't rank-4 (N, C, H, W), or the kernel is larger than the input (kernel_h > H or kernel_w > W).
Note
Device-generic: runs on Cpu, Cuda or Hip tensors (GPU-native-kernels Mission 4).

Implements pulsatrix::Module.

◆ op_type()

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

Pooling per charter's closed OpType set.

Implements pulsatrix::Module.

◆ propagate_relevance()

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

Winner-take-all LRP relevance propagation (Bach et al. 2015).

Parameters
relevance_outRelevance at this module's output. Must match the shape of the most recent forward() call's output.
configUnused – winner-take-all has no tunable parameter.
Returns
Relevance at this module's input: relevance_out's value at each window's cached argmax position, zero everywhere else. Conserves trivially by construction (no epsilon stabilizer needed).
Exceptions
std::logic_errorif forward() has never been called.
std::invalid_argumentif relevance_out's shape doesn't match the cached forward output shape.

Implements pulsatrix::Module.

◆ supports_lrp_rule()

bool pulsatrix::MaxPool2DModule::supports_lrp_rule ( LRPRule  ) const
inlineoverridevirtual

Winner-take-all ignores the config: the same under every rule, so supports all of them.

Reimplemented from pulsatrix::Module.


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