1pub fn scalar_operand<'a, T: crate::TensorScalar>(
44 tensor: &'a crate::Tensor,
45 op: &'static str,
46 backend: impl core::fmt::Display,
47) -> crate::Result<&'a crate::TypedTensor<T>> {
48 tensor.as_typed::<T>().ok_or_else(|| {
49 crate::Error::unsupported_dtype(
50 op,
51 tensor.dtype(),
52 format!("backend {backend} does not support an externally defined scalar"),
53 )
54 })
55}
56
57#[macro_export]
58macro_rules! with_scalar {
59 ($tensor:expr, all, backend = $backend:expr, op = $op:expr, |$typed:ident| $(-> $ret:ty)? $body:block) => {{
60 let _ = &$backend;
61 let _ = &$op;
62 match $tensor.dtype() {
63 $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>($tensor, $op, $backend)
64 .and_then(|$typed| $body),
65 $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>($tensor, $op, $backend)
66 .and_then(|$typed| $body),
67 $crate::DType::I32 => $crate::dispatch::scalar_operand::<i32>($tensor, $op, $backend)
68 .and_then(|$typed| $body),
69 $crate::DType::I64 => $crate::dispatch::scalar_operand::<i64>($tensor, $op, $backend)
70 .and_then(|$typed| $body),
71 $crate::DType::Bool => $crate::dispatch::scalar_operand::<bool>($tensor, $op, $backend)
72 .and_then(|$typed| $body),
73 $crate::DType::C32 => {
74 $crate::dispatch::scalar_operand::<$crate::Complex32>($tensor, $op, $backend)
75 .and_then(|$typed| $body)
76 }
77 $crate::DType::C64 => {
78 $crate::dispatch::scalar_operand::<$crate::Complex64>($tensor, $op, $backend)
79 .and_then(|$typed| $body)
80 }
81 $crate::DType::External(type_id) => Err($crate::Error::unsupported_dtype(
83 $op,
84 $crate::DType::External(type_id),
85 format!(
86 "backend {} does not support an externally defined scalar",
87 $backend
88 ),
89 )),
90 }
91 }};
92 ($tensor:expr, numeric, backend = $backend:expr, op = $op:expr, |$typed:ident| $(-> $ret:ty)? $body:block) => {{
93 match $tensor.dtype() {
94 $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>($tensor, $op, $backend)
95 .and_then(|$typed| $body),
96 $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>($tensor, $op, $backend)
97 .and_then(|$typed| $body),
98 $crate::DType::I32 => $crate::dispatch::scalar_operand::<i32>($tensor, $op, $backend)
99 .and_then(|$typed| $body),
100 $crate::DType::I64 => $crate::dispatch::scalar_operand::<i64>($tensor, $op, $backend)
101 .and_then(|$typed| $body),
102 $crate::DType::C32 => {
103 $crate::dispatch::scalar_operand::<$crate::Complex32>($tensor, $op, $backend)
104 .and_then(|$typed| $body)
105 }
106 $crate::DType::C64 => {
107 $crate::dispatch::scalar_operand::<$crate::Complex64>($tensor, $op, $backend)
108 .and_then(|$typed| $body)
109 }
110 $crate::DType::External(type_id) => Err($crate::Error::unsupported_dtype(
112 $op,
113 $crate::DType::External(type_id),
114 format!(
115 "backend {} does not support an externally defined scalar",
116 $backend
117 ),
118 )),
119 other => Err($crate::Error::unsupported_dtype(
120 $op,
121 other,
122 format!("backend {} does not support this operation/dtype", $backend),
123 )),
124 }
125 }};
126 ($tensor:expr, float_complex, backend = $backend:expr, op = $op:expr, |$typed:ident| $(-> $ret:ty)? $body:block) => {{
127 match $tensor.dtype() {
128 $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>($tensor, $op, $backend)
129 .and_then(|$typed| $body),
130 $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>($tensor, $op, $backend)
131 .and_then(|$typed| $body),
132 $crate::DType::C32 => {
133 $crate::dispatch::scalar_operand::<$crate::Complex32>($tensor, $op, $backend)
134 .and_then(|$typed| $body)
135 }
136 $crate::DType::C64 => {
137 $crate::dispatch::scalar_operand::<$crate::Complex64>($tensor, $op, $backend)
138 .and_then(|$typed| $body)
139 }
140 $crate::DType::External(type_id) => Err($crate::Error::unsupported_dtype(
142 $op,
143 $crate::DType::External(type_id),
144 format!(
145 "backend {} does not support an externally defined scalar",
146 $backend
147 ),
148 )),
149 other => Err($crate::Error::unsupported_dtype(
150 $op,
151 other,
152 format!("backend {} does not support this operation/dtype", $backend),
153 )),
154 }
155 }};
156 ($tensor:expr, float_only, backend = $backend:expr, op = $op:expr, |$typed:ident| $(-> $ret:ty)? $body:block) => {{
157 match $tensor.dtype() {
158 $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>($tensor, $op, $backend)
159 .and_then(|$typed| $body),
160 $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>($tensor, $op, $backend)
161 .and_then(|$typed| $body),
162 other => Err($crate::Error::unsupported_dtype(
163 $op,
164 other,
165 format!("backend {} does not support this operation/dtype", $backend),
166 )),
167 }
168 }};
169}
170
171#[macro_export]
191macro_rules! with_scalar_read {
192 ($read:expr, all, backend = $backend:expr, op = $op:expr, |$view:ident| $(-> $ret:ty)? $body:block) => {{
193 let _ = &$backend;
194 let _ = &$op;
195 match $read {
196 $crate::TensorRead::Tensor(tensor) => match tensor.dtype() {
197 $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>(
198 tensor, $op, $backend,
199 )
200 .and_then(|tensor| {
201 let $view = tensor.as_view();
202 $body
203 }),
204 $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>(
205 tensor, $op, $backend,
206 )
207 .and_then(|tensor| {
208 let $view = tensor.as_view();
209 $body
210 }),
211 $crate::DType::I32 => $crate::dispatch::scalar_operand::<i32>(
212 tensor, $op, $backend,
213 )
214 .and_then(|tensor| {
215 let $view = tensor.as_view();
216 $body
217 }),
218 $crate::DType::I64 => $crate::dispatch::scalar_operand::<i64>(
219 tensor, $op, $backend,
220 )
221 .and_then(|tensor| {
222 let $view = tensor.as_view();
223 $body
224 }),
225 $crate::DType::Bool => $crate::dispatch::scalar_operand::<bool>(
226 tensor, $op, $backend,
227 )
228 .and_then(|tensor| {
229 let $view = tensor.as_view();
230 $body
231 }),
232 $crate::DType::C32 => {
233 $crate::dispatch::scalar_operand::<$crate::Complex32>(tensor, $op, $backend)
234 .and_then(|tensor| {
235 let $view = tensor.as_view();
236 $body
237 })
238 }
239 $crate::DType::C64 => {
240 $crate::dispatch::scalar_operand::<$crate::Complex64>(tensor, $op, $backend)
241 .and_then(|tensor| {
242 let $view = tensor.as_view();
243 $body
244 })
245 }
246 $crate::DType::External(type_id) => Err($crate::Error::unsupported_dtype(
248 $op,
249 $crate::DType::External(type_id),
250 format!(
251 "backend {} does not support an externally defined scalar",
252 $backend
253 ),
254 )),
255 },
256 $crate::TensorRead::View(view) => match view {
257 $crate::TensorView::F32($view) => $body,
258 $crate::TensorView::F64($view) => $body,
259 $crate::TensorView::I32($view) => $body,
260 $crate::TensorView::I64($view) => $body,
261 $crate::TensorView::Bool($view) => $body,
262 $crate::TensorView::C32($view) => $body,
263 $crate::TensorView::C64($view) => $body,
264 },
265 }
266 }};
267 ($read:expr, numeric, backend = $backend:expr, op = $op:expr, |$view:ident| $(-> $ret:ty)? $body:block) => {{
268 let read = $read;
269 let dtype = read.dtype();
270 match read {
271 $crate::TensorRead::Tensor(tensor) => match tensor.dtype() {
272 $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>(
273 tensor, $op, $backend,
274 )
275 .and_then(|tensor| {
276 let $view = tensor.as_view();
277 $body
278 }),
279 $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>(
280 tensor, $op, $backend,
281 )
282 .and_then(|tensor| {
283 let $view = tensor.as_view();
284 $body
285 }),
286 $crate::DType::I32 => $crate::dispatch::scalar_operand::<i32>(
287 tensor, $op, $backend,
288 )
289 .and_then(|tensor| {
290 let $view = tensor.as_view();
291 $body
292 }),
293 $crate::DType::I64 => $crate::dispatch::scalar_operand::<i64>(
294 tensor, $op, $backend,
295 )
296 .and_then(|tensor| {
297 let $view = tensor.as_view();
298 $body
299 }),
300 $crate::DType::C32 => {
301 $crate::dispatch::scalar_operand::<$crate::Complex32>(tensor, $op, $backend)
302 .and_then(|tensor| {
303 let $view = tensor.as_view();
304 $body
305 })
306 }
307 $crate::DType::C64 => {
308 $crate::dispatch::scalar_operand::<$crate::Complex64>(tensor, $op, $backend)
309 .and_then(|tensor| {
310 let $view = tensor.as_view();
311 $body
312 })
313 }
314 $crate::DType::Bool => Err($crate::Error::unsupported_dtype(
315 $op,
316 dtype,
317 format!("backend {} does not support this operation/dtype", $backend),
318 )),
319 $crate::DType::External(type_id) => Err($crate::Error::unsupported_dtype(
320 $op,
321 $crate::DType::External(type_id),
322 format!(
323 "backend {} does not support an externally defined scalar",
324 $backend
325 ),
326 )),
327 },
328 $crate::TensorRead::View(view) => match view {
329 $crate::TensorView::F32($view) => $body,
330 $crate::TensorView::F64($view) => $body,
331 $crate::TensorView::I32($view) => $body,
332 $crate::TensorView::I64($view) => $body,
333 $crate::TensorView::C32($view) => $body,
334 $crate::TensorView::C64($view) => $body,
335 $crate::TensorView::Bool(_) => Err($crate::Error::unsupported_dtype(
336 $op,
337 dtype,
338 format!("backend {} does not support this operation/dtype", $backend),
339 )),
340 },
341 }
342 }};
343 ($read:expr, float_complex, backend = $backend:expr, op = $op:expr, |$view:ident| $(-> $ret:ty)? $body:block) => {{
344 let read = $read;
345 let dtype = read.dtype();
346 match read {
347 $crate::TensorRead::Tensor(tensor) => match tensor.dtype() {
348 $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>(
349 tensor, $op, $backend,
350 )
351 .and_then(|tensor| {
352 let $view = tensor.as_view();
353 $body
354 }),
355 $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>(
356 tensor, $op, $backend,
357 )
358 .and_then(|tensor| {
359 let $view = tensor.as_view();
360 $body
361 }),
362 $crate::DType::C32 => {
363 $crate::dispatch::scalar_operand::<$crate::Complex32>(tensor, $op, $backend)
364 .and_then(|tensor| {
365 let $view = tensor.as_view();
366 $body
367 })
368 }
369 $crate::DType::C64 => {
370 $crate::dispatch::scalar_operand::<$crate::Complex64>(tensor, $op, $backend)
371 .and_then(|tensor| {
372 let $view = tensor.as_view();
373 $body
374 })
375 }
376 $crate::DType::I32 | $crate::DType::I64 | $crate::DType::Bool => {
377 Err($crate::Error::unsupported_dtype(
378 $op,
379 dtype,
380 format!("backend {} does not support this operation/dtype", $backend),
381 ))
382 }
383 $crate::DType::External(type_id) => Err($crate::Error::unsupported_dtype(
384 $op,
385 $crate::DType::External(type_id),
386 format!(
387 "backend {} does not support an externally defined scalar",
388 $backend
389 ),
390 )),
391 },
392 $crate::TensorRead::View(view) => match view {
393 $crate::TensorView::F32($view) => $body,
394 $crate::TensorView::F64($view) => $body,
395 $crate::TensorView::C32($view) => $body,
396 $crate::TensorView::C64($view) => $body,
397 $crate::TensorView::I32(_)
398 | $crate::TensorView::I64(_)
399 | $crate::TensorView::Bool(_) => Err($crate::Error::unsupported_dtype(
400 $op,
401 dtype,
402 format!("backend {} does not support this operation/dtype", $backend),
403 )),
404 },
405 }
406 }};
407 ($read:expr, float_only, backend = $backend:expr, op = $op:expr, |$view:ident| $(-> $ret:ty)? $body:block) => {{
408 let read = $read;
409 let dtype = read.dtype();
410 match read {
411 $crate::TensorRead::Tensor(tensor) => match tensor.dtype() {
412 $crate::DType::F32 => $crate::dispatch::scalar_operand::<f32>(
413 tensor, $op, $backend,
414 )
415 .and_then(|tensor| {
416 let $view = tensor.as_view();
417 $body
418 }),
419 $crate::DType::F64 => $crate::dispatch::scalar_operand::<f64>(
420 tensor, $op, $backend,
421 )
422 .and_then(|tensor| {
423 let $view = tensor.as_view();
424 $body
425 }),
426 $crate::DType::I32
427 | $crate::DType::I64
428 | $crate::DType::Bool
429 | $crate::DType::C32
430 | $crate::DType::C64
431 | $crate::DType::External(_) => Err($crate::Error::unsupported_dtype(
432 $op,
433 dtype,
434 format!("backend {} does not support this operation/dtype", $backend),
435 )),
436 },
437 $crate::TensorRead::View(view) => match view {
438 $crate::TensorView::F32($view) => $body,
439 $crate::TensorView::F64($view) => $body,
440 $crate::TensorView::I32(_)
441 | $crate::TensorView::I64(_)
442 | $crate::TensorView::Bool(_)
443 | $crate::TensorView::C32(_)
444 | $crate::TensorView::C64(_) => Err($crate::Error::unsupported_dtype(
445 $op,
446 dtype,
447 format!("backend {} does not support this operation/dtype", $backend),
448 )),
449 },
450 }
451 }};
452}