Skip to main content

ruprim_host/binary/
mod.rs

1//! Binary tensor operations (add, sub, mul, div).
2
3use 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/// Operation type for SIMD dispatch.
16#[derive(Clone, Copy)]
17pub enum BinaryOp {
18    Add,
19    Sub,
20    Mul,
21    Div,
22}
23
24/// Apply a binary operation element-wise to two tensors.
25///
26/// Requires tensors to have the same shape. Uses SIMD acceleration for f32
27/// when available and both tensors are contiguous.
28///
29/// Pass `simd_hint` to enable direct SIMD dispatch for standard ops (add/sub/mul/div).
30/// Pass `None` for custom operations that have no SIMD fast path.
31pub 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    // Broadcast tensors to the same shape if needed
45    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/// Fallback when SIMD is disabled.
68#[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
81/// Binary operation with in-place optimization for Pod types.
82pub 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    // In-place fast path: lhs unique, contiguous at offset 0, rhs contiguous
90    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    // Allocating path
105    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        // Both contiguous (but lhs not at offset 0)
114        (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        // Fast path for 2D non-contiguous (common for transpose)
124        _ if lhs.layout().num_dims() == 2 => {
125            apply_2d_strided(lhs_storage, rhs_storage, lhs.layout(), rhs.layout(), op)
126        }
127        // General fallback
128        _ => {
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/// Fast 2D strided binary operation using row-based iteration.
142#[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
174/// Apply a scalar operation to each element of a tensor.
175///
176/// Attempts in-place mutation when tensor is contiguous at offset 0.
177pub 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    // In-place fast path: unique, contiguous at offset 0
216    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    // Allocating path
227    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
241/// Helper to construct a tensor from result data.
242fn 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
252/// Apply a binary operation element-wise to two integer tensors.
253///
254/// Supports all integer dtypes: I64, I32, I16, I8, U64, U32, U16, U8.
255pub 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    // Broadcast tensors to the same shape if needed
262    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        // u64 values > i64::MAX wrap to negative i64. This is correct for
272        // add/sub/mul/bitwise (two's complement). Div/rem are handled at the call site.
273        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
281/// Apply a scalar operation to each element of an integer tensor.
282/// Note: scalar is truncated to target dtype (matches PyTorch).
283pub 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// Tests kept here exercise flex-specific behavior of `binary_op` /
317// `scalar_op`: non-contiguous (transposed/narrowed/permuted) strides,
318// flex f16/bf16 half-precision storage paths, and broadcast patterns
319// that probe the flex layout system. Plain contiguous add/sub/mul/div
320// and scalar-op smoke tests have been dropped in favor of the
321// equivalent coverage in ruda-backend-tests, which exercises every
322// backend. When adding new tests, keep them here only if they probe
323// flex-internal dispatch; otherwise add them to
324// crates/ruda-backend-tests/tests/tensor/float/ops/.
325#[cfg(test)]
326mod tests;
327
328pub mod dispatch_float;
329pub mod dispatch_int;
330mod integer_power;