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

y = reshape(x, {N, x.numel()/N}), N = x.shape().dim(0). No parameters, no gradient math beyond reshaping. More...

#include <flatten_module.hpp>

Inheritance diagram for pulsatrix::FlattenModule:
Collaboration diagram for pulsatrix::FlattenModule:

Public Member Functions

 FlattenModule (DeviceBackend *backend)
 Constructs a flatten module.
 
Tensor backward (const Tensor &grad_output) override
 Reshapes the gradient back to the shape forward() last saw.
 
OpType op_type () const override
 This module's operation-category tag, for ComputationGraph node tagging.
 
Tensor propagate_relevance (const Tensor &relevance_out, const LRPRuleConfig &) override
 Pass-through LRP relevance propagation, reshaped back to the input's shape.
 
bool supports_lrp_rule (LRPRule) const override
 A reshape is the same under every rule: 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. Called by forward() after precondition checks.
 

Detailed Description

y = reshape(x, {N, x.numel()/N}), N = x.shape().dim(0). No parameters, no gradient math beyond reshaping.

Note
Migrated from the original fully-unbatched semantics (flatten to a single rank-1 vector, including what's now the batch dim) by campaign_exai_dl_library_batch_dimension_support – a genuine behavior change, not just a shape-contract generalization like most other modules' migrations: the old behavior would have flattened N together with the feature dims, which is wrong once N carries real per-example batch semantics rather than always being 1.
op_type() returns OpType::Elementwise – op_type.hpp's closed set has no dedicated Reshape tag, and Elementwise is the closest existing fit (identity over the same values, just re-viewed). A deliberate choice, not an ideal one; flagged here so a future reader doesn't wonder why a reshape module is tagged Elementwise.

Constructor & Destructor Documentation

◆ FlattenModule()

pulsatrix::FlattenModule::FlattenModule ( DeviceBackend *  backend)
inlineexplicit

Constructs a flatten module.

Parameters
backendBackend to compute through. Not owned; must outlive this module.

Member Function Documentation

◆ backward()

Tensor pulsatrix::FlattenModule::backward ( const Tensor &  grad_output)
inlineoverridevirtual

Reshapes the gradient back to the shape forward() last saw.

Parameters
grad_outputGradient w.r.t. this module's (flattened) output.
Returns
Gradient w.r.t. this module's (original-shape) input.
Note
Must be called after forward() – uses the shape cached from that call.

Implements pulsatrix::Module.

◆ compute_device()

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

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

Implements pulsatrix::Module.

◆ op_type()

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

This module's operation-category tag, for ComputationGraph node tagging.

Returns
This module's OpType (see op_type.hpp's closed set).
Note
Added in Phase 2 Mission 0 alongside backward() – lets graph-wiring code tag nodes generically through a Module* rather than switching on concrete subclass.

Implements pulsatrix::Module.

◆ propagate_relevance()

Tensor pulsatrix::FlattenModule::propagate_relevance ( const Tensor &  relevance_out,
const LRPRuleConfig &   
)
inlineoverridevirtual

Pass-through LRP relevance propagation, reshaped back to the input's shape.

Note
No weighted connections exist to redistribute relevance across – reshaping is the only thing there is to undo, same rationale as ReluModule's pass-through rule for a parameterless pointwise op.

Implements pulsatrix::Module.

◆ supports_lrp_rule()

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

A reshape is the same under every rule: supports all of them.

Reimplemented from pulsatrix::Module.


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