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}