pulsatrix
Loading...
Searching...
No Matches
weighted_linear_regression.hpp
Go to the documentation of this file.
1
6#pragma once
7
8#include <cstddef>
9#include <stdexcept>
10#include <utility>
11#include <vector>
12
13#include "pulsatrix/assert.hpp"
15
16namespace pulsatrix {
17
38[[nodiscard]] inline std::vector<float> fit_weighted_linear_regression(
39 const std::vector<std::vector<float>>& samples, const std::vector<float>& targets,
40 const std::vector<float>& weights, float l2_lambda) {
41 if (samples.empty()) {
42 throw std::invalid_argument("fit_weighted_linear_regression: samples must not be empty");
43 }
44 if (samples.size() != targets.size() || samples.size() != weights.size()) {
45 throw std::invalid_argument(
46 "fit_weighted_linear_regression: samples, targets, and weights must be the same size");
47 }
48
49 const auto n_features = static_cast<int64_t>(samples[0].size());
50 for (const auto& row : samples) {
51 if (static_cast<int64_t>(row.size()) != n_features) {
52 throw std::invalid_argument("fit_weighted_linear_regression: every sample row must be the same length");
53 }
54 }
55
56 std::vector<std::vector<float>> a(static_cast<size_t>(n_features),
57 std::vector<float>(static_cast<size_t>(n_features), 0.0f));
58 std::vector<float> b(static_cast<size_t>(n_features), 0.0f);
59
60 for (size_t k = 0; k < samples.size(); ++k) {
61 for (int64_t i = 0; i < n_features; ++i) {
62 b[static_cast<size_t>(i)] += weights[k] * samples[k][static_cast<size_t>(i)] * targets[k];
63 for (int64_t j = 0; j < n_features; ++j) {
64 a[static_cast<size_t>(i)][static_cast<size_t>(j)] +=
65 weights[k] * samples[k][static_cast<size_t>(i)] * samples[k][static_cast<size_t>(j)];
66 }
67 }
68 }
69 for (int64_t i = 0; i < n_features; ++i) {
70 a[static_cast<size_t>(i)][static_cast<size_t>(i)] += l2_lambda;
71 }
72
73 return SolveLinearSystem(std::move(a), std::move(b));
74}
75
76} // namespace pulsatrix
PULSATRIX_ASSERT – debug-only invariant check for programmer errors, distinct from throw (used for ca...
Small, dense linear-system solve – shared by weighted_linear_regression.hpp (Phase 3's LIME/KernelSHA...
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
std::vector< float > SolveLinearSystem(std::vector< std::vector< float > > a, std::vector< float > b)
Solves A*x = b via Gaussian elimination with partial pivoting.
Definition linear_algebra.hpp:36