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

y = (mask_i ? x_i / (1 - p) : 0) at training time (inverted dropout – scaling happens at training time so eval-time forward needs no rescaling); y = x at eval time or when p == 0. No parameters. More...

#include <dropout_module.hpp>

Inheritance diagram for pulsatrix::DropoutModule:
Collaboration diagram for pulsatrix::DropoutModule:

Public Member Functions

 DropoutModule (float p, DeviceBackend *backend, uint64_t seed)
 Constructs a dropout layer.
 
 DropoutModule (float p, DeviceBackend *backend)
 Seeded from the global seed stream (next_seed(), FND-7), so every layer built this way draws its own masks – reproducibly for a given set_seed().
 
Tensor backward (const Tensor &grad_output) override
 Computes the gradient w.r.t. this module's input: grad_output * mask * scale, the true gradient of the actual (masked/scaled) forward computation. If the most recent forward() ran in eval mode, mask is all-ones and scale == 1, so this correctly reduces to identity with no special-casing needed here.
 
OpType op_type () const override
 Elementwise per charter's closed OpType set – a per-element scale-or-zero operation, an honest fit for the existing category (no new OpType needed).
 
Tensor propagate_relevance (const Tensor &relevance_out, const LRPRuleConfig &config) override
 Unconditional identity LRP relevance propagation.
 
bool supports_lrp_rule (LRPRule) const override
 Pass-through relevance 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 – per-element RNG draw at training time, identity at eval time or p == 0.
 

Detailed Description

y = (mask_i ? x_i / (1 - p) : 0) at training time (inverted dropout – scaling happens at training time so eval-time forward needs no rescaling); y = x at eval time or when p == 0. No parameters.

Note
propagate_relevance is unconditional identity pass-through, reusing ReluModule's established precedent (Montavon et al. 2019: activation-/ regularization-like pointwise operations pass relevance through unchanged) – the same masked-backward/unconditional-identity-relevance split ReluModule already uses, not a new rule invented for this module. Dropout is conventionally disabled during inference/explanation in every mainstream framework, so at is_training() == false (the expected state when running LRP) forward is already identity, making this the mathematically exact treatment for that case, not just an approximation carried over from ReLU.

Constructor & Destructor Documentation

◆ DropoutModule() [1/2]

pulsatrix::DropoutModule::DropoutModule ( float  p,
DeviceBackend *  backend,
uint64_t  seed 
)

Constructs a dropout layer.

Parameters
pDrop probability, must be in [0, 1).
backendBackend to compute through. Not owned; must outlive this module.
seedRNG seed for the masks. The two-argument constructor draws one from the global seed stream instead (next_seed(), FND-7). Masks come from DeviceBackend::dropout_forward's counter-based generator (element k of the stream is a pure function of (seed, k)), so the same seed produces the same masks on every backend. Each training forward() advances the stream by numel().
Exceptions
std::invalid_argumentif p < 0 or p >= 1 – external boundary (p == 1 would make scale = 1/(1-p) diverge; construction arguments can originate from Phase 5's Python bindings with no upstream validation).
Note
No device parameter: the cached mask is allocated per forward() on the input's own device, same reasoning as ReluModule's forward_impl.

◆ DropoutModule() [2/2]

pulsatrix::DropoutModule::DropoutModule ( float  p,
DeviceBackend *  backend 
)

Seeded from the global seed stream (next_seed(), FND-7), so every layer built this way draws its own masks – reproducibly for a given set_seed().

Member Function Documentation

◆ backward()

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

Computes the gradient w.r.t. this module's input: grad_output * mask * scale, the true gradient of the actual (masked/scaled) forward computation. If the most recent forward() ran in eval mode, mask is all-ones and scale == 1, so this correctly reduces to identity with no special-casing needed here.

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 (GPU-native-kernels Mission 1b). The scale applied is the one the most recent forward() actually used: before Mission 1b this always applied 1/(1-p), so a gradient through an eval-mode forward came back wrongly scaled.

Implements pulsatrix::Module.

◆ compute_device()

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

The actual forward computation – per-element RNG draw at training time, identity at eval time or p == 0.

Note
Device-generic (GPU-native-kernels Mission 1b).

Implements pulsatrix::Module.

◆ op_type()

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

Elementwise per charter's closed OpType set – a per-element scale-or-zero operation, an honest fit for the existing category (no new OpType needed).

Implements pulsatrix::Module.

◆ propagate_relevance()

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

Unconditional identity LRP relevance propagation.

Parameters
relevance_outRelevance at this module's output. Must match the shape of the most recent forward() call's output.
configUnused.
Returns
relevance_out, unchanged – see the class-level note for why this is exact, not approximate, when is_training() == false.
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::DropoutModule::supports_lrp_rule ( LRPRule  ) const
inlineoverridevirtual

Pass-through relevance 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: