Skip to main content

mdarray_linalg/naive/contract/
context.rs

1use std::iter::Sum;
2use std::ops::AddAssign;
3
4use mdarray::{Array, Dim, DynRank, Layout, Shape, Slice};
5use num_complex::ComplexFloat;
6use num_traits::{MulAdd, One, Zero};
7
8use super::simple::naive_matmul;
9use crate::{
10    Naive,
11    contract::{
12        _contract, _hypercontract, einsum_to_contract_axes, Axes, Contract, ContractBuilder,
13        MatmulBuilder,
14    },
15};
16
17struct NaiveMatmulBuilder<'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}
29
30struct NaiveContractBuilder<'a, T, Sa, Sb, La, Lb>
31where
32    La: Layout,
33    Lb: Layout,
34    Sa: Shape,
35    Sb: Shape,
36{
37    alpha: T,
38    a: &'a Slice<T, Sa, La>,
39    b: &'a Slice<T, Sb, Lb>,
40    axes: Axes<'a>,
41    einsum: bool,
42    einsum_axes_a: Option<Vec<usize>>,
43    einsum_axes_b: Option<Vec<usize>>,
44    current_output_labels: Option<Vec<u8>>,
45    requested_output_labels: Option<Vec<u8>>,
46}
47
48impl<'a, T, D0, D1, D2, La, Lb> MatmulBuilder<'a, T, D0, D1, D2, La, Lb>
49    for NaiveMatmulBuilder<'a, T, D0, D1, D2, La, Lb>
50where
51    La: Layout,
52    Lb: Layout,
53    T: ComplexFloat + Zero + One + MulAdd<Output = T>,
54    D0: Dim,
55    D1: Dim,
56    D2: Dim,
57{
58    fn scale(mut self, factor: T) -> Self {
59        self.alpha = self.alpha * factor;
60        self
61    }
62
63    fn eval(self) -> Array<T, (D0, D2)> {
64        let (m, _) = *self.a.shape();
65        let (_, n) = *self.b.shape();
66        let mut c = Array::from_elem((m, n), T::zero());
67        naive_matmul(self.alpha, self.a, self.b, T::zero(), &mut c);
68        c
69    }
70
71    fn write<Lc: Layout>(self, c: &mut Slice<T, (D0, D2), Lc>) {
72        naive_matmul(self.alpha, self.a, self.b, T::zero(), c);
73    }
74
75    fn add_to<Lc: Layout>(self, c: &mut Slice<T, (D0, D2), Lc>) {
76        naive_matmul(self.alpha, self.a, self.b, T::one(), c);
77    }
78
79    fn add_to_scaled<Lc: Layout>(self, c: &mut Slice<T, (D0, D2), Lc>, beta: T) {
80        naive_matmul(self.alpha, self.a, self.b, beta, c);
81    }
82}
83
84impl<'a, T, Sa, Sb, La, Lb> ContractBuilder<'a, T, Sa, Sb, La, Lb>
85    for NaiveContractBuilder<'a, T, Sa, Sb, La, Lb>
86where
87    La: Layout,
88    Lb: Layout,
89    T: ComplexFloat + Zero + One + MulAdd<Output = T> + AddAssign + Sum,
90    Sa: Shape,
91    Sb: Shape,
92{
93    fn scale(mut self, factor: T) -> Self {
94        self.alpha = self.alpha * factor;
95        self
96    }
97
98    fn eval(self) -> Array<T, DynRank> {
99        if self.einsum {
100            let a = self.a.to_array().into_dyn();
101            let b = self.b.to_array().into_dyn();
102
103            // TODO: it is very likely that this copy is useless. It
104            // should be removed in a near future (9th april
105            // 2026). However due to tricky typing error I did not
106            // manage to give the view to _hypercontract without
107            // removing the genericity toward Layout and shape :(
108
109            let axes_a = self
110                .einsum_axes_a
111                .as_deref()
112                .expect("missing einsum axis labels for A");
113            let axes_b = self
114                .einsum_axes_b
115                .as_deref()
116                .expect("missing einsum axis labels for B");
117
118            let mut result = _hypercontract(Naive, a.expr(), b.expr(), axes_a, axes_b);
119
120            if let (Some(current), Some(requested)) = (
121                self.current_output_labels.as_deref(),
122                self.requested_output_labels.as_deref(),
123            )
124                && current != requested
125            {
126                    let perm: Vec<usize> = requested
127                        .iter()
128                        .map(|label| {
129                            current
130                                .iter()
131                                .position(|cur| cur == label)
132                                .expect("output label not present in contraction result")
133                        })
134                        .collect();
135
136                    result = result.permute(perm).to_tensor().into_dyn();
137            }
138
139            if self.alpha != T::one() {
140                result = result.map(|x| x * self.alpha).into_dyn();
141            }
142
143            result
144        } else {
145            _contract(Naive, self.a, self.b, self.axes, self.alpha)
146        }
147    }
148
149    fn write<Sc: Shape, Lc: Layout>(self, c: &mut Slice<T, Sc, Lc>) {
150        let result = self.eval();
151        assert_eq!(c.rank(), result.rank(), "output rank mismatch");
152        for i in 0..c.rank() {
153            assert_eq!(
154                c.dim(i),
155                result.dim(i),
156                "output shape mismatch on axis {i}: expected {}, got {}",
157                result.dim(i),
158                c.dim(i)
159            );
160        }
161        for (dst, src) in c.iter_mut().zip(result.iter()) {
162            *dst = *src;
163        }
164    }
165
166    fn add_to<Sc: Shape, Lc: Layout>(self, c: &mut Slice<T, Sc, Lc>) {
167        self.add_to_scaled(c, T::one())
168    }
169
170    fn add_to_scaled<Sc: Shape, Lc: Layout>(self, c: &mut Slice<T, Sc, Lc>, beta: T) {
171        let result = self.eval();
172        assert_eq!(c.rank(), result.rank(), "output rank mismatch");
173        for i in 0..c.rank() {
174            assert_eq!(
175                c.dim(i),
176                result.dim(i),
177                "output shape mismatch on axis {i}: expected {}, got {}",
178                result.dim(i),
179                c.dim(i)
180            );
181        }
182        for (dst, src) in c.iter_mut().zip(result.iter()) {
183            *dst = beta * *dst + *src;
184        }
185    }
186}
187
188impl<T> Contract<T> for Naive
189where
190    T: ComplexFloat + Zero + One + MulAdd<Output = T> + AddAssign + Sum,
191{
192    fn matmul<'a, D0, D1, D2, La, Lb>(
193        &self,
194        a: &'a Slice<T, (D0, D1), La>,
195        b: &'a Slice<T, (D1, D2), Lb>,
196    ) -> impl MatmulBuilder<'a, T, D0, D1, D2, La, Lb>
197    where
198        La: Layout,
199        Lb: Layout,
200        D0: Dim,
201        D1: Dim,
202        D2: Dim,
203    {
204        NaiveMatmulBuilder {
205            alpha: T::one(),
206            a,
207            b,
208        }
209    }
210
211    fn contract_all<'a, Sa, Sb, La, Lb>(
212        &self,
213        a: &'a Slice<T, Sa, La>,
214        b: &'a Slice<T, Sb, Lb>,
215    ) -> T
216    where
217        T: 'a,
218        Sa: Shape,
219        Sb: Shape,
220        La: Layout,
221        Lb: Layout,
222    {
223        _contract(Naive, a, b, Axes::All, T::one()).into_scalar()
224    }
225
226    fn contract_n<'a, Sa, Sb, La, Lb>(
227        &self,
228        a: &'a Slice<T, Sa, La>,
229        b: &'a Slice<T, Sb, Lb>,
230        n: usize,
231    ) -> impl ContractBuilder<'a, T, Sa, Sb, La, Lb>
232    where
233        T: 'a,
234        Sa: Shape,
235        Sb: Shape,
236        La: Layout,
237        Lb: Layout,
238    {
239        NaiveContractBuilder {
240            alpha: T::one(),
241            a,
242            b,
243            axes: Axes::LastFirst { k: n },
244            einsum: false,
245            einsum_axes_a: None,
246            einsum_axes_b: None,
247            current_output_labels: None,
248            requested_output_labels: None,
249        }
250    }
251
252    fn contract_pairs<'a, Sa, Sb, La, Lb>(
253        &self,
254        a: &'a Slice<T, Sa, La>,
255        b: &'a Slice<T, Sb, Lb>,
256        axes_a: &'a [usize],
257        axes_b: &'a [usize],
258    ) -> impl ContractBuilder<'a, T, Sa, Sb, La, Lb>
259    where
260        T: 'a,
261        Sa: Shape,
262        Sb: Shape,
263        La: Layout,
264        Lb: Layout,
265    {
266        NaiveContractBuilder {
267            alpha: T::one(),
268            a,
269            b,
270            axes: Axes::Specific(axes_a, axes_b),
271            einsum: false,
272            einsum_axes_a: None,
273            einsum_axes_b: None,
274            current_output_labels: None,
275            requested_output_labels: None,
276        }
277    }
278
279    fn contract<'a, Sa, Sb, La, Lb>(
280        &self,
281        a: &'a Slice<T, Sa, La>,
282        b: &'a Slice<T, Sb, Lb>,
283        indices_a: &'a [u8],
284        indices_b: &'a [u8],
285        indices_c: &'a [u8],
286    ) -> impl ContractBuilder<'a, T, Sa, Sb, La, Lb>
287    where
288        T: 'a,
289        Sa: Shape,
290        Sb: Shape,
291        La: Layout,
292        Lb: Layout,
293    {
294        assert_eq!(
295            indices_a.len(),
296            a.rank(),
297            "einsum indices_a length ({}) must match A rank ({})",
298            indices_a.len(),
299            a.rank()
300        );
301        assert_eq!(
302            indices_b.len(),
303            b.rank(),
304            "einsum indices_b length ({}) must match B rank ({})",
305            indices_b.len(),
306            b.rank()
307        );
308
309        let free: std::collections::HashSet<u8> = indices_c.iter().copied().collect();
310        let current_output_labels: Vec<u8> = indices_a
311            .iter()
312            .chain(indices_b.iter())
313            .copied()
314            .filter(|label| free.contains(label))
315            .collect();
316        let (einsum_axes_a, einsum_axes_b) =
317            einsum_to_contract_axes(indices_a, indices_b, indices_c);
318
319        NaiveContractBuilder {
320            alpha: T::one(),
321            a,
322            b,
323            axes: Axes::SpecificOwned(Vec::new(), Vec::new()),
324            einsum: true,
325            einsum_axes_a: Some(einsum_axes_a),
326            einsum_axes_b: Some(einsum_axes_b),
327            current_output_labels: Some(current_output_labels),
328            requested_output_labels: Some(indices_c.to_vec()),
329        }
330    }
331}