pulsatrix
Loading...
Searching...
No Matches
conjunction_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
42class ConjunctionModule : public Module {
43public:
45 enum class TNorm {
46 Product,
48 Godel
49 };
50
57
58 using Module::forward;
59
70 [[nodiscard]] Tensor forward(const Tensor& a, const Tensor& b);
71
82 [[nodiscard]] static Tensor stack_operands(const Tensor& a, const Tensor& b, DeviceBackend* backend);
83
98 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
99
101 [[nodiscard]] OpType op_type() const override { return OpType::Elementwise; }
102
130 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
131
132
134 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
135
136protected:
148 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
149
150private:
151 DeviceBackend* backend_;
152 TNorm t_norm_;
153 Tensor last_input_;
154 Tensor last_output_;
155 bool has_forwarded_ = false;
156};
157
158} // namespace pulsatrix
y = a T b for a selected t-norm T, over two independent fuzzy-truth-valued operand tensors (values in...
Definition conjunction_module.hpp:42
Tensor forward_impl(const Tensor &input) override
Splits input (leading dim 2) into the two operands and computes the selected t-norm elementwise.
Tensor backward(const Tensor &grad_output) override
Gradient w.r.t. this module's (stacked) input, via the selected t-norm's closed-form partial derivati...
Tensor forward(const Tensor &a, const Tensor &b)
Convenience two-operand entry point: builds the stacked input via stack_operands() and delegates to M...
OpType op_type() const override
Elementwise per this module's own op_type() convention (ReluModule/ResidualModule).
Definition conjunction_module.hpp:101
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition conjunction_module.hpp:134
ConjunctionModule(DeviceBackend *backend, TNorm t_norm=TNorm::Product)
Constructs a conjunction module.
TNorm
Which t-norm this instance computes. Product is the campaign's primary case.
Definition conjunction_module.hpp:45
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
LRP relevance propagation for the selected t-norm โ€“ genuinely novel, no prior art (research_2026_neur...
static Tensor stack_operands(const Tensor &a, const Tensor &b, DeviceBackend *backend)
Combines two independent operand tensors into the leading-dim-2 stacked tensor forward_impl()/backwar...
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
Tensor forward(const Tensor &input)
Runs this module's forward computation.
Definition module.hpp:73
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