pulsatrix
Loading...
Searching...
No Matches
dropout_module.hpp
Go to the documentation of this file.
1
6#pragma once
7
8
9#include <optional>
10#include <cstdint>
11
12#include "pulsatrix/module.hpp"
13
14namespace pulsatrix {
15
30class DropoutModule : public Module {
31public:
48 DropoutModule(float p, DeviceBackend* backend, uint64_t seed);
49
52 DropoutModule(float p, DeviceBackend* backend);
53
69 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
70
73 [[nodiscard]] OpType op_type() const override { return OpType::Elementwise; }
74
86 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
87
89 [[nodiscard]] bool supports_lrp_rule(LRPRule) const override { return true; }
90
91
93 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
94
95protected:
101 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
102
103private:
104 float p_;
105 float scale_;
106 DeviceBackend* backend_;
107 uint64_t seed_;
108 uint64_t draws_ = 0; // elements consumed from the (seed_, k) stream so far
109 Shape last_shape_ = Shape({0});
110 Tensor last_mask_; // 1.0 (kept) or 0.0 (dropped); meaningful only when !last_was_identity_
111 bool last_was_identity_ = true;
112 bool has_forwarded_ = false;
113};
114
115} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
virtual DeviceType device() const noexcept=0
Which device this backend's buffers reside on.
y = (mask_i ? x_i / (1 - p) : 0) at training time (inverted dropout – scaling happens at training tim...
Definition dropout_module.hpp:30
Tensor forward_impl(const Tensor &input) override
The actual forward computation – per-element RNG draw at training time, identity at eval time or p ==...
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 th...
DropoutModule(float p, DeviceBackend *backend, uint64_t seed)
Constructs a dropout layer.
bool supports_lrp_rule(LRPRule) const override
Pass-through relevance is the same under every rule: supports all of them.
Definition dropout_module.hpp:89
DropoutModule(float p, DeviceBackend *backend)
Seeded from the global seed stream (next_seed(), FND-7), so every layer built this way draws its own ...
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
Unconditional identity LRP relevance propagation.
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition dropout_module.hpp:93
OpType op_type() const override
Elementwise per charter's closed OpType set – a per-element scale-or-zero operation,...
Definition dropout_module.hpp:73
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
An N-dimensional shape. A plain aggregate of dimensions with no invariant beyond "non-negative dimens...
Definition shape.hpp:24
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Abstract base every layer subclasses – NVI forward(), pure-virtual LRP contract.
Definition acquisition_functions.hpp:16
OpType
The op-type tag a Node carries. Charter Part 2 §3: nodes are tagged by a small closed set of op types...
Definition op_type.hpp:19
LRPRule
The LRP rule family a module applies. Semantics follow Zennit 1.0.0 exactly (Anders et al....
Definition lrp_rule_config.hpp:19
Configuration for LRP relevance propagation: which rule a module applies and its hyperparameters....
Definition lrp_rule_config.hpp:57