1use std::hash::{Hash, Hasher};
2use std::sync::Arc;
3
4#[cfg(all(test, feature = "autodiff"))]
5use crate::ad::{ADRuleResult, PrimitiveTransposeInput};
6#[cfg(all(test, feature = "autodiff"))]
7use computegraph::types::{LocalValueId, OperationRole, ValueKey};
8use computegraph::GraphOperation;
9use num_complex::{Complex32, Complex64};
10
11use crate::dim_expr::DimExpr;
12use crate::ext_op::{ext_op_eq, hash_extension, ExtensionOp};
13use crate::input_key::TensorInputKey;
14use tenferro_tensor::{
15 CompareDir, DType, DotGeneralConfig, GatherConfig, PadConfig, ScatterConfig, SliceConfig,
16 TensorScalar,
17};
18
19pub trait ConstantScalar: TensorScalar + private::Sealed {
29 fn constant_bytes(self) -> Vec<u8>;
39}
40
41mod private {
42 pub trait Sealed {}
43
44 impl Sealed for f64 {}
45 impl Sealed for f32 {}
46 impl Sealed for i64 {}
47 impl Sealed for i32 {}
48 impl Sealed for bool {}
49 impl Sealed for num_complex::Complex64 {}
50 impl Sealed for num_complex::Complex32 {}
51}
52
53impl ConstantScalar for f64 {
54 fn constant_bytes(self) -> Vec<u8> {
55 self.to_le_bytes().to_vec()
56 }
57}
58
59impl ConstantScalar for f32 {
60 fn constant_bytes(self) -> Vec<u8> {
61 self.to_le_bytes().to_vec()
62 }
63}
64
65impl ConstantScalar for i64 {
66 fn constant_bytes(self) -> Vec<u8> {
67 self.to_le_bytes().to_vec()
68 }
69}
70
71impl ConstantScalar for i32 {
72 fn constant_bytes(self) -> Vec<u8> {
73 self.to_le_bytes().to_vec()
74 }
75}
76
77impl ConstantScalar for bool {
78 fn constant_bytes(self) -> Vec<u8> {
79 vec![u8::from(self)]
80 }
81}
82
83impl ConstantScalar for Complex64 {
84 fn constant_bytes(self) -> Vec<u8> {
85 let mut bytes = Vec::with_capacity(16);
86 bytes.extend_from_slice(&self.re.to_le_bytes());
87 bytes.extend_from_slice(&self.im.to_le_bytes());
88 bytes
89 }
90}
91
92impl ConstantScalar for Complex32 {
93 fn constant_bytes(self) -> Vec<u8> {
94 let mut bytes = Vec::with_capacity(8);
95 bytes.extend_from_slice(&self.re.to_le_bytes());
96 bytes.extend_from_slice(&self.im.to_le_bytes());
97 bytes
98 }
99}
100
101tenferro_core_ops::define_std_tensor_op!();
102
103impl StdTensorOp {
104 pub fn constant<T: ConstantScalar>(value: T) -> Self {
120 Self::Constant {
121 dtype: T::dtype(),
122 bytes: value.constant_bytes(),
123 }
124 }
125}
126
127impl PartialEq for StdTensorOp {
128 fn eq(&self, other: &Self) -> bool {
129 if std::mem::discriminant(self) != std::mem::discriminant(other) {
130 return false;
131 }
132 match (self, other) {
133 (Self::Add, Self::Add)
134 | (Self::Sub, Self::Sub)
135 | (Self::Mul, Self::Mul)
136 | (Self::Neg, Self::Neg)
137 | (Self::Conj, Self::Conj)
138 | (Self::Div, Self::Div)
139 | (Self::Rem, Self::Rem)
140 | (Self::Abs, Self::Abs)
141 | (Self::Sign, Self::Sign)
142 | (Self::Maximum, Self::Maximum)
143 | (Self::Minimum, Self::Minimum)
144 | (Self::Select, Self::Select)
145 | (Self::Clamp, Self::Clamp)
146 | (Self::Exp, Self::Exp)
147 | (Self::Log, Self::Log)
148 | (Self::Sin, Self::Sin)
149 | (Self::Cos, Self::Cos)
150 | (Self::Tanh, Self::Tanh)
151 | (Self::Sqrt, Self::Sqrt)
152 | (Self::Rsqrt, Self::Rsqrt)
153 | (Self::Pow, Self::Pow)
154 | (Self::Expm1, Self::Expm1)
155 | (Self::Log1p, Self::Log1p)
156 | (Self::Erf, Self::Erf)
157 | (Self::DynamicUpdateSlice, Self::DynamicUpdateSlice) => true,
158 (Self::DotGeneral { config: a }, Self::DotGeneral { config: b }) => a == b,
159 (Self::Transpose { perm: a }, Self::Transpose { perm: b }) => a == b,
160 (Self::Reshape { to_shape: a }, Self::Reshape { to_shape: b }) => a == b,
161 (
162 Self::BroadcastInDim {
163 shape: sa,
164 dims: da,
165 },
166 Self::BroadcastInDim {
167 shape: sb,
168 dims: db,
169 },
170 ) => sa == sb && da == db,
171 (Self::Convert { from: fa, to: ta }, Self::Convert { from: fb, to: tb }) => {
172 fa == fb && ta == tb
173 }
174 (
175 Self::Constant {
176 dtype: da,
177 bytes: ba,
178 },
179 Self::Constant {
180 dtype: db,
181 bytes: bb,
182 },
183 ) => da == db && ba == bb,
184 (Self::ReduceSum { axes: a }, Self::ReduceSum { axes: b })
185 | (Self::ReduceSumSquares { axes: a }, Self::ReduceSumSquares { axes: b })
186 | (Self::ReduceProd { axes: a }, Self::ReduceProd { axes: b })
187 | (Self::ReduceMax { axes: a }, Self::ReduceMax { axes: b })
188 | (Self::ReduceMin { axes: a }, Self::ReduceMin { axes: b })
189 | (Self::Reverse { axes: a }, Self::Reverse { axes: b }) => a == b,
190 (Self::Compare(a), Self::Compare(b)) => a == b,
191 (
192 Self::ExtractDiag {
193 axis_a: aa,
194 axis_b: ba,
195 },
196 Self::ExtractDiag {
197 axis_a: ab,
198 axis_b: bb,
199 },
200 )
201 | (
202 Self::EmbedDiag {
203 axis_a: aa,
204 axis_b: ba,
205 },
206 Self::EmbedDiag {
207 axis_a: ab,
208 axis_b: bb,
209 },
210 ) => aa == ab && ba == bb,
211 (Self::Tril { k: a }, Self::Tril { k: b })
212 | (Self::Triu { k: a }, Self::Triu { k: b }) => a == b,
213 (Self::Gather(a), Self::Gather(b)) => a == b,
214 (
215 Self::GatherDynamicSliceSizes {
216 offset_dims: oa,
217 collapsed_slice_dims: ca,
218 start_index_map: sa,
219 index_vector_dim: ia,
220 slice_sizes: za,
221 },
222 Self::GatherDynamicSliceSizes {
223 offset_dims: ob,
224 collapsed_slice_dims: cb,
225 start_index_map: sb,
226 index_vector_dim: ib,
227 slice_sizes: zb,
228 },
229 ) => oa == ob && ca == cb && sa == sb && ia == ib && za == zb,
230 (Self::Scatter(a), Self::Scatter(b)) => a == b,
231 (Self::Slice(a), Self::Slice(b)) => a == b,
232 (Self::DynamicSlice { slice_sizes: a }, Self::DynamicSlice { slice_sizes: b }) => {
233 a == b
234 }
235 (Self::Pad(a), Self::Pad(b)) => a == b,
236 (
237 Self::Concatenate {
238 axis: a,
239 input_count: na,
240 },
241 Self::Concatenate {
242 axis: b,
243 input_count: nb,
244 },
245 ) => a == b && na == nb,
246 (Self::ShapeOf { axis: a }, Self::ShapeOf { axis: b })
247 | (Self::DynamicTruncate { axis: a }, Self::DynamicTruncate { axis: b })
248 | (Self::PadToMatch { axis: a }, Self::PadToMatch { axis: b }) => a == b,
249 (Self::Extension(a), Self::Extension(b)) => ext_op_eq(a.as_ref(), b.as_ref()),
250 _ => false,
251 }
252 }
253}
254
255impl Eq for StdTensorOp {}
256
257impl Hash for StdTensorOp {
258 fn hash<H: Hasher>(&self, state: &mut H) {
259 std::mem::discriminant(self).hash(state);
260 match self {
261 Self::Add
262 | Self::Sub
263 | Self::Mul
264 | Self::Neg
265 | Self::Conj
266 | Self::Div
267 | Self::Rem
268 | Self::Abs
269 | Self::Sign
270 | Self::Maximum
271 | Self::Minimum
272 | Self::Select
273 | Self::Clamp
274 | Self::Exp
275 | Self::Log
276 | Self::Sin
277 | Self::Cos
278 | Self::Tanh
279 | Self::Sqrt
280 | Self::Rsqrt
281 | Self::Pow
282 | Self::Expm1
283 | Self::Log1p
284 | Self::Erf => {}
285 Self::DotGeneral { config } => {
286 config.hash(state);
287 }
288 Self::Transpose { perm } => perm.hash(state),
289 Self::Reshape { to_shape } => {
290 to_shape.hash(state);
291 }
292 Self::BroadcastInDim { shape, dims } => {
293 shape.hash(state);
294 dims.hash(state);
295 }
296 Self::Convert { from, to } => {
297 from.hash(state);
298 to.hash(state);
299 }
300 Self::Constant { dtype, bytes } => {
301 dtype.hash(state);
302 bytes.hash(state);
303 }
304 Self::ReduceSum { axes } | Self::ReduceSumSquares { axes } => {
305 axes.hash(state);
306 }
307 Self::Compare(dir) => dir.hash(state),
308 Self::ExtractDiag { axis_a, axis_b } | Self::EmbedDiag { axis_a, axis_b } => {
309 axis_a.hash(state);
310 axis_b.hash(state);
311 }
312 Self::Tril { k } | Self::Triu { k } => k.hash(state),
313 Self::Gather(config) => config.hash(state),
314 Self::GatherDynamicSliceSizes {
315 offset_dims,
316 collapsed_slice_dims,
317 start_index_map,
318 index_vector_dim,
319 slice_sizes,
320 } => {
321 offset_dims.hash(state);
322 collapsed_slice_dims.hash(state);
323 start_index_map.hash(state);
324 index_vector_dim.hash(state);
325 slice_sizes.hash(state);
326 }
327 Self::Scatter(config) => config.hash(state),
328 Self::Slice(config) => config.hash(state),
329 Self::DynamicSlice { slice_sizes } => slice_sizes.hash(state),
330 Self::DynamicUpdateSlice => {}
331 Self::Pad(config) => config.hash(state),
332 Self::Concatenate { axis, input_count } => {
333 axis.hash(state);
334 input_count.hash(state);
335 }
336 Self::Reverse { axes } => axes.hash(state),
337 Self::ShapeOf { axis } | Self::DynamicTruncate { axis } | Self::PadToMatch { axis } => {
338 axis.hash(state)
339 }
340 Self::ReduceProd { axes } | Self::ReduceMax { axes } | Self::ReduceMin { axes } => {
341 axes.hash(state);
342 }
343 Self::Extension(op) => hash_extension(op.as_ref(), state),
344 }
345 }
346}
347
348fn n_inputs_from_dim_exprs(min_inputs: usize, exprs: &[&[DimExpr]]) -> usize {
349 let max_idx = exprs
350 .iter()
351 .flat_map(|exprs| exprs.iter())
352 .filter_map(DimExpr::max_input_idx)
353 .max()
354 .map_or(0, |max_idx| max_idx + 1);
355 max_idx.max(min_inputs)
356}
357
358impl GraphOperation for StdTensorOp {
359 type Operand = ();
362 type Context = ();
363 type InputKey = TensorInputKey;
364
365 fn input_count(&self) -> usize {
366 match self {
367 Self::Add | Self::Sub | Self::Mul | Self::DotGeneral { .. } | Self::Gather(_) => 2,
368 Self::GatherDynamicSliceSizes { slice_sizes, .. } => {
369 n_inputs_from_dim_exprs(2, &[slice_sizes])
370 }
371 Self::Neg
372 | Self::Conj
373 | Self::Transpose { .. }
374 | Self::Convert { .. }
375 | Self::ExtractDiag { .. }
376 | Self::EmbedDiag { .. }
377 | Self::Tril { .. }
378 | Self::Triu { .. }
379 | Self::Slice(_)
380 | Self::Pad(_)
381 | Self::Reverse { .. }
382 | Self::ShapeOf { .. } => 1,
383 Self::DynamicTruncate { .. } | Self::PadToMatch { .. } => 2,
384 Self::Reshape { to_shape } => n_inputs_from_dim_exprs(1, &[to_shape]),
385 Self::BroadcastInDim { shape, .. } => n_inputs_from_dim_exprs(1, &[shape]),
386 Self::ReduceSum { .. }
387 | Self::ReduceSumSquares { .. }
388 | Self::ReduceProd { .. }
389 | Self::ReduceMax { .. }
390 | Self::ReduceMin { .. } => 1,
391 Self::Div
392 | Self::Rem
393 | Self::Maximum
394 | Self::Minimum
395 | Self::Pow
396 | Self::DynamicSlice { .. } => 2,
397 Self::Constant { .. } => 0,
398 Self::Scatter(_) | Self::DynamicUpdateSlice => 3,
399 Self::Concatenate { input_count, .. } => *input_count,
400 Self::Abs
401 | Self::Sign
402 | Self::Exp
403 | Self::Log
404 | Self::Sin
405 | Self::Cos
406 | Self::Tanh
407 | Self::Sqrt
408 | Self::Rsqrt
409 | Self::Expm1
410 | Self::Log1p
411 | Self::Erf => 1,
412 Self::Select | Self::Clamp => 3,
413 Self::Compare(_) => 2,
414 Self::Extension(op) => ExtensionOp::input_count(op.as_ref()),
415 }
416 }
417
418 fn output_count(&self) -> usize {
419 match self {
420 Self::Add
421 | Self::Sub
422 | Self::Mul
423 | Self::Neg
424 | Self::Conj
425 | Self::DotGeneral { .. }
426 | Self::Transpose { .. }
427 | Self::Reshape { .. }
428 | Self::BroadcastInDim { .. }
429 | Self::Convert { .. }
430 | Self::ReduceSum { .. }
431 | Self::ReduceSumSquares { .. }
432 | Self::Div
433 | Self::Rem
434 | Self::Abs
435 | Self::Sign
436 | Self::Maximum
437 | Self::Minimum
438 | Self::Compare(_)
439 | Self::Select
440 | Self::Clamp
441 | Self::Constant { .. }
442 | Self::Exp
443 | Self::Log
444 | Self::Sin
445 | Self::Cos
446 | Self::Tanh
447 | Self::Sqrt
448 | Self::Rsqrt
449 | Self::Pow
450 | Self::Expm1
451 | Self::Log1p
452 | Self::Erf
453 | Self::ExtractDiag { .. }
454 | Self::EmbedDiag { .. }
455 | Self::Tril { .. }
456 | Self::Triu { .. }
457 | Self::Gather(_)
458 | Self::GatherDynamicSliceSizes { .. }
459 | Self::Scatter(_)
460 | Self::Slice(_)
461 | Self::DynamicSlice { .. }
462 | Self::DynamicUpdateSlice
463 | Self::Pad(_)
464 | Self::Reverse { .. }
465 | Self::ShapeOf { .. }
466 | Self::DynamicTruncate { .. }
467 | Self::PadToMatch { .. }
468 | Self::ReduceProd { .. }
469 | Self::ReduceMax { .. }
470 | Self::ReduceMin { .. } => 1,
471 Self::Concatenate { .. } => 1,
472 Self::Extension(op) => ExtensionOp::output_count(op.as_ref()),
473 }
474 }
475}
476
477#[cfg(all(test, feature = "autodiff"))]
478impl StdTensorOp {
479 pub(crate) fn jvp_rule(
480 &self,
481 builder: &mut computegraph::graph::GraphBuilder<Self>,
482 primal_in: &[ValueKey<Self>],
483 primal_out: &[ValueKey<Self>],
484 tangent_in: &[Option<LocalValueId>],
485 ctx: &mut crate::ad::context::ShapeGuardContext,
486 ) -> ADRuleResult<Vec<Option<LocalValueId>>> {
487 crate::ad::linearize(self, builder, primal_in, primal_out, tangent_in, ctx)
488 }
489
490 pub(crate) fn transpose_rule(
491 &self,
492 builder: &mut computegraph::graph::GraphBuilder<Self>,
493 cotangent_out: &[Option<LocalValueId>],
494 inputs: &[computegraph::ValueRef<Self>],
495 mode: &OperationRole,
496 ctx: &mut crate::ad::context::ShapeGuardContext,
497 ) -> ADRuleResult<Vec<Option<LocalValueId>>> {
498 let inputs = inputs
499 .iter()
500 .map(|input| match input {
501 computegraph::ValueRef::Local(local_id) => {
502 let key = builder.global_key(*local_id).clone();
503 PrimitiveTransposeInput::Residual(key)
504 }
505 computegraph::ValueRef::External(key) => {
506 PrimitiveTransposeInput::Residual(key.clone())
507 }
508 })
509 .collect::<Vec<_>>();
510 crate::ad::transpose_rule(self, builder, cotangent_out, inputs.as_slice(), mode, ctx)
511 }
512}