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

down_proj(silu(gate_proj(x)) * up_proj(x)), the gated feedforward block used in place of a plain two-linear-layer MLP in most modern transformers. Rank-agnostic over (..., d_model) -> (..., d_model), matching MultiHeadAttentionModule's I/O contract so Phase 3 Mission 5 (TransformerBlock) can chain them directly. More...

#include <swiglu_module.hpp>

Inheritance diagram for pulsatrix::SwiGLUModule:
Collaboration diagram for pulsatrix::SwiGLUModule:

Public Member Functions

 SwiGLUModule (int64_t d_model, int64_t d_ff, DeviceBackend *backend)
 Constructs a SwiGLU block with zero-initialized projections.
 
Tensor backward (const Tensor &grad_output) override
 Gradient w.r.t. this module's input; sub-module parameter gradients accumulate inside gate_proj_/up_proj_/down_proj_ (reachable through parameters()).
 
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: down_proj_'s epsilon rule, then the diagonal Eq. 15 split into gate/up shares, then SiLU's identity pass-through, then gate_proj_'s/up_proj_'s epsilon rules summed.
 
std::vector< NamedParamRef > named_parameters () override
 gate_proj_'s, up_proj_'s, and down_proj_'s parameters, flattened.
 
int64_t d_model () const
 
int64_t d_ff () const
 
std::optional< DeviceType > compute_device () const override
 Where this layer computes, so forward() rejects an input on another device (FND-8).
 
Sub-module access – weight initialization from tests/loaders, and inspection.
LinearModule & gate_proj ()
 
LinearModule & up_proj ()
 
LinearModule & down_proj ()
 
- 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).
 
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
 Runs: project (gate, up) -> silu(gate) -> gate*up -> project (down).
 

Detailed Description

down_proj(silu(gate_proj(x)) * up_proj(x)), the gated feedforward block used in place of a plain two-linear-layer MLP in most modern transformers. Rank-agnostic over (..., d_model) -> (..., d_model), matching MultiHeadAttentionModule's I/O contract so Phase 3 Mission 5 (TransformerBlock) can chain them directly.

Composition over inheritance, same principle MultiHeadAttentionModule established: gate_proj_/up_proj_/down_proj_ are real LinearModule sub-objects driven through their own forward()/backward()/propagate_relevance(). The only hand-written operation is the elementwise silu(gate) * up gate – its LRP rule is a diagonal specialization of AttnLRP's Eq. 15 (Achtibat et al. 2024): for c[i] = a[i]*b[i], the sum_k over a matmul's shared contraction index collapses to a single term since there is no contraction, giving R_a[i] += (a[i]*b[i]/(2*c[i]+eps*sign(c[i]))) * R_c[i], symmetrically for R_b[i] – the same rule MultiHeadAttentionModule uses for its two batched matmuls, specialized to a diagonal case instead of a real contraction.

Note
SiLU itself gets this codebase's existing pointwise-nonlinearity identity pass-through treatment (Montavon et al. 2019 – the same convention already applied to ReluModule/DropoutModule/RMSNorm's-and-LayerNorm's Eq. 19 identity rule), extended here to a smooth, non-piecewise-linear nonlinearity for the first time in this codebase. A known approximation (not an exact DTD derivation for SiLU specifically), consistent with this codebase's own established simplification for every other pointwise nonlinearity so far – see mission_swiglu.md's Recon.
Conserves near-exactly (unlike SoftmaxModule's/MultiHeadAttentionModule's large by-design gap): the diagonal split above satisfies R_a[i]+R_b[i] = 2*(c[i]/(2c[i]+eps))*R_c[i] ~= R_c[i] (exact at eps=0), and SiLU's pass-through is exact by construction – so the whole composed module should conserve up to an epsilon residual only, same category as RoPEModule. Measured and confirmed in tests/swiglu_module_test.cpp, not just asserted.
op_type() reuses OpType::Elementwise – same resolution RoPEModule made (the hand-written novel part is an elementwise gate multiply, not a new operation category). Contrast MultiHeadAttentionModule, which earned a new category for genuinely novel cross-position mixing this module does not do – every position here is processed independently.

Constructor & Destructor Documentation

◆ SwiGLUModule()

pulsatrix::SwiGLUModule::SwiGLUModule ( int64_t  d_model,
int64_t  d_ff,
DeviceBackend *  backend 
)

Constructs a SwiGLU block with zero-initialized projections.

Parameters
d_modelInput/output width.
d_ffHidden (gate/up) width.
backendBackend to allocate/compute through. Not owned; must outlive this module.
Exceptions
std::invalid_argumentif d_model <= 0 or d_ff <= 0 – external boundary (construction arguments can originate from Phase 5's Python bindings with no upstream validation), same convention as every other constructor.

Member Function Documentation

◆ backward()

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

Gradient w.r.t. this module's input; sub-module parameter gradients accumulate inside gate_proj_/up_proj_/down_proj_ (reachable through parameters()).

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 as grad_output.
Exceptions
std::logic_errorif called before any forward().
std::invalid_argumentif grad_output's shape differs from the cached forward shape.
Note
Device-generic: the gate derivative is DeviceBackend::elementwise_backward(Silu), the products DeviceBackend::mul (GPU-native-kernels Mission 1).

Implements pulsatrix::Module.

◆ compute_device()

std::optional< DeviceType > pulsatrix::SwiGLUModule::compute_device ( ) const
inlineoverridevirtual

Where this layer computes, so forward() rejects an input on another device (FND-8).

Reimplemented from pulsatrix::Module.

◆ d_ff()

int64_t pulsatrix::SwiGLUModule::d_ff ( ) const
inline

◆ d_model()

int64_t pulsatrix::SwiGLUModule::d_model ( ) const
inline

◆ down_proj()

LinearModule & pulsatrix::SwiGLUModule::down_proj ( )
inline

◆ forward_impl()

Tensor pulsatrix::SwiGLUModule::forward_impl ( const Tensor &  input)
overrideprotectedvirtual

Runs: project (gate, up) -> silu(gate) -> gate*up -> project (down).

Parameters
input(..., d_model), rank >= 2, any device.
Returns
(..., d_model), same shape as input.
Exceptions
std::invalid_argumentif input's rank < 2 or final dimension != d_model.

Implements pulsatrix::Module.

◆ gate_proj()

LinearModule & pulsatrix::SwiGLUModule::gate_proj ( )
inline

◆ named_parameters()

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

gate_proj_'s, up_proj_'s, and down_proj_'s parameters, flattened.

Reimplemented from pulsatrix::Module.

◆ op_type()

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

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

Implements pulsatrix::Module.

◆ propagate_relevance()

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

LRP relevance propagation: down_proj_'s epsilon rule, then the diagonal Eq. 15 split into gate/up shares, then SiLU's identity pass-through, then gate_proj_'s/up_proj_'s epsilon rules summed.

Parameters
relevance_outRelevance at this module's output, matching the cached forward shape.
configSupplies the epsilon stabilizer for both the diagonal split and the three LinearModule epsilon rules.
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.

◆ up_proj()

LinearModule & pulsatrix::SwiGLUModule::up_proj ( )
inline

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