Skip to main content

ruprim_host/comparison/
mod.rs

1//! Comparison operations returning boolean tensors.
2
3use 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
17/// Comparison operation type for SIMD dispatch.
18pub use simd::CmpOp as CompareOp;
19
20/// Compare two tensors element-wise, returning a boolean tensor with the
21/// requested output dtype.
22pub 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    // Broadcast to same shape if needed
37    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/// Specialized comparison for f32 with SIMD fast path.
55#[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    // SIMD fast path: both tensors contiguous
67    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    // Optimized broadcast path for outer-product style broadcasting
86    // Pattern: [N, 1] vs [1, M] -> [N, M] where one has stride 0 in inner dim
87    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    // Fallback to generic path
95    compare_typed(lhs, rhs, out_dtype, cmp)
96}
97
98/// Try optimized outer-product style broadcast comparison.
99/// Returns Some((result, shape)) if the pattern matches.
100#[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    // Pattern 1: lhs has stride 0 in dim 1 (column broadcast), rhs contiguous
114    // lhs[i,j] = lhs_data[i*stride], rhs[i,j] = rhs_data[i*cols + j]
115    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    // Pattern 2: rhs has stride 0 in dim 0 (row broadcast), lhs contiguous
139    // lhs[i,j] = lhs_data[i*cols + j], rhs[i,j] = rhs_data[j*stride]
140    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        // Build the broadcast rhs values once
148        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            // Compare row with broadcast values
158            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    // Pattern 3: Outer product - lhs stride 0 in dim 1, rhs stride 0 in dim 0
173    // This is the [N,1] vs [1,M] case
174    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        // Build the broadcast rhs row once
183        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/// Swap comparison operation for reversed operand order.
205#[cfg(feature = "simd")]
206fn swap_cmp_op(op: simd::CmpOp) -> simd::CmpOp {
207    match op {
208        simd::CmpOp::Gt => simd::CmpOp::Lt, // a > b becomes b < a
209        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, // symmetric
213        simd::CmpOp::Ne => simd::CmpOp::Ne,
214    }
215}
216
217/// Fallback when SIMD is disabled.
218#[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
232/// Compare tensor with scalar, returning a boolean tensor with the requested
233/// output dtype.
234pub 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/// Specialized scalar comparison for f32 with SIMD fast path.
268#[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    // SIMD fast path: tensor is contiguous
280    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    // Fallback to generic path
294    compare_elem_typed(lhs, rhs, out_dtype, cmp)
295}
296
297/// Fallback when SIMD is disabled.
298#[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        // Fast path for 2D non-contiguous (common for transpose)
340        _ 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
381/// Build a bool `FlexTensor` from a `Vec<u8>` of 0/1 bytes, tagged with the
382/// requested output dtype.
383///
384/// ruda-tensor-host stores bools as 1 byte per element, so only Native and U8 are
385/// supported. `Bool(U32)` would require 4-byte-per-element storage throughout
386/// the backend; `dtype_usage` declares it unsupported and this function panics
387/// if it's requested.
388pub 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
401// Specific comparison functions
402
403pub 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
541// ============================================================================
542// Helpers for any/all
543// ============================================================================
544
545fn 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
560/// Reduce along a dimension producing a bool tensor.
561///
562/// The `is_nonzero` closure reads the data slice at a given index and returns
563/// whether the element is nonzero.
564fn 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
600/// Reduce along a dimension producing a bool tensor (for float any/all_dim).
601fn 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
638/// Reduce along a dimension producing a bool tensor (for int any/all_dim).
639fn 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
666/// Reduce along a dimension producing a bool tensor (for bool any/all_dim).
667fn 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// Tests kept here probe flex-internal `reduce_bool_dim_with` dispatch on
680// non-contiguous inputs (stale-pointer-read regression, see prior incident
681// in `any_float_dim`). Plain comparison ops and stride variants (flipped
682// / transposed / narrowed) have been migrated to ruda-backend-tests at
683// tensor/{float,int}/ops/comparison.rs so every backend is exercised.
684#[cfg(test)]
685mod tests;
686
687pub mod dispatch;