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

y = x + inner->forward(x) for an arbitrary already-built Module. The classic ResNet shortcut connection, owning its own native relevance-split rule – the charter's explicitly named failure mode to avoid is Captum/Zennit's "Canonizer surgery" (an external post-hoc graph rewrite for residual connections). More...

#include <residual_module.hpp>

Inheritance diagram for pulsatrix::ResidualModule:
Collaboration diagram for pulsatrix::ResidualModule:

Public Member Functions

 ResidualModule (Module *inner, DeviceBackend *backend)
 Constructs a residual wrapper around an existing Module.
 
Tensor backward (const Tensor &grad_output) override
 Gradient w.r.t. this module's input: both paths receive grad_output unchanged (real gradient of a plain sum), then inner_'s own backward() adds its contribution.
 
OpType op_type () const override
 Elementwise per this module's own op_type() note above.
 
Tensor propagate_relevance (const Tensor &relevance_out, const LRPRuleConfig &config) override
 LRP relevance propagation: the residual epsilon/z-rule split between x and inner_->forward(x), then inner_'s own propagate_relevance() for its share.
 
std::vector< NamedParamRef > named_parameters () override
 inner_'s own named_parameters(), prefixed inner. – this module owns none of its own.
 
void set_training (bool training) override
 Cascades to inner_, the same way SequentialModule/MultiHeadAttentionModule do.
 
Module & inner ()
 
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).
 
bool is_training () const
 Whether this module is currently in training mode.
 

Protected Member Functions

Tensor forward_impl (const Tensor &input) override
 y = x + inner_->forward(x).
 

Detailed Description

y = x + inner->forward(x) for an arbitrary already-built Module. The classic ResNet shortcut connection, owning its own native relevance-split rule – the charter's explicitly named failure mode to avoid is Captum/Zennit's "Canonizer surgery" (an external post-hoc graph rewrite for residual connections).

Not a new LRP rule: reuses this codebase's own two-term weighted-sum epsilon/z-rule (LSTMModule's c_t = f_t*c_{t-1}+i_t*g_t, GRUModule's analogous carry-split, TransformerBlock's own two residual adds), weight fixed at 1 – resolved at Phase 3 activation, not re-derived here.

Note
inner_ is held by non-owning pointer – "not owned; must outlive this object", the same convention as SequentialModule's layers_. Any already-composed Module can be wrapped: a SequentialModule chaining Conv2DModule/ BatchNormModule/ReluModule for a real ResNet basic block, a bare Conv2DModule, a LinearModule, even another ResidualModule.
op_type() reuses OpType::Elementwise – this module's entire purpose is the residual add itself (fixed weight 1, computed independently per element), the same category TransformerBlock's own residual adds already use. inner_'s own op_type() is unaffected.
inner_->forward(x) must return the same shape as x – this module cannot validate that structurally in advance (it does not know inner_'s internals); a mismatch surfaces as Tensor's own shape-mismatch error at the add or at the relevance split.

Constructor & Destructor Documentation

◆ ResidualModule()

pulsatrix::ResidualModule::ResidualModule ( Module *  inner,
DeviceBackend *  backend 
)

Constructs a residual wrapper around an existing Module.

Parameters
innerThe wrapped function F. Not owned; must outlive this object.
backendBackend to allocate/compute through. Not owned; must outlive this module. Unlike inner, not null-checked – every DeviceBackend* member elsewhere in this codebase follows the same "not owned, must outlive, not validated" convention (LinearModule, RoPEModule, etc. never null-check their own backend argument either); inner is validated because, unlike backend, this constructor's member-initializer-list Tensor construction doesn't dereference it before any body-level check could run.
Exceptions
std::invalid_argumentif inner is null – external boundary, same convention as SequentialModule's null-entry check.

Member Function Documentation

◆ backward()

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

Gradient w.r.t. this module's input: both paths receive grad_output unchanged (real gradient of a plain sum), then inner_'s own backward() adds its contribution.

Parameters
grad_outputGradient w.r.t. this module's output, matching the cached forward shape.
Returns
Gradient w.r.t. this module's input, same shape.
Exceptions
std::logic_errorif called before any forward().
std::invalid_argumentif grad_output's shape differs from the cached forward shape.
Note
Device-generic: inner backward plus DeviceBackend::add (GPU-native-kernels Mission 1); runs on a GPU tensor whenever inner_ does.

Implements pulsatrix::Module.

◆ compute_device()

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

y = x + inner_->forward(x).

Parameters
inputAny shape inner_ accepts; inner_->forward(input) must return the same shape.
Returns
Same shape as input.

Implements pulsatrix::Module.

◆ inner()

Module & pulsatrix::ResidualModule::inner ( )
inline

◆ named_parameters()

std::vector< NamedParamRef > pulsatrix::ResidualModule::named_parameters ( )
overridevirtual

inner_'s own named_parameters(), prefixed inner. – this module owns none of its own.

Reimplemented from pulsatrix::Module.

◆ op_type()

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

Elementwise per this module's own op_type() note above.

Implements pulsatrix::Module.

◆ propagate_relevance()

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

LRP relevance propagation: the residual epsilon/z-rule split between x and inner_->forward(x), then inner_'s own propagate_relevance() for its share.

Parameters
relevance_outRelevance at this module's output, matching the cached forward shape.
configSupplies the epsilon stabilizer for the residual split and inner_'s own rule.
Returns
Relevance at this module's input, same shape as relevance_out.
Exceptions
std::logic_errorif called before any forward().
std::invalid_argumentif relevance_out's shape differs from the cached forward shape.
Note
Device-generic: runs on Cpu, Cuda or Hip tensors (GPU-native-kernels Mission 3).

Implements pulsatrix::Module.

◆ set_training()

void pulsatrix::ResidualModule::set_training ( bool  training)
overridevirtual

Cascades to inner_, the same way SequentialModule/MultiHeadAttentionModule do.

Reimplemented from pulsatrix::Module.


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