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 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}