|
pulsatrix
|
Softmax over the tensor's last dimension, applied independently to every "row" (every fixed combination of all leading dimensions). More...
#include <softmax_module.hpp>


Public Member Functions | |
| SoftmaxModule (DeviceBackend *backend) | |
| Constructs a softmax module. | |
| Tensor | backward (const Tensor &grad_output) override |
| Computes the gradient w.r.t. this module's input. | |
| OpType | op_type () const override |
| Activation per charter's closed OpType set (same category as ReluModule). | |
| Tensor | propagate_relevance (const Tensor &relevance_out, const LRPRuleConfig &config) override |
| AttnLRP's softmax relevance rule (Achtibat et al. 2024, Eq. 13 – Deep Taylor Decomposition), applied per row: R_in[i] = x[i] * (R_out[i] - s[i] * sum_j(R_out[j])), with x the cached forward input and s the cached forward output. | |
| 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< 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 |
| Numerically stable softmax over the last axis, per row (subtract the row max before exponentiating – same convention as CrossEntropyLoss::forward). | |
Softmax over the tensor's last dimension, applied independently to every "row" (every fixed combination of all leading dimensions).
|
explicit |
Constructs a softmax module.
| backend | Backend to allocate through. Not owned; must outlive this module. |
Computes the gradient w.r.t. this module's input.
| grad_output | Gradient w.r.t. this module's output. Must match the shape of the most recent forward() call's output. |
Implements pulsatrix::Module.
|
inlineoverridevirtual |
Where this layer computes, so forward() rejects an input on another device (FND-8).
Reimplemented from pulsatrix::Module.
Numerically stable softmax over the last axis, per row (subtract the row max before exponentiating – same convention as CrossEntropyLoss::forward).
| input | Input tensor. Must be rank >= 1; any device. |
Implements pulsatrix::Module.
|
inlineoverridevirtual |
Activation per charter's closed OpType set (same category as ReluModule).
Implements pulsatrix::Module.
|
overridevirtual |
AttnLRP's softmax relevance rule (Achtibat et al. 2024, Eq. 13 – Deep Taylor Decomposition), applied per row: R_in[i] = x[i] * (R_out[i] - s[i] * sum_j(R_out[j])), with x the cached forward input and s the cached forward output.
| relevance_out | Relevance at this module's output. Must match forward()'s shape. |
| config | Unused – Eq. 13 takes no epsilon/gamma parameter. Deliberately NOT epsilon-stabilized (see the conservation note below). |
Implements pulsatrix::Module.