Skip to main content

mdarray_linalg_faer/
contract.rs

1use std::iter::Sum;
2use std::ops::AddAssign;
3
4use faer::{Accum, Par, linalg::matmul::matmul};
5use faer_traits::ComplexField;
6use mdarray::{Array, Dim, DynRank, Layout, Shape, Slice};
7use mdarray_linalg::contract::{
8    _contract, _hypercontract, einsum_to_contract_axes, extract_axes, Axes, Contract,
9    ContractAxes, ContractBuilder, MatmulBuilder,
10};
11use mdarray_linalg::{finish_contraction, prepare_contraction};
12use num_complex::ComplexFloat;
13use num_traits::{MulAdd, One, Zero};
14
15use crate::{Faer, into_faer, into_faer_mut};
16
17struct FaerMatmulBuilder<'a, T, D0, D1, D2, La, Lb>
18where
19    La: Layout,
20    Lb: Layout,
21    D0: Dim,
22    D1: Dim,
23    D2: Dim,
24{
25    alpha: T,
26    a: &'a Slice<T, (D0, D1), La>,
27    b: &'a Slice<T, (D1, D2), Lb>,
28    par: Par,
29}
30
31struct FaerContractBuilder<'a, T, Sa, Sb, La, Lb>
32where
33    La: Layout,
34    Lb: Layout,
35    Sa: Shape,
36    Sb: Shape,
37{
38    alpha: T,
39    a: &'a Slice<T, Sa, La>,
40    b: &'a Slice<T, Sb, Lb>,
41    axes: Axes<'a>,
42    par: Par,
43    einsum: bool,
44    einsum_axes_a: Option<Vec<usize>>,
45    einsum_axes_b: Option<Vec<usize>>,
46    current_output_labels: Option<Vec<u8>>,
47    requested_output_labels: Option<Vec<u8>>,
48}
49
50impl<'a, T, D0, D1, D2, La, Lb> MatmulBuilder<'a, T, D0, D1, D2, La, Lb>
51    for FaerMatmulBuilder<'a, T, D0, D1, D2, La, Lb>
52where
53    La: Layout,
54    Lb: Layout,
55    D0: Dim,
56    D1: Dim,
57    D2: Dim,
58    T: ComplexFloat + ComplexField + One + Zero + 'static,
59{
60    fn scale(mut self, factor: T) -> Self {
61        self.alpha *= factor;
62        self
63    }
64
65    fn eval(self) -> Array<T, (D0, D2)> {
66        let (ma, _) = *self.a.shape();
67        let (_, nb) = *self.b.shape();
68
69        let a_faer = into_faer(self.a);
70        let b_faer = into_faer(self.b);
71
72        let mut c = Array::<T, (D0, D2)>::from_elem((ma, nb), T::zero());
73        let mut c_faer = into_faer_mut(&mut c);
74
75        matmul(
76            &mut c_faer,
77            Accum::Replace,
78            a_faer,
79            b_faer,
80            self.alpha,
81            self.par,
82        );
83
84        c
85    }
86
87    fn write<Lc: Layout>(self, c: &mut Slice<T, (D0, D2), Lc>) {
88        let mut c_faer = into_faer_mut(c);
89        matmul(
90            &mut c_faer,
91            Accum::Replace,
92            into_faer(self.a),
93            into_faer(self.b),
94            self.alpha,
95            self.par,
96        );
97    }
98
99    fn add_to<Lc: Layout>(self, c: &mut Slice<T, (D0, D2), Lc>) {
100        let mut c_faer = into_faer_mut(c);
101        matmul(
102            &mut c_faer,
103            Accum::Add,
104            into_faer(self.a),
105            into_faer(self.b),
106            self.alpha,
107            self.par,
108        );
109    }
110
111    fn add_to_scaled<Lc: Layout>(self, c: &mut Slice<T, (D0, D2), Lc>, beta: T) {
112        for value in c.iter_mut() {
113            *value = beta * *value;
114        }
115
116        let mut c_faer = into_faer_mut(c);
117        matmul(
118            &mut c_faer,
119            Accum::Add,
120            into_faer(self.a),
121            into_faer(self.b),
122            self.alpha,
123            self.par,
124        );
125    }
126}
127
128impl<'a, T, Sa, Sb, La, Lb> ContractBuilder<'a, T, Sa, Sb, La, Lb>
129    for FaerContractBuilder<'a, T, Sa, Sb, La, Lb>
130where
131    La: Layout,
132    Lb: Layout,
133    T: ComplexFloat + ComplexField + Zero + One + 'static + MulAdd<Output = T> + AddAssign + Sum,
134    Sa: Shape,
135    Sb: Shape,
136{
137    fn scale(mut self, factor: T) -> Self {
138        self.alpha *= factor;
139        self
140    }
141
142    fn eval(self) -> Array<T, DynRank> {
143        if self.einsum {
144            let a = self.a.to_array().into_dyn();
145            let b = self.b.to_array().into_dyn();
146
147            let axes_a = self
148                .einsum_axes_a
149                .as_deref()
150                .expect("missing einsum axis labels for A");
151            let axes_b = self
152                .einsum_axes_b
153                .as_deref()
154                .expect("missing einsum axis labels for B");
155
156            let mut result = _hypercontract(Faer::default(), a.expr(), b.expr(), axes_a, axes_b);
157
158            if let (Some(current), Some(requested)) = (
159                self.current_output_labels.as_deref(),
160                self.requested_output_labels.as_deref(),
161            )
162                && current != requested
163            {
164                    let perm: Vec<usize> = requested
165                        .iter()
166                        .map(|label| {
167                            current
168                                .iter()
169                                .position(|cur| cur == label)
170                                .expect("output label not present in contraction result")
171                        })
172                        .collect();
173                    result = result.permute(perm).to_tensor().into_dyn();
174            }
175
176            if self.alpha != T::one() {
177                result = result.map(|x| x * self.alpha).into_dyn();
178            }
179
180            result
181        } else {
182            let (a_2d, b_2d, keep_shape_a, keep_shape_b) =
183                prepare_contraction!(self.axes, self.a, self.b);
184            let a_faer = into_faer(&a_2d);
185            let b_faer = into_faer(&b_2d);
186
187            let (m, _) = *a_2d.shape();
188            let (_, n) = *b_2d.shape();
189            let mut c = Array::<T, (usize, usize)>::from_elem((m, n), T::zero());
190            let mut c_faer = into_faer_mut(&mut c);
191
192            matmul(
193                &mut c_faer,
194                Accum::Replace,
195                a_faer,
196                b_faer,
197                self.alpha,
198                self.par,
199            );
200
201            finish_contraction!(c, keep_shape_a, keep_shape_b)
202        }
203    }
204
205    fn write<Sc: Shape, Lc: Layout>(self, c: &mut Slice<T, Sc, Lc>) {
206        let result = self.eval();
207        assert_eq!(c.rank(), result.rank(), "output rank mismatch");
208        for i in 0..c.rank() {
209            assert_eq!(c.dim(i), result.dim(i), "output shape mismatch on axis {i}");
210        }
211        for (dst, src) in c.iter_mut().zip(result.iter()) {
212            *dst = *src;
213        }
214    }
215
216    fn add_to<Sc: Shape, Lc: Layout>(self, c: &mut Slice<T, Sc, Lc>) {
217        self.add_to_scaled(c, T::one())
218    }
219
220    fn add_to_scaled<Sc: Shape, Lc: Layout>(self, c: &mut Slice<T, Sc, Lc>, beta: T) {
221        let result = self.eval();
222        assert_eq!(c.rank(), result.rank(), "output rank mismatch");
223        for i in 0..c.rank() {
224            assert_eq!(c.dim(i), result.dim(i), "output shape mismatch on axis {i}");
225        }
226        for (dst, src) in c.iter_mut().zip(result.iter()) {
227            *dst = beta * *dst + *src;
228        }
229    }
230}
231
232impl<T> Contract<T> for Faer
233where
234    T: ComplexFloat + ComplexField + Zero + One + 'static + MulAdd<Output = T> + AddAssign + Sum,
235{
236    fn matmul<'a, D0, D1, D2, La, Lb>(
237        &self,
238        a: &'a Slice<T, (D0, D1), La>,
239        b: &'a Slice<T, (D1, D2), Lb>,
240    ) -> impl MatmulBuilder<'a, T, D0, D1, D2, La, Lb>
241    where
242        La: Layout,
243        Lb: Layout,
244        D0: Dim,
245        D1: Dim,
246        D2: Dim,
247    {
248        FaerMatmulBuilder {
249            alpha: T::one(),
250            a,
251            b,
252            par: faer::get_global_parallelism(),
253        }
254    }
255
256    fn contract_all<'a, Sa, Sb, La, Lb>(
257        &self,
258        a: &'a Slice<T, Sa, La>,
259        b: &'a Slice<T, Sb, Lb>,
260    ) -> T
261    where
262        T: 'a,
263        Sa: Shape,
264        Sb: Shape,
265        La: Layout,
266        Lb: Layout,
267    {
268        _contract(Faer, a, b, Axes::All, T::one()).into_scalar()
269    }
270
271    fn contract_n<'a, Sa, Sb, La, Lb>(
272        &self,
273        a: &'a Slice<T, Sa, La>,
274        b: &'a Slice<T, Sb, Lb>,
275        n: usize,
276    ) -> impl ContractBuilder<'a, T, Sa, Sb, La, Lb>
277    where
278        T: 'a,
279        Sa: Shape,
280        Sb: Shape,
281        La: Layout,
282        Lb: Layout,
283    {
284        FaerContractBuilder {
285            alpha: T::one(),
286            a,
287            b,
288            axes: Axes::LastFirst { k: n },
289            par: faer::get_global_parallelism(),
290            einsum: false,
291            einsum_axes_a: None,
292            einsum_axes_b: None,
293            current_output_labels: None,
294            requested_output_labels: None,
295        }
296    }
297
298    fn contract_pairs<'a, Sa, Sb, La, Lb>(
299        &self,
300        a: &'a Slice<T, Sa, La>,
301        b: &'a Slice<T, Sb, Lb>,
302        axes_a: &'a [usize],
303        axes_b: &'a [usize],
304    ) -> impl ContractBuilder<'a, T, Sa, Sb, La, Lb>
305    where
306        T: 'a,
307        Sa: Shape,
308        Sb: Shape,
309        La: Layout,
310        Lb: Layout,
311    {
312        FaerContractBuilder {
313            alpha: T::one(),
314            a,
315            b,
316            axes: Axes::Specific(axes_a, axes_b),
317            par: faer::get_global_parallelism(),
318            einsum: false,
319            einsum_axes_a: None,
320            einsum_axes_b: None,
321            current_output_labels: None,
322            requested_output_labels: None,
323        }
324    }
325
326    fn contract<'a, Sa, Sb, La, Lb>(
327        &self,
328        a: &'a Slice<T, Sa, La>,
329        b: &'a Slice<T, Sb, Lb>,
330        indices_a: &'a [u8],
331        indices_b: &'a [u8],
332        indices_c: &'a [u8],
333    ) -> impl ContractBuilder<'a, T, Sa, Sb, La, Lb>
334    where
335        T: 'a,
336        Sa: Shape,
337        Sb: Shape,
338        La: Layout,
339        Lb: Layout,
340    {
341        assert_eq!(indices_a.len(), a.rank(), "einsum indices_a length must match A rank");
342        assert_eq!(indices_b.len(), b.rank(), "einsum indices_b length must match B rank");
343
344        let free: std::collections::HashSet<u8> = indices_c.iter().copied().collect();
345        let current_output_labels: Vec<u8> = indices_a
346            .iter()
347            .chain(indices_b.iter())
348            .copied()
349            .filter(|label| free.contains(label))
350            .collect();
351        let (einsum_axes_a, einsum_axes_b) =
352            einsum_to_contract_axes(indices_a, indices_b, indices_c);
353
354        FaerContractBuilder {
355            alpha: T::one(),
356            a,
357            b,
358            axes: Axes::SpecificOwned(Vec::new(), Vec::new()),
359            par: faer::get_global_parallelism(),
360            einsum: true,
361            einsum_axes_a: Some(einsum_axes_a),
362            einsum_axes_b: Some(einsum_axes_b),
363            current_output_labels: Some(current_output_labels),
364            requested_output_labels: Some(indices_c.to_vec()),
365        }
366    }
367}