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}