pulsatrix
Loading...
Searching...
No Matches
lime.hpp
Go to the documentation of this file.
1
6#pragma once
7
8#include <cmath>
9#include <functional>
10#include <random>
11#include <stdexcept>
12#include <string>
13#include <vector>
14
15#include "pulsatrix/assert.hpp"
18#include "pulsatrix/tensor.hpp"
20
21namespace pulsatrix {
22
41class LIME {
42public:
63 [[nodiscard]] Attribution explain(const std::function<Tensor(const Tensor&)>& predict, const Tensor& input,
64 int64_t target_index, int64_t num_samples, float sigma, float l2_lambda,
65 unsigned seed, DeviceBackend* backend) const {
66 if (num_samples <= 0) {
67 throw std::invalid_argument("LIME::explain: num_samples must be positive");
68 }
69 if (sigma <= 0.0f) {
70 throw std::invalid_argument("LIME::explain: sigma must be positive");
71 }
72
73 Tensor base_output = predict(input);
74 PULSATRIX_ASSERT(target_index >= 0 && target_index < base_output.numel());
75 float base_value = base_output.read_element(target_index);
76
77 int64_t n_features = input.numel();
78 const std::vector<float> input_values = input.to_host_vector();
79 DeviceBackend* input_backend = explainer_detail::backend_beside(input, backend);
80 std::mt19937 rng(seed);
81 std::normal_distribution<float> noise(0.0f, sigma);
82
83 std::vector<std::vector<float>> samples;
84 std::vector<float> targets;
85 std::vector<float> weights;
86 samples.reserve(static_cast<size_t>(num_samples));
87 targets.reserve(static_cast<size_t>(num_samples));
88 weights.reserve(static_cast<size_t>(num_samples));
89
90 for (int64_t s = 0; s < num_samples; ++s) {
91 std::vector<float> perturbed_values(static_cast<size_t>(n_features));
92 std::vector<float> delta(static_cast<size_t>(n_features));
93 float squared_distance = 0.0f;
94 for (int64_t i = 0; i < n_features; ++i) {
95 float d = noise(rng);
96 delta[static_cast<size_t>(i)] = d;
97 perturbed_values[static_cast<size_t>(i)] = input_values[static_cast<size_t>(i)] + d;
98 squared_distance += d * d;
99 }
100 Tensor perturbed(input.shape(), input_backend, perturbed_values, input.device());
101
102 Tensor output = predict(perturbed);
103 float centered_target = output.read_element(target_index) - base_value;
104 float weight = std::exp(-squared_distance / (2.0f * sigma * sigma));
105
106 samples.push_back(std::move(delta));
107 targets.push_back(centered_target);
108 weights.push_back(weight);
109 }
110
111 std::vector<float> coefficients = fit_weighted_linear_regression(samples, targets, weights, l2_lambda);
112
113 coefficients.resize(static_cast<size_t>(n_features));
114 Tensor values(input.shape(), input_backend, coefficients, input.device());
115
116 return Attribution{"lime", std::move(values),
117 {{"target_index", std::to_string(target_index)},
118 {"num_samples", std::to_string(num_samples)},
119 {"sigma", std::to_string(sigma)}}};
120 }
121};
122
123} // namespace pulsatrix
PULSATRIX_ASSERT – debug-only invariant check for programmer errors, distinct from throw (used for ca...
#define PULSATRIX_ASSERT(cond)
Aborts with a diagnostic message if cond is false. Debug-only – use for conditions that indicate a bu...
Definition assert.hpp:22
First-class explanation result type – values, method, and metadata together.
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
Fits a locality-weighted linear surrogate around one input: perturb x with Gaussian noise,...
Definition lime.hpp:41
Attribution explain(const std::function< Tensor(const Tensor &)> &predict, const Tensor &input, int64_t target_index, int64_t num_samples, float sigma, float l2_lambda, unsigned seed, DeviceBackend *backend) const
Computes the LIME local surrogate explanation for one output index.
Definition lime.hpp:63
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 numel() const
Total element count – shape().numel().
Definition tensor.hpp:116
const Shape & shape() const
This tensor's shape.
Definition tensor.hpp:113
float read_element(int64_t flat_index) const
Reads one element to the host through the owning backend, on any device.
std::vector< float > to_host_vector() const
Copies the whole buffer to a host vector through the owning backend, on any device.
Abstract interface isolating vendor-specific memory/compute operations from Tensor/ComputationGraph.
DeviceBackend * backend_beside(const Tensor &like, DeviceBackend *backend)
The backend to allocate a tensor through that must live beside like.
Definition attribution.hpp:49
Definition acquisition_functions.hpp:16
std::vector< float > fit_weighted_linear_regression(const std::vector< std::vector< float > > &samples, const std::vector< float > &targets, const std::vector< float > &weights, float l2_lambda)
Fits w* = argmin_w sum_i weight_i*(target_i - w^T sample_i)^2 + l2_lambda*||w||^2 via the normal equa...
Definition weighted_linear_regression.hpp:38
An explanation result: the raw attribution values, the method that produced them, and any relevant me...
Definition attribution.hpp:23
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).
Weighted least squares via normal equations – the shared fitting primitive Phase 3's LIME and KernelS...