Skip to main content

mdarray_linalg_blas/contract/
context.rs

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