1use alloc::boxed::Box;
4#[cfg(feature = "simd")]
5use alloc::vec;
6use alloc::vec::Vec;
7use ruda_core::tensor::{DType, element::Element};
8use ruda_core::{bytes::Bytes, tensor::{BoolDType, BoolStore, Shape}};
9use half::{bf16, f16};
10use bytemuck::Pod;
11
12use ruda_core::tensor::host::strided_index::StridedIter;
13use ruda_core::tensor::host::{HostTensor, Layout};
14
15use crate::simd;
16
17pub use simd::CmpOp as CompareOp;
19
20pub fn compare<F32Cmp, F64Cmp>(
23 lhs: HostTensor,
24 rhs: HostTensor,
25 out_dtype: BoolDType,
26 f32_cmp: F32Cmp,
27 f64_cmp: F64Cmp,
28 simd_hint: Option<CompareOp>,
29) -> HostTensor
30where
31 F32Cmp: Fn(f32, f32) -> bool + Copy,
32 F64Cmp: Fn(f64, f64) -> bool + Copy,
33{
34 debug_assert_eq!(lhs.dtype(), rhs.dtype(), "compare: dtype mismatch");
35
36 let (lhs, rhs) = crate::expand::broadcast_binary(lhs, rhs);
38
39 let dtype = lhs.dtype();
40
41 match dtype {
42 DType::F32 => compare_f32(lhs, &rhs, out_dtype, f32_cmp, simd_hint),
43 DType::F64 => compare_typed(lhs, &rhs, out_dtype, f64_cmp),
44 DType::F16 => compare_typed(lhs, &rhs, out_dtype, |a: f16, b: f16| {
45 f32_cmp(a.to_f32(), b.to_f32())
46 }),
47 DType::BF16 => compare_typed(lhs, &rhs, out_dtype, |a: bf16, b: bf16| {
48 f32_cmp(a.to_f32(), b.to_f32())
49 }),
50 _ => panic!("compare: unsupported dtype {:?}", dtype),
51 }
52}
53
54#[cfg(feature = "simd")]
56fn compare_f32<Cmp>(
57 lhs: HostTensor,
58 rhs: &HostTensor,
59 out_dtype: BoolDType,
60 cmp: Cmp,
61 simd_hint: Option<CompareOp>,
62) -> HostTensor
63where
64 Cmp: Fn(f32, f32) -> bool,
65{
66 if let (Some((l_start, l_end)), Some((r_start, r_end))) = (
68 lhs.layout().contiguous_offsets(),
69 rhs.layout().contiguous_offsets(),
70 ) && let Some(simd_op) = simd_hint
71 {
72 let shape = lhs.layout().shape().clone();
73 let lhs_storage: &[f32] = lhs.storage();
74 let rhs_storage: &[f32] = rhs.storage();
75
76 let l_slice = &lhs_storage[l_start..l_end];
77 let r_slice = &rhs_storage[r_start..r_end];
78
79 let mut result = vec![0u8; l_slice.len()];
80 simd::cmp_f32(l_slice, r_slice, &mut result, simd_op);
81
82 return make_bool_tensor(result, shape, out_dtype);
83 }
84
85 if lhs.layout().num_dims() == 2
88 && let Some(simd_op) = simd_hint
89 && let Some((result, shape)) = try_broadcast_cmp_f32(&lhs, rhs, simd_op)
90 {
91 return make_bool_tensor(result, shape, out_dtype);
92 }
93
94 compare_typed(lhs, rhs, out_dtype, cmp)
96}
97
98#[cfg(feature = "simd")]
101fn try_broadcast_cmp_f32(
102 lhs: &HostTensor,
103 rhs: &HostTensor,
104 op: simd::CmpOp,
105) -> Option<(Vec<u8>, Shape)> {
106 let lhs_strides = lhs.layout().strides();
107 let rhs_strides = rhs.layout().strides();
108 let shape = lhs.layout().shape().clone();
109 let [rows, cols] = shape[..] else {
110 return None;
111 };
112
113 if lhs_strides[1] == 0 && rhs_strides == [cols as isize, 1] {
116 let lhs_storage: &[f32] = lhs.storage();
117 let rhs_storage: &[f32] = rhs.storage();
118 let l_offset = lhs.layout().start_offset() as isize;
119 let l_stride = lhs_strides[0];
120 let r_offset = rhs.layout().start_offset();
121
122 let mut result = vec![0u8; rows * cols];
123 for row in 0..rows {
124 let a_val = lhs_storage[(l_offset + row as isize * l_stride) as usize];
125 let r_row_start = r_offset + row * cols;
126 let r_slice = &rhs_storage[r_row_start..r_row_start + cols];
127 let out_start = row * cols;
128 simd::cmp_scalar_f32(
129 r_slice,
130 a_val,
131 &mut result[out_start..out_start + cols],
132 swap_cmp_op(op),
133 );
134 }
135 return Some((result, shape));
136 }
137
138 if rhs_strides[0] == 0 && lhs_strides == [cols as isize, 1] {
141 let lhs_storage: &[f32] = lhs.storage();
142 let rhs_storage: &[f32] = rhs.storage();
143 let l_offset = lhs.layout().start_offset();
144 let r_offset = rhs.layout().start_offset() as isize;
145 let r_stride = rhs_strides[1];
146
147 let rhs_row: Vec<f32> = (0..cols)
149 .map(|j| rhs_storage[(r_offset + j as isize * r_stride) as usize])
150 .collect();
151
152 let mut result = vec![0u8; rows * cols];
153 for row in 0..rows {
154 let l_row_start = l_offset + row * cols;
155 let l_slice = &lhs_storage[l_row_start..l_row_start + cols];
156 let out_start = row * cols;
157 for (j, (&lv, &rv)) in l_slice.iter().zip(rhs_row.iter()).enumerate() {
159 result[out_start + j] = match op {
160 simd::CmpOp::Gt => (lv > rv) as u8,
161 simd::CmpOp::Ge => (lv >= rv) as u8,
162 simd::CmpOp::Lt => (lv < rv) as u8,
163 simd::CmpOp::Le => (lv <= rv) as u8,
164 simd::CmpOp::Eq => (lv == rv) as u8,
165 simd::CmpOp::Ne => (lv != rv) as u8,
166 };
167 }
168 }
169 return Some((result, shape));
170 }
171
172 if lhs_strides[1] == 0 && rhs_strides[0] == 0 {
175 let lhs_storage: &[f32] = lhs.storage();
176 let rhs_storage: &[f32] = rhs.storage();
177 let l_offset = lhs.layout().start_offset() as isize;
178 let l_stride = lhs_strides[0];
179 let r_offset = rhs.layout().start_offset() as isize;
180 let r_stride = rhs_strides[1];
181
182 let rhs_row: Vec<f32> = (0..cols)
184 .map(|j| rhs_storage[(r_offset + j as isize * r_stride) as usize])
185 .collect();
186
187 let mut result = vec![0u8; rows * cols];
188 for row in 0..rows {
189 let a_val = lhs_storage[(l_offset + row as isize * l_stride) as usize];
190 let out_start = row * cols;
191 simd::cmp_scalar_f32(
192 &rhs_row,
193 a_val,
194 &mut result[out_start..out_start + cols],
195 swap_cmp_op(op),
196 );
197 }
198 return Some((result, shape));
199 }
200
201 None
202}
203
204#[cfg(feature = "simd")]
206fn swap_cmp_op(op: simd::CmpOp) -> simd::CmpOp {
207 match op {
208 simd::CmpOp::Gt => simd::CmpOp::Lt, simd::CmpOp::Ge => simd::CmpOp::Le,
210 simd::CmpOp::Lt => simd::CmpOp::Gt,
211 simd::CmpOp::Le => simd::CmpOp::Ge,
212 simd::CmpOp::Eq => simd::CmpOp::Eq, simd::CmpOp::Ne => simd::CmpOp::Ne,
214 }
215}
216
217#[cfg(not(feature = "simd"))]
219fn compare_f32<Cmp>(
220 lhs: HostTensor,
221 rhs: &HostTensor,
222 out_dtype: BoolDType,
223 cmp: Cmp,
224 _simd_hint: Option<CompareOp>,
225) -> HostTensor
226where
227 Cmp: Fn(f32, f32) -> bool,
228{
229 compare_typed(lhs, rhs, out_dtype, cmp)
230}
231
232pub fn compare_elem<F32Cmp, F64Cmp>(
235 lhs: HostTensor,
236 rhs: f64,
237 out_dtype: BoolDType,
238 f32_cmp: F32Cmp,
239 f64_cmp: F64Cmp,
240 simd_hint: Option<CompareOp>,
241) -> HostTensor
242where
243 F32Cmp: Fn(f32, f32) -> bool + Copy,
244 F64Cmp: Fn(f64, f64) -> bool + Copy,
245{
246 let dtype = lhs.dtype();
247
248 match dtype {
249 DType::F32 => compare_elem_f32(lhs, rhs as f32, out_dtype, f32_cmp, simd_hint),
250 DType::F64 => compare_elem_typed(lhs, rhs, out_dtype, f64_cmp),
251 DType::F16 => {
252 let scalar = f16::from_f64(rhs);
253 compare_elem_typed(lhs, scalar, out_dtype, |a: f16, b: f16| {
254 f32_cmp(a.to_f32(), b.to_f32())
255 })
256 }
257 DType::BF16 => {
258 let scalar = bf16::from_f64(rhs);
259 compare_elem_typed(lhs, scalar, out_dtype, |a: bf16, b: bf16| {
260 f32_cmp(a.to_f32(), b.to_f32())
261 })
262 }
263 _ => panic!("compare_elem: unsupported dtype {:?}", dtype),
264 }
265}
266
267#[cfg(feature = "simd")]
269fn compare_elem_f32<Cmp>(
270 lhs: HostTensor,
271 rhs: f32,
272 out_dtype: BoolDType,
273 cmp: Cmp,
274 simd_hint: Option<CompareOp>,
275) -> HostTensor
276where
277 Cmp: Fn(f32, f32) -> bool,
278{
279 if let Some((start, end)) = lhs.layout().contiguous_offsets()
281 && let Some(simd_op) = simd_hint
282 {
283 let shape = lhs.layout().shape().clone();
284 let lhs_storage: &[f32] = lhs.storage();
285 let l_slice = &lhs_storage[start..end];
286
287 let mut result = vec![0u8; l_slice.len()];
288 simd::cmp_scalar_f32(l_slice, rhs, &mut result, simd_op);
289
290 return make_bool_tensor(result, shape, out_dtype);
291 }
292
293 compare_elem_typed(lhs, rhs, out_dtype, cmp)
295}
296
297#[cfg(not(feature = "simd"))]
299fn compare_elem_f32<Cmp>(
300 lhs: HostTensor,
301 rhs: f32,
302 out_dtype: BoolDType,
303 cmp: Cmp,
304 _simd_hint: Option<CompareOp>,
305) -> HostTensor
306where
307 Cmp: Fn(f32, f32) -> bool,
308{
309 compare_elem_typed(lhs, rhs, out_dtype, cmp)
310}
311
312fn compare_typed<E, Cmp>(
313 lhs: HostTensor,
314 rhs: &HostTensor,
315 out_dtype: BoolDType,
316 cmp: Cmp,
317) -> HostTensor
318where
319 E: Element + Pod,
320 Cmp: Fn(E, E) -> bool,
321{
322 let shape = lhs.layout().shape().clone();
323 let lhs_storage: &[E] = lhs.storage();
324 let rhs_storage: &[E] = rhs.storage();
325
326 let result: Vec<u8> = match (
327 lhs.layout().contiguous_offsets(),
328 rhs.layout().contiguous_offsets(),
329 ) {
330 (Some((l_start, l_end)), Some((r_start, r_end))) => {
331 let l_slice = &lhs_storage[l_start..l_end];
332 let r_slice = &rhs_storage[r_start..r_end];
333 l_slice
334 .iter()
335 .zip(r_slice)
336 .map(|(&a, &b)| cmp(a, b) as u8)
337 .collect()
338 }
339 _ if lhs.layout().num_dims() == 2 => crate::binary::apply_2d_strided(
341 lhs_storage,
342 rhs_storage,
343 lhs.layout(),
344 rhs.layout(),
345 |a, b| cmp(a, b) as u8,
346 ),
347 _ => {
348 let lhs_iter = StridedIter::new(lhs.layout());
349 let rhs_iter = StridedIter::new(rhs.layout());
350 lhs_iter
351 .zip(rhs_iter)
352 .map(|(li, ri)| cmp(lhs_storage[li], rhs_storage[ri]) as u8)
353 .collect()
354 }
355 };
356
357 make_bool_tensor(result, shape, out_dtype)
358}
359
360fn compare_elem_typed<E, Cmp>(lhs: HostTensor, rhs: E, out_dtype: BoolDType, cmp: Cmp) -> HostTensor
361where
362 E: Element + Pod + Copy,
363 Cmp: Fn(E, E) -> bool,
364{
365 let shape = lhs.layout().shape().clone();
366 let lhs_storage: &[E] = lhs.storage();
367
368 let result: Vec<u8> = match lhs.layout().contiguous_offsets() {
369 Some((start, end)) => lhs_storage[start..end]
370 .iter()
371 .map(|&a| cmp(a, rhs) as u8)
372 .collect(),
373 None => StridedIter::new(lhs.layout())
374 .map(|idx| cmp(lhs_storage[idx], rhs) as u8)
375 .collect(),
376 };
377
378 make_bool_tensor(result, shape, out_dtype)
379}
380
381pub fn make_bool_tensor(data: Vec<u8>, shape: Shape, out_dtype: BoolDType) -> HostTensor {
389 let store = match out_dtype {
390 BoolDType::Native => BoolStore::Native,
391 BoolDType::U8 => BoolStore::U8,
392 BoolDType::U32 => panic!(
393 "ruda-tensor-host does not support Bool(U32) storage (only Native and U8). \
394 Use a backend that declares Bool(U32) support, or work with Bool(Native)/Bool(U8)."
395 ),
396 };
397 let bytes = Bytes::from_elems(data);
398 HostTensor::new(bytes, Layout::contiguous(shape), DType::Bool(store))
399}
400
401pub fn greater(lhs: HostTensor, rhs: HostTensor, out_dtype: BoolDType) -> HostTensor {
404 compare(
405 lhs,
406 rhs,
407 out_dtype,
408 |a, b| a > b,
409 |a, b| a > b,
410 Some(CompareOp::Gt),
411 )
412}
413
414pub fn greater_elem(lhs: HostTensor, rhs: f64, out_dtype: BoolDType) -> HostTensor {
415 compare_elem(
416 lhs,
417 rhs,
418 out_dtype,
419 |a, b| a > b,
420 |a, b| a > b,
421 Some(CompareOp::Gt),
422 )
423}
424
425pub fn greater_equal(lhs: HostTensor, rhs: HostTensor, out_dtype: BoolDType) -> HostTensor {
426 compare(
427 lhs,
428 rhs,
429 out_dtype,
430 |a, b| a >= b,
431 |a, b| a >= b,
432 Some(CompareOp::Ge),
433 )
434}
435
436pub fn greater_equal_elem(lhs: HostTensor, rhs: f64, out_dtype: BoolDType) -> HostTensor {
437 compare_elem(
438 lhs,
439 rhs,
440 out_dtype,
441 |a, b| a >= b,
442 |a, b| a >= b,
443 Some(CompareOp::Ge),
444 )
445}
446
447pub fn lower(lhs: HostTensor, rhs: HostTensor, out_dtype: BoolDType) -> HostTensor {
448 compare(
449 lhs,
450 rhs,
451 out_dtype,
452 |a, b| a < b,
453 |a, b| a < b,
454 Some(CompareOp::Lt),
455 )
456}
457
458pub fn lower_elem(lhs: HostTensor, rhs: f64, out_dtype: BoolDType) -> HostTensor {
459 compare_elem(
460 lhs,
461 rhs,
462 out_dtype,
463 |a, b| a < b,
464 |a, b| a < b,
465 Some(CompareOp::Lt),
466 )
467}
468
469pub fn lower_equal(lhs: HostTensor, rhs: HostTensor, out_dtype: BoolDType) -> HostTensor {
470 compare(
471 lhs,
472 rhs,
473 out_dtype,
474 |a, b| a <= b,
475 |a, b| a <= b,
476 Some(CompareOp::Le),
477 )
478}
479
480pub fn lower_equal_elem(lhs: HostTensor, rhs: f64, out_dtype: BoolDType) -> HostTensor {
481 compare_elem(
482 lhs,
483 rhs,
484 out_dtype,
485 |a, b| a <= b,
486 |a, b| a <= b,
487 Some(CompareOp::Le),
488 )
489}
490
491pub fn equal(lhs: HostTensor, rhs: HostTensor, out_dtype: BoolDType) -> HostTensor {
492 compare(
493 lhs,
494 rhs,
495 out_dtype,
496 |a, b| a == b,
497 |a, b| a == b,
498 Some(CompareOp::Eq),
499 )
500}
501
502pub fn equal_elem(lhs: HostTensor, rhs: f64, out_dtype: BoolDType) -> HostTensor {
503 compare_elem(
504 lhs,
505 rhs,
506 out_dtype,
507 |a, b| a == b,
508 |a, b| a == b,
509 Some(CompareOp::Eq),
510 )
511}
512
513pub fn not_equal(lhs: HostTensor, rhs: HostTensor, out_dtype: BoolDType) -> HostTensor {
514 compare(
515 lhs,
516 rhs,
517 out_dtype,
518 |a, b| a != b,
519 |a, b| a != b,
520 Some(CompareOp::Ne),
521 )
522}
523
524pub fn not_equal_elem(lhs: HostTensor, rhs: f64, out_dtype: BoolDType) -> HostTensor {
525 compare_elem(
526 lhs,
527 rhs,
528 out_dtype,
529 |a, b| a != b,
530 |a, b| a != b,
531 Some(CompareOp::Ne),
532 )
533}
534
535mod integer;
536pub use integer::*;
537
538mod predicate_reduce;
539pub use predicate_reduce::*;
540
541fn bool_scalar(val: bool, out_dtype: BoolDType) -> HostTensor {
546 let byte: u8 = if val { 1 } else { 0 };
547 make_bool_tensor(alloc::vec![byte], Shape::from(alloc::vec![1]), out_dtype)
548}
549
550fn iter_elements<'a, E: Element + Pod + 'a>(
551 tensor: &'a HostTensor,
552) -> Box<dyn Iterator<Item = E> + 'a> {
553 let data: &[E] = tensor.storage();
554 match tensor.layout().contiguous_offsets() {
555 Some((start, end)) => Box::new(data[start..end].iter().copied()),
556 None => Box::new(StridedIter::new(tensor.layout()).map(move |idx| data[idx])),
557 }
558}
559
560fn reduce_bool_dim_with(
565 tensor: &HostTensor,
566 dim: usize,
567 init: bool,
568 combine: fn(bool, bool) -> bool,
569 out_dtype: BoolDType,
570 is_nonzero: impl Fn(usize) -> bool,
571) -> HostTensor {
572 debug_assert!(tensor.is_contiguous() && tensor.layout().start_offset() == 0);
573 let shape = tensor.layout().shape();
574 let ndims = shape.num_dims();
575 assert!(dim < ndims);
576
577 let dim_size = shape[dim];
578 let mut out_shape: Vec<usize> = shape.to_vec();
579 out_shape[dim] = 1;
580 let outer_size: usize = shape[..dim].iter().product();
581 let inner_size: usize = shape[dim + 1..].iter().product();
582
583 let out_size = outer_size.max(1) * inner_size.max(1);
584 let mut result: Vec<u8> = Vec::with_capacity(out_size);
585
586 for outer in 0..outer_size.max(1) {
587 for inner in 0..inner_size.max(1) {
588 let mut acc = init;
589 for d in 0..dim_size {
590 let idx = outer * dim_size * inner_size + d * inner_size + inner;
591 acc = combine(acc, is_nonzero(idx));
592 }
593 result.push(if acc { 1 } else { 0 });
594 }
595 }
596
597 make_bool_tensor(result, Shape::from(out_shape), out_dtype)
598}
599
600fn reduce_bool_dim(
602 tensor: &HostTensor,
603 dim: usize,
604 init: bool,
605 combine: fn(bool, bool) -> bool,
606 out_dtype: BoolDType,
607) -> HostTensor {
608 let tensor = tensor.to_contiguous();
609 match tensor.dtype() {
610 DType::F32 => {
611 let data: &[f32] = tensor.storage();
612 reduce_bool_dim_with(&tensor, dim, init, combine, out_dtype, |idx| {
613 data[idx] != 0.0
614 })
615 }
616 DType::F64 => {
617 let data: &[f64] = tensor.storage();
618 reduce_bool_dim_with(&tensor, dim, init, combine, out_dtype, |idx| {
619 data[idx] != 0.0
620 })
621 }
622 DType::F16 => {
623 let data: &[f16] = tensor.storage();
624 reduce_bool_dim_with(&tensor, dim, init, combine, out_dtype, |idx| {
625 data[idx].to_f32() != 0.0
626 })
627 }
628 DType::BF16 => {
629 let data: &[bf16] = tensor.storage();
630 reduce_bool_dim_with(&tensor, dim, init, combine, out_dtype, |idx| {
631 data[idx].to_f32() != 0.0
632 })
633 }
634 _ => panic!("reduce_bool_dim: unsupported dtype {:?}", tensor.dtype()),
635 }
636}
637
638fn reduce_bool_dim_int(
640 tensor: &HostTensor,
641 dim: usize,
642 init: bool,
643 combine: fn(bool, bool) -> bool,
644 out_dtype: BoolDType,
645) -> HostTensor {
646 let tensor = tensor.to_contiguous();
647 macro_rules! dispatch {
648 ($ty:ty) => {{
649 let data: &[$ty] = tensor.storage();
650 reduce_bool_dim_with(&tensor, dim, init, combine, out_dtype, |idx| data[idx] != 0)
651 }};
652 }
653 match tensor.dtype() {
654 DType::I64 => dispatch!(i64),
655 DType::I32 => dispatch!(i32),
656 DType::I16 => dispatch!(i16),
657 DType::I8 => dispatch!(i8),
658 DType::U64 => dispatch!(u64),
659 DType::U32 => dispatch!(u32),
660 DType::U16 => dispatch!(u16),
661 DType::U8 => dispatch!(u8),
662 other => panic!("reduce_bool_dim_int: unsupported dtype {:?}", other),
663 }
664}
665
666fn reduce_bool_dim_raw(
668 tensor: &HostTensor,
669 dim: usize,
670 init: bool,
671 combine: fn(bool, bool) -> bool,
672 out_dtype: BoolDType,
673) -> HostTensor {
674 let tensor = tensor.to_contiguous();
675 let data: &[u8] = tensor.bytes();
676 reduce_bool_dim_with(&tensor, dim, init, combine, out_dtype, |idx| data[idx] != 0)
677}
678
679#[cfg(test)]
685mod tests;
686
687pub mod dispatch;