Skip to main content

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}