pulsatrix
Loading...
Searching...
No Matches
linear_algebra.hpp
Go to the documentation of this file.
1
17#pragma once
18
19#include <cmath>
20#include <cstdint>
21#include <stdexcept>
22#include <utility>
23#include <vector>
24
25namespace pulsatrix {
26
36[[nodiscard]] inline std::vector<float> SolveLinearSystem(std::vector<std::vector<float>> a,
37 std::vector<float> b) {
38 const auto n = static_cast<int64_t>(b.size());
39 for (int64_t col = 0; col < n; ++col) {
40 int64_t pivot_row = col;
41 float pivot_magnitude = std::fabs(a[static_cast<size_t>(col)][static_cast<size_t>(col)]);
42 for (int64_t row = col + 1; row < n; ++row) {
43 float magnitude = std::fabs(a[static_cast<size_t>(row)][static_cast<size_t>(col)]);
44 if (magnitude > pivot_magnitude) {
45 pivot_magnitude = magnitude;
46 pivot_row = row;
47 }
48 }
49 if (pivot_magnitude < 1e-8f) {
50 throw std::runtime_error("SolveLinearSystem: near-singular system");
51 }
52 std::swap(a[static_cast<size_t>(col)], a[static_cast<size_t>(pivot_row)]);
53 std::swap(b[static_cast<size_t>(col)], b[static_cast<size_t>(pivot_row)]);
54
55 for (int64_t row = col + 1; row < n; ++row) {
56 float factor = a[static_cast<size_t>(row)][static_cast<size_t>(col)] /
57 a[static_cast<size_t>(col)][static_cast<size_t>(col)];
58 for (int64_t c = col; c < n; ++c) {
59 a[static_cast<size_t>(row)][static_cast<size_t>(c)] -=
60 factor * a[static_cast<size_t>(col)][static_cast<size_t>(c)];
61 }
62 b[static_cast<size_t>(row)] -= factor * b[static_cast<size_t>(col)];
63 }
64 }
65
66 std::vector<float> x(static_cast<size_t>(n), 0.0f);
67 for (int64_t row = n - 1; row >= 0; --row) {
68 float sum = b[static_cast<size_t>(row)];
69 for (int64_t c = row + 1; c < n; ++c) {
70 sum -= a[static_cast<size_t>(row)][static_cast<size_t>(c)] * x[static_cast<size_t>(c)];
71 }
72 x[static_cast<size_t>(row)] = sum / a[static_cast<size_t>(row)][static_cast<size_t>(row)];
73 }
74 return x;
75}
76
77} // namespace pulsatrix
Definition acquisition_functions.hpp:16
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