1use tenferro_ops::broadcast::{broadcast_error_to_validation, broadcast_shape, broadcast_shapes};
8use tenferro_tensor::validate::matmul_config_for_shapes;
9use tenferro_tensor::{BackendSession, CompareDir, DType, Error, Result, TensorRead};
10
11use crate::typed_tensor::{broadcast_to_in_read, ReadInput};
12
13use crate::TensorSessionOpsExt;
14use tenferro_tensor::Tensor;
15
16impl TensorSessionOpsExt for Tensor {
17 fn add(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
18 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
19 session.add_read(lhs.tensor_read(), rhs.tensor_read())
20 }
21
22 fn mul(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
23 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
24 session.mul_read(lhs.tensor_read(), rhs.tensor_read())
25 }
26
27 fn exp(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
28 session.exp(self)
29 }
30
31 fn reduce_sum(&self, axes: &[usize], session: &mut dyn BackendSession) -> Result<Tensor> {
32 session.reduce_sum(self, axes)
33 }
34
35 fn convert(&self, to: DType, session: &mut dyn BackendSession) -> Result<Tensor> {
36 session.convert(self, to)
37 }
38
39 fn cast(&self, to: DType, session: &mut dyn BackendSession) -> Result<Tensor> {
40 session.cast(self, to)
41 }
42
43 fn sub(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
44 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
45 session.sub_read(lhs.tensor_read(), rhs.tensor_read())
46 }
47
48 fn div(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
49 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
50 session.div_read(lhs.tensor_read(), rhs.tensor_read())
51 }
52
53 fn rem(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
54 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
55 session.rem_read(lhs.tensor_read(), rhs.tensor_read())
56 }
57
58 fn pow(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
59 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
60 session.pow_read(lhs.tensor_read(), rhs.tensor_read())
61 }
62
63 fn maximum(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
64 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
65 session.maximum_read(lhs.tensor_read(), rhs.tensor_read())
66 }
67
68 fn minimum(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
69 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
70 session.minimum_read(lhs.tensor_read(), rhs.tensor_read())
71 }
72
73 fn neg(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
74 session.neg(self)
75 }
76
77 fn abs(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
78 session.abs(self)
79 }
80
81 fn sign(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
82 session.sign(self)
83 }
84
85 fn conj(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
86 session.conj(self)
87 }
88
89 fn log(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
90 session.log(self)
91 }
92
93 fn expm1(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
94 session.expm1(self)
95 }
96
97 fn log1p(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
98 session.log1p(self)
99 }
100
101 fn sin(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
102 session.sin(self)
103 }
104
105 fn cos(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
106 session.cos(self)
107 }
108
109 fn tanh(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
110 session.tanh(self)
111 }
112
113 fn sqrt(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
114 session.sqrt(self)
115 }
116
117 fn rsqrt(&self, session: &mut dyn BackendSession) -> Result<Tensor> {
118 session.rsqrt(self)
119 }
120
121 fn compare(
122 &self,
123 rhs: &Tensor,
124 dir: CompareDir,
125 session: &mut dyn BackendSession,
126 ) -> Result<Tensor> {
127 let (lhs, rhs) = broadcast_binary_in(self, rhs, session)?;
128 session.compare_read(lhs.tensor_read(), rhs.tensor_read(), &dir)
129 }
130
131 fn where_select(
132 &self,
133 on_true: &Tensor,
134 on_false: &Tensor,
135 session: &mut dyn BackendSession,
136 ) -> Result<Tensor> {
137 let (condition, on_true, on_false) =
138 broadcast_ternary_in(self, on_true, on_false, session)?;
139 session.select_read(
140 condition.tensor_read(),
141 on_true.tensor_read(),
142 on_false.tensor_read(),
143 )
144 }
145
146 fn clamp(
147 &self,
148 lower: &Tensor,
149 upper: &Tensor,
150 session: &mut dyn BackendSession,
151 ) -> Result<Tensor> {
152 let (input, lower, upper) = broadcast_ternary_in(self, lower, upper, session)?;
153 session.clamp_read(
154 input.tensor_read(),
155 lower.tensor_read(),
156 upper.tensor_read(),
157 )
158 }
159
160 fn matmul(&self, rhs: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
161 let config = matmul_config_for_shapes("matmul", self.shape(), rhs.shape())?;
162 session.dot_general(self, rhs, &config)
163 }
164
165 fn reshape(&self, shape: &[usize], session: &mut dyn BackendSession) -> Result<Tensor> {
166 session.reshape(self, shape)
167 }
168
169 fn transpose(&self, perm: &[usize], session: &mut dyn BackendSession) -> Result<Tensor> {
170 session.transpose(self, perm)
171 }
172}
173
174fn broadcast_binary_in<'a>(
175 lhs: &'a Tensor,
176 rhs: &'a Tensor,
177 session: &mut dyn BackendSession,
178) -> Result<(ReadInput<'a>, ReadInput<'a>)> {
179 let shape = broadcast_shape(lhs.shape(), rhs.shape()).map_err(broadcast_error)?;
180 Ok((
181 broadcast_to_in_read(TensorRead::from_tensor(lhs), &shape, session)?,
182 broadcast_to_in_read(TensorRead::from_tensor(rhs), &shape, session)?,
183 ))
184}
185
186fn broadcast_ternary_in<'a>(
187 first: &'a Tensor,
188 second: &'a Tensor,
189 third: &'a Tensor,
190 session: &mut dyn BackendSession,
191) -> Result<(ReadInput<'a>, ReadInput<'a>, ReadInput<'a>)> {
192 let shape = broadcast_shapes([first.shape(), second.shape(), third.shape()])
193 .map_err(broadcast_error)?;
194 Ok((
195 broadcast_to_in_read(TensorRead::from_tensor(first), &shape, session)?,
196 broadcast_to_in_read(TensorRead::from_tensor(second), &shape, session)?,
197 broadcast_to_in_read(TensorRead::from_tensor(third), &shape, session)?,
198 ))
199}
200
201fn broadcast_error(err: tenferro_ops::broadcast::BroadcastError) -> Error {
202 Error::validation("broadcast", broadcast_error_to_validation(err))
203}