1use alloc::vec::Vec;
4use ruda_core::tensor::{DType, element::Element};
5use ruda_core::{bytes::Bytes, tensor::Shape};
6use half::{bf16, f16};
7
8use ruda_core::tensor::host::HostTensor;
9use ruda_core::tensor::host::layout::Layout;
10use ruda_core::tensor::host::strided_index::StridedIter;
11
12#[cfg(feature = "simd")]
13use crate::simd;
14
15#[derive(Clone, Copy)]
17pub enum BinaryOp {
18 Add,
19 Sub,
20 Mul,
21 Div,
22}
23
24pub fn binary_op<F32Op, F64Op>(
32 lhs: HostTensor,
33 rhs: HostTensor,
34 f32_op: F32Op,
35 f64_op: F64Op,
36 simd_hint: Option<BinaryOp>,
37) -> HostTensor
38where
39 F32Op: Fn(f32, f32) -> f32 + Copy,
40 F64Op: Fn(f64, f64) -> f64 + Copy,
41{
42 debug_assert_eq!(lhs.dtype(), rhs.dtype(), "binary_op: dtype mismatch");
43
44 let (lhs, rhs) = crate::expand::broadcast_binary(lhs, rhs);
46
47 let dtype = lhs.dtype();
48
49 match dtype {
50 DType::F32 => binary_op_f32(lhs, &rhs, f32_op, simd_hint),
51 DType::F64 => binary_op_typed(lhs, &rhs, f64_op),
52 DType::F16 => binary_op_typed(lhs, &rhs, |a: f16, b: f16| {
53 f16::from_f32(f32_op(a.to_f32(), b.to_f32()))
54 }),
55 DType::BF16 => binary_op_typed(lhs, &rhs, |a: bf16, b: bf16| {
56 bf16::from_f32(f32_op(a.to_f32(), b.to_f32()))
57 }),
58 _ => panic!("binary_op: unsupported dtype {:?}", dtype),
59 }
60}
61
62#[cfg(feature = "simd")]
63mod broadcast;
64#[cfg(feature = "simd")]
65use broadcast::*;
66
67#[cfg(not(feature = "simd"))]
69fn binary_op_f32<Op>(
70 lhs: HostTensor,
71 rhs: &HostTensor,
72 op: Op,
73 _simd_hint: Option<BinaryOp>,
74) -> HostTensor
75where
76 Op: Fn(f32, f32) -> f32,
77{
78 binary_op_typed(lhs, rhs, op)
79}
80
81pub fn binary_op_typed<E, Op>(mut lhs: HostTensor, rhs: &HostTensor, op: Op) -> HostTensor
83where
84 E: Element + bytemuck::Pod,
85 Op: Fn(E, E) -> E,
86{
87 let rhs_storage: &[E] = rhs.storage();
88
89 if lhs.is_unique()
91 && let (Some((0, l_end)), Some((r_start, r_end))) = (
92 lhs.layout().contiguous_offsets(),
93 rhs.layout().contiguous_offsets(),
94 )
95 {
96 let lhs_storage: &mut [E] = lhs.storage_mut();
97 let r_slice = &rhs_storage[r_start..r_end];
98 for (l, &r) in lhs_storage[..l_end].iter_mut().zip(r_slice) {
99 *l = op(*l, r);
100 }
101 return lhs;
102 }
103
104 let shape = lhs.layout().shape().clone();
106 let dtype = lhs.dtype();
107 let lhs_storage: &[E] = lhs.storage();
108
109 let result: Vec<E> = match (
110 lhs.layout().contiguous_offsets(),
111 rhs.layout().contiguous_offsets(),
112 ) {
113 (Some((l_start, l_end)), Some((r_start, r_end))) => {
115 let l_slice = &lhs_storage[l_start..l_end];
116 let r_slice = &rhs_storage[r_start..r_end];
117 l_slice
118 .iter()
119 .zip(r_slice)
120 .map(|(&a, &b)| op(a, b))
121 .collect()
122 }
123 _ if lhs.layout().num_dims() == 2 => {
125 apply_2d_strided(lhs_storage, rhs_storage, lhs.layout(), rhs.layout(), op)
126 }
127 _ => {
129 let lhs_iter = StridedIter::new(lhs.layout());
130 let rhs_iter = StridedIter::new(rhs.layout());
131 lhs_iter
132 .zip(rhs_iter)
133 .map(|(li, ri)| op(lhs_storage[li], rhs_storage[ri]))
134 .collect()
135 }
136 };
137
138 make_tensor(result, shape, dtype)
139}
140
141#[inline]
143pub(crate) fn apply_2d_strided<E, R, Op>(
144 lhs: &[E],
145 rhs: &[E],
146 lhs_layout: &Layout,
147 rhs_layout: &Layout,
148 op: Op,
149) -> Vec<R>
150where
151 E: Copy,
152 Op: Fn(E, E) -> R,
153{
154 let (rows, cols, l_row_stride, l_col_stride) = lhs_layout.as_2d_strides().unwrap();
155 let (_, _, r_row_stride, r_col_stride) = rhs_layout.as_2d_strides().unwrap();
156 let l_offset = lhs_layout.start_offset() as isize;
157 let r_offset = rhs_layout.start_offset() as isize;
158
159 let mut result = Vec::with_capacity(rows * cols);
160
161 for row in 0..rows {
162 let l_row_start = l_offset + row as isize * l_row_stride;
163 let r_row_start = r_offset + row as isize * r_row_stride;
164 for col in 0..cols {
165 let l_idx = (l_row_start + col as isize * l_col_stride) as usize;
166 let r_idx = (r_row_start + col as isize * r_col_stride) as usize;
167 result.push(op(lhs[l_idx], rhs[r_idx]));
168 }
169 }
170
171 result
172}
173
174pub fn scalar_op<F32Op, F64Op>(
178 tensor: HostTensor,
179 scalar: f64,
180 f32_op: F32Op,
181 f64_op: F64Op,
182) -> HostTensor
183where
184 F32Op: Fn(f32, f32) -> f32 + Copy,
185 F64Op: Fn(f64, f64) -> f64 + Copy,
186{
187 let dtype = tensor.dtype();
188
189 match dtype {
190 DType::F32 => scalar_op_typed(tensor, scalar as f32, f32_op),
191 DType::F64 => scalar_op_typed(tensor, scalar, f64_op),
192 DType::F16 => {
193 let scalar_f16 = f16::from_f32(scalar as f32);
194 let s = scalar_f16.to_f32();
195 scalar_op_typed(tensor, scalar_f16, |a: f16, _| {
196 f16::from_f32(f32_op(a.to_f32(), s))
197 })
198 }
199 DType::BF16 => {
200 let scalar_bf16 = bf16::from_f32(scalar as f32);
201 let s = scalar_bf16.to_f32();
202 scalar_op_typed(tensor, scalar_bf16, |a: bf16, _| {
203 bf16::from_f32(f32_op(a.to_f32(), s))
204 })
205 }
206 _ => panic!("scalar_op: unsupported dtype {:?}", dtype),
207 }
208}
209
210pub fn scalar_op_typed<E, Op>(mut tensor: HostTensor, scalar: E, op: Op) -> HostTensor
211where
212 E: Element + bytemuck::Pod,
213 Op: Fn(E, E) -> E,
214{
215 if tensor.is_unique()
217 && let Some((0, end)) = tensor.layout().contiguous_offsets()
218 {
219 let storage: &mut [E] = tensor.storage_mut();
220 for x in storage[..end].iter_mut() {
221 *x = op(*x, scalar);
222 }
223 return tensor;
224 }
225
226 let shape = tensor.layout().shape().clone();
228 let dtype = tensor.dtype();
229 let storage: &[E] = tensor.storage();
230
231 let result: Vec<E> = match tensor.layout().contiguous_offsets() {
232 Some((start, end)) => storage[start..end].iter().map(|&x| op(x, scalar)).collect(),
233 None => StridedIter::new(tensor.layout())
234 .map(|i| op(storage[i], scalar))
235 .collect(),
236 };
237
238 make_tensor(result, shape, dtype)
239}
240
241fn make_tensor<E: bytemuck::Pod + Send + Sync>(
243 data: Vec<E>,
244 shape: Shape,
245 dtype: DType,
246) -> HostTensor {
247 let bytes = Bytes::from_elems(data);
248 let layout = Layout::contiguous(shape);
249 HostTensor::new(bytes, layout, dtype)
250}
251
252pub fn int_binary_op<Op>(lhs: HostTensor, rhs: HostTensor, op: Op) -> HostTensor
256where
257 Op: Fn(i64, i64) -> i64 + Copy,
258{
259 debug_assert_eq!(lhs.dtype(), rhs.dtype(), "int_binary_op: dtype mismatch");
260
261 let (lhs, rhs) = crate::expand::broadcast_binary(lhs, rhs);
263
264 let dtype = lhs.dtype();
265
266 match dtype {
267 DType::I64 => binary_op_typed(lhs, &rhs, op),
268 DType::I32 => binary_op_typed(lhs, &rhs, |a: i32, b: i32| op(a as i64, b as i64) as i32),
269 DType::I16 => binary_op_typed(lhs, &rhs, |a: i16, b: i16| op(a as i64, b as i64) as i16),
270 DType::I8 => binary_op_typed(lhs, &rhs, |a: i8, b: i8| op(a as i64, b as i64) as i8),
271 DType::U64 => binary_op_typed(lhs, &rhs, |a: u64, b: u64| op(a as i64, b as i64) as u64),
274 DType::U32 => binary_op_typed(lhs, &rhs, |a: u32, b: u32| op(a as i64, b as i64) as u32),
275 DType::U16 => binary_op_typed(lhs, &rhs, |a: u16, b: u16| op(a as i64, b as i64) as u16),
276 DType::U8 => binary_op_typed(lhs, &rhs, |a: u8, b: u8| op(a as i64, b as i64) as u8),
277 _ => panic!("int_binary_op: unsupported dtype {:?}", dtype),
278 }
279}
280
281pub fn int_scalar_op<Op>(tensor: HostTensor, scalar: i64, op: Op) -> HostTensor
284where
285 Op: Fn(i64, i64) -> i64 + Copy,
286{
287 let dtype = tensor.dtype();
288
289 match dtype {
290 DType::I64 => scalar_op_typed(tensor, scalar, op),
291 DType::I32 => scalar_op_typed(tensor, scalar as i32, |a: i32, b: i32| {
292 op(a as i64, b as i64) as i32
293 }),
294 DType::I16 => scalar_op_typed(tensor, scalar as i16, |a: i16, b: i16| {
295 op(a as i64, b as i64) as i16
296 }),
297 DType::I8 => scalar_op_typed(tensor, scalar as i8, |a: i8, b: i8| {
298 op(a as i64, b as i64) as i8
299 }),
300 DType::U64 => scalar_op_typed(tensor, scalar as u64, |a: u64, b: u64| {
301 op(a as i64, b as i64) as u64
302 }),
303 DType::U32 => scalar_op_typed(tensor, scalar as u32, |a: u32, b: u32| {
304 op(a as i64, b as i64) as u32
305 }),
306 DType::U16 => scalar_op_typed(tensor, scalar as u16, |a: u16, b: u16| {
307 op(a as i64, b as i64) as u16
308 }),
309 DType::U8 => scalar_op_typed(tensor, scalar as u8, |a: u8, b: u8| {
310 op(a as i64, b as i64) as u8
311 }),
312 _ => panic!("int_scalar_op: unsupported dtype {:?}", dtype),
313 }
314}
315
316#[cfg(test)]
326mod tests;
327
328pub mod dispatch_float;
329pub mod dispatch_int;
330mod integer_power;