pulsatrix
Loading...
Searching...
No Matches
flatten_module.hpp
Go to the documentation of this file.
1
7#pragma once
8
9#include <optional>
10#include "pulsatrix/module.hpp"
11
12namespace pulsatrix {
13
28class FlattenModule : public Module {
29public:
34 explicit FlattenModule(DeviceBackend* backend) : backend_(backend), last_input_shape_(Shape({0})) {}
35
42 [[nodiscard]] Tensor backward(const Tensor& grad_output) override {
43 require_device(grad_output, backend_->device(), "FlattenModule::backward");
44 Tensor grad_input(grad_output);
45 grad_input.reshape(last_input_shape_);
46 return grad_input;
47 }
48
49 [[nodiscard]] OpType op_type() const override { return OpType::Elementwise; }
50
57 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig&) override {
58 require_device(relevance_out, backend_->device(), "FlattenModule::propagate_relevance");
59 Tensor relevance_in(relevance_out);
60 relevance_in.reshape(last_input_shape_);
61 return relevance_in;
62 }
63
65 [[nodiscard]] bool supports_lrp_rule(LRPRule) const override { return true; }
66
67
69 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
70
71protected:
72 [[nodiscard]] Tensor forward_impl(const Tensor& input) override {
73 last_input_shape_ = input.shape();
74 int64_t N = input.shape().dim(0);
75 Tensor output(input);
76 output.reshape(Shape({N, input.numel() / N}));
77 return output;
78 }
79
80private:
81 DeviceBackend* backend_;
82 Shape last_input_shape_;
83};
84
85} // 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 = reshape(x, {N, x.numel()/N}), N = x.shape().dim(0). No parameters, no gradient math beyond reshap...
Definition flatten_module.hpp:28
Tensor backward(const Tensor &grad_output) override
Reshapes the gradient back to the shape forward() last saw.
Definition flatten_module.hpp:42
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &) override
Pass-through LRP relevance propagation, reshaped back to the input's shape.
Definition flatten_module.hpp:57
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition flatten_module.hpp:69
FlattenModule(DeviceBackend *backend)
Constructs a flatten module.
Definition flatten_module.hpp:34
OpType op_type() const override
This module's operation-category tag, for ComputationGraph node tagging.
Definition flatten_module.hpp:49
bool supports_lrp_rule(LRPRule) const override
A reshape is the same under every rule: supports all of them.
Definition flatten_module.hpp:65
Tensor forward_impl(const Tensor &input) override
The actual forward computation. Called by forward() after precondition checks.
Definition flatten_module.hpp:72
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
int64_t dim(size_t index) const
Size of a single dimension.
Definition shape.hpp:93
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Tensor & reshape(Shape new_shape)
Reinterprets this tensor's dimensions in place – same buffer, new shape.
int64_t numel() const
Total element count – shape().numel().
Definition tensor.hpp:116
const Shape & shape() const
This tensor's shape.
Definition tensor.hpp:113
Abstract base every layer subclasses – NVI forward(), pure-virtual LRP contract.
Definition acquisition_functions.hpp:16
void require_device(const Tensor &t, DeviceType expected, const char *where)
Throws unless t lives on expected – the check every module and loss runs on the tensors handed to it ...
Definition tensor.hpp:285
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