pulsatrix
Loading...
Searching...
No Matches
max_pool2d_module.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <optional>
8#include <vector>
9
10#include "pulsatrix/module.hpp"
11
12namespace pulsatrix {
13
26class MaxPool2DModule : public Module {
27public:
37 MaxPool2DModule(int64_t kernel_h, int64_t kernel_w, DeviceBackend* backend);
38
50 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
51
53 [[nodiscard]] OpType op_type() const override { return OpType::Pooling; }
54
67 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
68
70 [[nodiscard]] bool supports_lrp_rule(LRPRule) const override { return true; }
71
72
74 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
75
76protected:
84 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
85
86private:
87 int64_t kernel_h_;
88 int64_t kernel_w_;
89 DeviceBackend* backend_;
90 Shape last_input_shape_ = Shape({0});
91 int64_t last_out_h_ = 0;
92 int64_t last_out_w_ = 0;
93 // Flat (within-window) offset of the argmax for each output element, indexed the same
94 // way as the output buffer (n, c, oh, ow) -- row-major flat index into (H, W) input
95 // plane, i.e. ih * W + iw. Sized N*C*out_h*out_w after a real forward() call.
96 // Stored as whole-number floats on the input's device (exact below 2^24 -- far above any
97 // plane this module pools), consumed by DeviceBackend::max_unpool (GPU-native-kernels
98 // Mission 4).
99 Tensor argmax_flat_index_ = Tensor(Shape({0}), backend_);
100 bool has_forwarded_ = false;
101};
102
103} // 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.
Max pooling, rank-4 (N, channels, H, W), matching Conv2DModule's convention. Stride fixed equal to ke...
Definition max_pool2d_module.hpp:26
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
Winner-take-all LRP relevance propagation (Bach et al. 2015).
Tensor forward_impl(const Tensor &input) override
The actual forward computation – per-window max, argmax cached per output element for backward()/prop...
bool supports_lrp_rule(LRPRule) const override
Winner-take-all ignores the config: the same under every rule, so supports all of them.
Definition max_pool2d_module.hpp:70
Tensor backward(const Tensor &grad_output) override
Computes the gradient w.r.t. this module's input – only the cached argmax position within each window...
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition max_pool2d_module.hpp:74
OpType op_type() const override
Pooling per charter's closed OpType set.
Definition max_pool2d_module.hpp:53
MaxPool2DModule(int64_t kernel_h, int64_t kernel_w, DeviceBackend *backend)
Constructs a max-pool layer.
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