pulsatrix
Loading...
Searching...
No Matches
lrp.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <cstdint>
8#include <functional>
9#include <memory>
10#include <stdexcept>
11#include <string>
12#include <utility>
13#include <vector>
14
21#include "pulsatrix/tensor.hpp"
22
23namespace pulsatrix {
24
26enum class LRPSeed {
32 OneHot
33};
34
43struct LRPTarget {
44 std::vector<int64_t> targets;
45 std::vector<int64_t> contrasts = {};
47};
48
58using LRPComposite = std::function<LRPRuleConfig(size_t layer_index, const Module& module)>;
59
69namespace lrp_composite {
70
71namespace detail {
72inline bool is_conv(const Module& module) { return dynamic_cast<const Conv2DModule*>(&module) != nullptr; }
73inline LRPRuleConfig zennit_epsilon(float epsilon) {
74 LRPRuleConfig config{epsilon};
75 config.epsilon_bias_in_denominator = true;
76 return config;
77}
78inline LRPRuleConfig alpha_beta(float alpha, float beta, float epsilon) {
79 LRPRuleConfig config{epsilon};
80 config.rule = LRPRule::AlphaBeta;
81 config.alpha = alpha;
82 config.beta = beta;
83 return config;
84}
85} // namespace detail
86
88inline LRPComposite epsilon_plus(float epsilon = 1e-6f) {
89 return [epsilon](size_t, const Module& module) {
90 return detail::is_conv(module) ? detail::alpha_beta(1.0f, 0.0f, epsilon) : detail::zennit_epsilon(epsilon);
91 };
92}
93
95inline LRPComposite epsilon_alpha2_beta1(float epsilon = 1e-6f) {
96 return [epsilon](size_t, const Module& module) {
97 return detail::is_conv(module) ? detail::alpha_beta(2.0f, 1.0f, epsilon) : detail::zennit_epsilon(epsilon);
98 };
99}
100
110inline LRPComposite epsilon_gamma_box(float low, float high, float gamma = 0.25f, float epsilon = 1e-6f) {
111 auto first_conv_seen = std::make_shared<bool>(false);
112 return [=](size_t layer_index, const Module& module) {
113 if (layer_index == 0) {
114 *first_conv_seen = false;
115 }
116 if (!detail::is_conv(module)) {
117 return detail::zennit_epsilon(epsilon);
118 }
119 LRPRuleConfig config{epsilon};
120 if (!*first_conv_seen) {
121 *first_conv_seen = true;
122 config.rule = LRPRule::ZBox;
123 config.low = low;
124 config.high = high;
125 } else {
126 config.rule = LRPRule::Gamma;
127 config.gamma = gamma;
128 }
129 return config;
130 };
131}
132
133} // namespace lrp_composite
134
150class LRP {
151public:
153 explicit LRP(LRPRuleConfig config = LRPRuleConfig{}) : config_(config) {}
154
159 explicit LRP(LRPComposite composite, std::string name = "custom")
160 : config_(), composite_(std::move(composite)), composite_name_(std::move(name)) {
161 if (!composite_) {
162 throw std::invalid_argument("LRP: composite must not be empty");
163 }
164 }
165
167 [[nodiscard]] static LRP epsilon_plus(float epsilon = 1e-6f) {
168 return LRP(lrp_composite::epsilon_plus(epsilon), "epsilon_plus");
169 }
171 [[nodiscard]] static LRP epsilon_gamma_box(float low, float high, float gamma = 0.25f, float epsilon = 1e-6f) {
172 return LRP(lrp_composite::epsilon_gamma_box(low, high, gamma, epsilon), "epsilon_gamma_box");
173 }
175 [[nodiscard]] static LRP epsilon_alpha2_beta1(float epsilon = 1e-6f) {
176 return LRP(lrp_composite::epsilon_alpha2_beta1(epsilon), "epsilon_alpha2_beta1");
177 }
178
180 [[nodiscard]] Attribution explain(ExplainerContext& ctx, const Tensor& input, int64_t target_index,
181 DeviceBackend* backend) const {
182 return explain(ctx, input, LRPTarget{{target_index}}, backend);
183 }
184
197 [[nodiscard]] Attribution explain(ExplainerContext& ctx, const Tensor& input, const LRPTarget& target,
198 DeviceBackend* backend) const {
199 Tensor output = ctx.forward_pass(input);
200 if (output.rank() != 2) {
201 throw std::invalid_argument("LRP::explain: network output must be rank-2 (N, num_classes)");
202 }
203 const int64_t N = output.shape().dim(0);
204 const int64_t C = output.shape().dim(1);
205 const std::vector<int64_t> targets = per_row(target.targets, N, C, "targets");
206 const std::vector<int64_t> contrasts =
207 target.contrasts.empty() ? std::vector<int64_t>{} : per_row(target.contrasts, N, C, "contrasts");
208 for (size_t n = 0; n < contrasts.size(); ++n) {
209 if (contrasts[n] == targets[n]) {
210 throw std::invalid_argument(
211 "LRP::explain: a row's contrast equals its target (the seed would be zero)");
212 }
213 }
214
215 // Seed on the host, then upload: valid whatever device the output lives on.
216 std::vector<float> y;
217 if (target.seed == LRPSeed::OutputValue) {
218 y.resize(static_cast<size_t>(N * C));
219 output.backend()->copy(y.data(), output.data(), y.size() * sizeof(float),
222 }
223 auto seed_value = [&](int64_t n, int64_t c) {
224 return target.seed == LRPSeed::OutputValue ? y[static_cast<size_t>(n * C + c)] : 1.0f;
225 };
226 std::vector<float> seed(static_cast<size_t>(N * C), 0.0f);
227 float relevance_out_sum = 0.0f;
228 for (int64_t n = 0; n < N; ++n) {
229 const int64_t t = targets[static_cast<size_t>(n)];
230 seed[static_cast<size_t>(n * C + t)] = seed_value(n, t);
231 relevance_out_sum += seed_value(n, t);
232 if (!contrasts.empty()) {
233 const int64_t c = contrasts[static_cast<size_t>(n)];
234 seed[static_cast<size_t>(n * C + c)] = -seed_value(n, c);
235 relevance_out_sum -= seed_value(n, c);
236 }
237 }
238 Tensor seed_tensor(output.shape(), backend, seed, output.device());
239
240 std::vector<LRPRuleConfig> configs;
241 configs.reserve(ctx.modules().size());
242 for (size_t i = 0; i < ctx.modules().size(); ++i) {
243 configs.push_back(composite_ ? composite_(i, *ctx.modules()[i]) : config_);
244 }
245 std::string rules;
246 for (size_t i = 0; i < configs.size(); ++i) {
247 rules += (i == 0 ? "" : ",") + lrp_rule_name(configs[i].rule);
248 }
249 Tensor relevance = ctx.relevance_pass(seed_tensor, configs);
250 const float relevance_in_sum =
251 relevance.backend()->sum(relevance.data(), static_cast<size_t>(relevance.numel()));
252
253 return Attribution{"lrp",
254 std::move(relevance),
255 {{"rule", composite_ ? "composite:" + composite_name_ : lrp_rule_name(config_.rule)},
256 {"rules", rules},
257 {"epsilon", std::to_string(config_.epsilon)},
258 {"seed", target.seed == LRPSeed::OutputValue ? "output_value" : "one_hot"},
259 {"targets", join(target.targets)},
260 {"contrasts", join(target.contrasts)},
261 {"relevance_out_sum", std::to_string(relevance_out_sum)},
262 {"relevance_in_sum", std::to_string(relevance_in_sum)}}};
263 }
264
266 [[nodiscard]] const LRPRuleConfig& config() const { return config_; }
267
269 [[nodiscard]] const LRPComposite& composite() const { return composite_; }
270
271private:
272 // Expands a 1-entry list to N rows and validates every index against C classes.
273 static std::vector<int64_t> per_row(const std::vector<int64_t>& indices, int64_t N, int64_t C,
274 const char* what) {
275 if (indices.size() != 1 && static_cast<int64_t>(indices.size()) != N) {
276 throw std::invalid_argument(std::string("LRP::explain: ") + what +
277 " must have 1 entry or one per row of the output");
278 }
279 std::vector<int64_t> rows(static_cast<size_t>(N));
280 for (int64_t n = 0; n < N; ++n) {
281 const int64_t index = indices.size() == 1 ? indices[0] : indices[static_cast<size_t>(n)];
282 if (index < 0 || index >= C) {
283 throw std::invalid_argument(std::string("LRP::explain: ") + what + " index out of range");
284 }
285 rows[static_cast<size_t>(n)] = index;
286 }
287 return rows;
288 }
289
290 static std::string join(const std::vector<int64_t>& v) {
291 std::string out;
292 for (size_t i = 0; i < v.size(); ++i) {
293 out += (i == 0 ? "" : ",") + std::to_string(v[i]);
294 }
295 return out;
296 }
297
298 LRPRuleConfig config_;
299 LRPComposite composite_;
300 std::string composite_name_;
301};
302
303} // namespace pulsatrix
First-class explanation result type – values, method, and metadata together.
2D convolution, batched (input/output are rank-4: N x channels x H x W) – migrated from the original ...
Definition conv2d_module.hpp:31
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
virtual float sum(const float *in, size_t n)=0
sum_i in[i], returned to the host. Same reduction order as dot(). Synchronizes.
virtual void copy(void *dst, const void *src, size_t bytes, CopyDirection dir)=0
Copies bytes between buffers.
Wraps an ordered chain of Modules, running them via Module::forward_traced to build a real Computatio...
Definition explainer_context.hpp:67
Tensor forward_pass(const Tensor &input)
Runs the full module chain forward, building a fresh graph and caching every node's activation value ...
Definition explainer_context.hpp:95
Tensor relevance_pass(const Tensor &output_relevance, const LRPRuleConfig &config)
Propagates LRP relevance from the network output back to the input through every module's own propaga...
Definition explainer_context.hpp:234
const std::vector< Module * > & modules() const
The module chain, in forward order (not owned).
Definition explainer_context.hpp:269
Whole-model LRP: runs the forward pass, seeds relevance at the chosen output(s), and propagates it to...
Definition lrp.hpp:150
LRP(LRPRuleConfig config=LRPRuleConfig{})
Uniform rule: every module applies config.
Definition lrp.hpp:153
static LRP epsilon_alpha2_beta1(float epsilon=1e-6f)
Zennit EpsilonAlpha2Beta1 preset (lrp_composite::epsilon_alpha2_beta1).
Definition lrp.hpp:175
LRP(LRPComposite composite, std::string name="custom")
Per-layer rules from a composite; name is reported as "composite:<name>".
Definition lrp.hpp:159
static LRP epsilon_plus(float epsilon=1e-6f)
Zennit EpsilonPlus preset (lrp_composite::epsilon_plus).
Definition lrp.hpp:167
static LRP epsilon_gamma_box(float low, float high, float gamma=0.25f, float epsilon=1e-6f)
Zennit EpsilonGammaBox preset (lrp_composite::epsilon_gamma_box).
Definition lrp.hpp:171
const LRPComposite & composite() const
The composite, or an empty function for a uniform-rule LRP.
Definition lrp.hpp:269
Attribution explain(ExplainerContext &ctx, const Tensor &input, const LRPTarget &target, DeviceBackend *backend) const
Explains the given target(s) / contrast(s).
Definition lrp.hpp:197
Attribution explain(ExplainerContext &ctx, const Tensor &input, int64_t target_index, DeviceBackend *backend) const
Explains target_index for every row, OutputValue seed.
Definition lrp.hpp:180
const LRPRuleConfig & config() const
The uniform config (default-constructed when this LRP uses a composite).
Definition lrp.hpp:266
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
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
DeviceType device() const
Which device this tensor's buffer conceptually resides on.
Definition tensor.hpp:122
int64_t rank() const
Number of dimensions – shape().rank().
Definition tensor.hpp:119
DeviceBackend * backend() const
The backend that owns this tensor's buffer. Not owned by the Tensor.
Definition tensor.hpp:130
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
const float * data() const
Raw buffer access. nullptr iff numel() == 0.
Definition tensor.hpp:167
2D convolution – implemented via im2col + DeviceBackend::gemm (no new backend primitive).
Abstract interface isolating vendor-specific memory/compute operations from Tensor/ComputationGraph.
Stable interface every explainer gets, regardless of type (charter Part 2 SS2).
Dense/fully-connected layer – the reference Module implementation.
Selects which LRP rule variant a Module::propagate_relevance() call uses.
LRPRuleConfig alpha_beta(float alpha, float beta, float epsilon)
Definition lrp.hpp:78
bool is_conv(const Module &module)
Definition lrp.hpp:72
LRPRuleConfig zennit_epsilon(float epsilon)
Definition lrp.hpp:73
LRPComposite epsilon_gamma_box(float low, float high, float gamma=0.25f, float epsilon=1e-6f)
Zennit EpsilonGammaBox: ZBox(low, high) for the first Conv2D layer (lowest index),...
Definition lrp.hpp:110
LRPComposite epsilon_alpha2_beta1(float epsilon=1e-6f)
Zennit EpsilonAlpha2Beta1: Epsilon for Linear, AlphaBeta(2, 1) for Conv2D.
Definition lrp.hpp:95
LRPComposite epsilon_plus(float epsilon=1e-6f)
Zennit EpsilonPlus: Epsilon for Linear, ZPlus (AlphaBeta 1, 0) for Conv2D.
Definition lrp.hpp:88
Definition acquisition_functions.hpp:16
LRPSeed
How LRP seeds relevance at the target output.
Definition lrp.hpp:26
std::string lrp_rule_name(LRPRule rule)
Lower-case rule name ("epsilon", "gamma", "alpha_beta", "zbox") for messages/metadata.
Definition lrp_rule_config.hpp:35
@ AlphaBeta
Alpha-beta rule (Bach et al. 2015) with alpha - beta == 1; alpha 1, beta 0 is ZPlus.
std::function< LRPRuleConfig(size_t layer_index, const Module &module)> LRPComposite
Per-layer LRP rule choice: maps (top-level module index in forward order, module) to the LRPRuleConfi...
Definition lrp.hpp:58
An explanation result: the raw attribution values, the method that produced them, and any relevant me...
Definition attribution.hpp:23
Configuration for LRP relevance propagation: which rule a module applies and its hyperparameters....
Definition lrp_rule_config.hpp:57
float epsilon
Stabilizer added to every rule's denominator: z + epsilon * sign(z), sign(0) = +1.
Definition lrp_rule_config.hpp:59
bool epsilon_bias_in_denominator
Epsilon rule only, honoured by LinearModule / Conv2DModule: include the bias in the denominator,...
Definition lrp_rule_config.hpp:77
LRPRule rule
Which rule to apply.
Definition lrp_rule_config.hpp:61
What LRP explains: one target class (and optionally one contrast class) per row of the network's (N,...
Definition lrp.hpp:43
std::vector< int64_t > contrasts
Definition lrp.hpp:45
LRPSeed seed
Definition lrp.hpp:46
std::vector< int64_t > targets
Definition lrp.hpp:44
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).