pulsatrix
Loading...
Searching...
No Matches
avg_pool2d_module.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <optional>
9
10namespace pulsatrix {
11
25class AvgPool2DModule : public Module {
26public:
37 AvgPool2DModule(int64_t kernel_h, int64_t kernel_w, DeviceBackend* backend, float eps = 1e-6f);
38
51 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
52
54 [[nodiscard]] OpType op_type() const override { return OpType::Pooling; }
55
67 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
68
69
71 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
72
73protected:
80 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
81
82private:
83 int64_t kernel_h_;
84 int64_t kernel_w_;
85 DeviceBackend* backend_;
86 float eps_;
87 Tensor last_input_; // (N, C, H, W)
88 int64_t last_out_h_ = 0;
89 int64_t last_out_w_ = 0;
90 bool has_forwarded_ = false;
91};
92
93} // namespace pulsatrix
Average pooling, rank-4 (N, channels, H, W), matching Conv2DModule's convention. Stride fixed equal t...
Definition avg_pool2d_module.hpp:25
OpType op_type() const override
Pooling per charter's closed OpType set.
Definition avg_pool2d_module.hpp:54
AvgPool2DModule(int64_t kernel_h, int64_t kernel_w, DeviceBackend *backend, float eps=1e-6f)
Constructs an average-pool layer.
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition avg_pool2d_module.hpp:71
Tensor forward_impl(const Tensor &input) override
The actual forward computation – per-window mean.
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
Epsilon/z-rule LRP relevance propagation, weight = 1/K (Bach et al. 2015).
Tensor backward(const Tensor &grad_output) override
Computes the gradient w.r.t. this module's input – uniform 1/K per input position in each window (the...
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.
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
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
Configuration for LRP relevance propagation: which rule a module applies and its hyperparameters....
Definition lrp_rule_config.hpp:57