igraph/linalg/blas.rs
1//! The BLAS interface (`igraph_blas.h`).
2//!
3//! Degenerate (empty) dimensions are handled on the Rust side, because the
4//! bundled BLAS terminates the process on invalid leading dimensions.
5
6use super::to_c_int;
7use crate::{
8 error::{Error, Result},
9 ffi::*,
10 igraph_call,
11 matrix::Matrix,
12 vector::Vector,
13};
14
15fn dgemv_dims(transpose: bool, a: &Matrix, x: usize, y: usize) -> Result<()> {
16 to_c_int(a.nrow(), "number of rows")?;
17 to_c_int(a.ncol(), "number of columns")?;
18 let (xin, yout) = if transpose {
19 (a.nrow(), a.ncol())
20 } else {
21 (a.ncol(), a.nrow())
22 };
23 if x != xin || y != yout {
24 return Err(Error::invalid(format!(
25 "dgemv: a {:?} matrix{} needs x of length {xin} and y of length {yout}, got {x} and {y}",
26 a.shape(),
27 if transpose { " (transposed)" } else { "" }
28 )));
29 }
30 Ok(())
31}
32
33/// Matrix-vector product `alpha * op(A) * x + beta * y`, where `op(A)` is
34/// `A` or its transpose (`igraph_blas_dgemv`); the result is returned as a
35/// new vector. When `beta` is zero, `y` only provides the length.
36///
37/// Time complexity: O(nk) for an `n` × `k` matrix.
38///
39/// Binds [`igraph_blas_dgemv`](https://igraph.org/c/html/latest/igraph-Linalg.html#igraph_blas_dgemv).
40///
41/// # Errors
42/// If the lengths of `x` and `y` do not match `op(A)`.
43///
44/// # Examples
45///
46/// ```
47/// use igraph::{linalg::blas_dgemv, prelude::*};
48/// let a = Matrix::from_rows(&[[1.0, 2.0], [3.0, 4.0]]).unwrap();
49/// assert_eq!(blas_dgemv(false, 1.0, &a, &[1.0, 1.0], 0.0, &[0.0, 0.0]).unwrap(), vec![3.0, 7.0]);
50/// assert_eq!(blas_dgemv(true, 2.0, &a, &[1.0, 1.0], 1.0, &[1.0, 1.0]).unwrap(), vec![9.0, 13.0]);
51/// ```
52pub fn blas_dgemv(
53 transpose: bool,
54 alpha: f64,
55 a: &Matrix,
56 x: &[f64],
57 beta: f64,
58 y: &[f64],
59) -> Result<Vec<f64>> {
60 dgemv_dims(transpose, a, x.len(), y.len())?;
61 let mut out = Vector::from_slice(y);
62 if a.is_empty() {
63 // op(A) x is an empty sum; BLAS would reject the zero leading dimension.
64 for yi in out.iter_mut() {
65 *yi = if beta == 0.0 { 0.0 } else { beta * *yi };
66 }
67 return Ok(out.into());
68 }
69 let xv = Vector::view(x);
70 igraph_call!(igraph_blas_dgemv(
71 transpose,
72 alpha,
73 a,
74 xv.as_ptr(),
75 beta,
76 &mut out
77 ))?;
78 Ok(out.into())
79}
80
81/// In-place matrix-vector product `y <- alpha * op(A) * x + beta * y` on
82/// plain slices (`igraph_blas_dgemv_array`).
83///
84/// Binds [`igraph_blas_dgemv_array`](https://igraph.org/c/html/latest/igraph-Linalg.html#igraph_blas_dgemv_array).
85///
86/// # Errors
87/// If the lengths of `x` and `y` do not match `op(A)`.
88pub fn blas_dgemv_array(
89 transpose: bool,
90 alpha: f64,
91 a: &Matrix,
92 x: &[f64],
93 beta: f64,
94 y: &mut [f64],
95) -> Result<()> {
96 dgemv_dims(transpose, a, x.len(), y.len())?;
97 if a.is_empty() {
98 // op(A) x is an empty sum; BLAS would reject the zero leading dimension.
99 for yi in y.iter_mut() {
100 *yi = if beta == 0.0 { 0.0 } else { beta * *yi };
101 }
102 return Ok(());
103 }
104 igraph_call!(igraph_blas_dgemv_array(
105 transpose,
106 alpha,
107 a,
108 x.as_ptr(),
109 beta,
110 y.as_mut_ptr()
111 ))
112}
113
114/// Matrix-matrix product `alpha * op(A) * op(B) + beta * C`, where `op(X)`
115/// is `X` or its transpose (`igraph_blas_dgemm`). `c` may be `None` when
116/// `beta` is zero (or to mean a zero matrix).
117///
118/// igraph 1.0.0 and 1.0.1 check the shape of `C` against the wrong dimension when
119/// `beta != 0`; this wrapper validates the shapes itself and adds `beta * C`
120/// on the Rust side when igraph would reject a correct call.
121///
122/// Time complexity: O(nmk) for an `n` × `k` times `k` × `m` product.
123///
124/// Binds [`igraph_blas_dgemm`](https://igraph.org/c/html/latest/igraph-Linalg.html#igraph_blas_dgemm).
125///
126/// # Examples
127///
128/// ```
129/// use igraph::{linalg::blas_dgemm, prelude::*};
130/// let a = Matrix::from_rows(&[[1.0, 2.0, 3.0]]).unwrap(); // 1x3
131/// let b = Matrix::from_rows(&[[1.0], [1.0], [1.0]]).unwrap(); // 3x1
132/// let ab = blas_dgemm(false, false, 1.0, &a, &b, 0.0, None).unwrap();
133/// assert_eq!(ab.to_rows(), vec![vec![6.0]]);
134/// // The outer product B A is 3x3.
135/// let ba = blas_dgemm(false, false, 1.0, &b, &a, 0.0, None).unwrap();
136/// assert_eq!(ba.shape(), (3, 3));
137/// // A' A with transposition flags.
138/// let ata = blas_dgemm(true, false, 1.0, &a, &a, 0.0, None).unwrap();
139/// assert_eq!(ata[(2, 1)], 6.0);
140/// ```
141pub fn blas_dgemm(
142 transpose_a: bool,
143 transpose_b: bool,
144 alpha: f64,
145 a: &Matrix,
146 b: &Matrix,
147 beta: f64,
148 c: Option<&Matrix>,
149) -> Result<Matrix> {
150 for m in [a, b] {
151 to_c_int(m.nrow(), "number of rows")?;
152 to_c_int(m.ncol(), "number of columns")?;
153 }
154 let (m, k) = if transpose_a {
155 (a.ncol(), a.nrow())
156 } else {
157 (a.nrow(), a.ncol())
158 };
159 let (kb, n) = if transpose_b {
160 (b.ncol(), b.nrow())
161 } else {
162 (b.nrow(), b.ncol())
163 };
164 if k != kb {
165 return Err(Error::invalid(format!(
166 "dgemm: {m}-by-{k} and {kb}-by-{n} matrices cannot be multiplied"
167 )));
168 }
169 let c = c.filter(|_| beta != 0.0);
170 if let Some(c) = c
171 && c.shape() != (m, n)
172 {
173 return Err(Error::invalid(format!(
174 "dgemm: C is {:?}, expected {:?}",
175 c.shape(),
176 (m, n)
177 )));
178 }
179 let mut res = Matrix::zeros(m, n);
180 if m > 0 && n > 0 && k > 0 {
181 // igraph 1.0.0 and 1.0.1 compare C's shape with (m, k) when beta != 0: only let it
182 // add beta * C when that (buggy) check agrees with the real shape.
183 if let Some(c) = c
184 && k == n
185 {
186 res = c.clone();
187 igraph_call!(igraph_blas_dgemm(
188 transpose_a,
189 transpose_b,
190 alpha,
191 a,
192 b,
193 beta,
194 &mut res
195 ))?;
196 return Ok(res);
197 }
198 igraph_call!(igraph_blas_dgemm(
199 transpose_a,
200 transpose_b,
201 alpha,
202 a,
203 b,
204 0.0,
205 &mut res
206 ))?;
207 }
208 if let Some(c) = c {
209 for (r, &cv) in res.as_mut_slice().iter_mut().zip(c.as_slice()) {
210 *r += beta * cv;
211 }
212 }
213 Ok(res)
214}
215
216/// Euclidean norm of a vector (`igraph_blas_dnrm2`), computed without
217/// undue overflow or underflow.
218///
219/// Binds [`igraph_blas_dnrm2`](https://igraph.org/c/html/latest/igraph-Linalg.html#igraph_blas_dnrm2).
220///
221/// ```
222/// use igraph::linalg::blas_dnrm2;
223/// assert_eq!(blas_dnrm2(&[3.0, 4.0]).unwrap(), 5.0);
224/// ```
225pub fn blas_dnrm2(v: &[f64]) -> Result<f64> {
226 to_c_int(v.len(), "vector length")?;
227 crate::error::ensure_init();
228 let view = Vector::view(v);
229 Ok(unsafe { igraph_blas_dnrm2(view.as_ptr()) })
230}
231
232/// Dot product of two vectors of the same length (`igraph_blas_ddot`).
233///
234/// Binds [`igraph_blas_ddot`](https://igraph.org/c/html/latest/igraph-Linalg.html#igraph_blas_ddot).
235///
236/// ```
237/// use igraph::linalg::blas_ddot;
238/// assert_eq!(blas_ddot(&[1.0, 2.0, 3.0], &[4.0, 5.0, 6.0]).unwrap(), 32.0);
239/// ```
240///
241/// # Errors
242/// If the lengths differ.
243pub fn blas_ddot(a: &[f64], b: &[f64]) -> Result<f64> {
244 to_c_int(a.len(), "vector length")?;
245 if a.len() != b.len() {
246 return Err(Error::invalid(
247 "dot product of vectors with different lengths",
248 ));
249 }
250 let (va, vb) = (Vector::view(a), Vector::view(b));
251 let mut res = 0.0;
252 igraph_call!(igraph_blas_ddot(va.as_ptr(), vb.as_ptr(), &mut res))?;
253 Ok(res)
254}