hurray_core/layout/addressing/
coo.rs1use std::cmp::Ordering;
10
11use crate::{Error, Result};
12
13pub fn element_offset(query: &[u64], is_sorted: bool, indices: &[u64]) -> Result<Option<u64>> {
49 let rank = query.len();
50 if rank == 0 {
51 return Err(Error::IndexRankMismatch {
52 index_rank: 0,
53 shape_rank: 0,
54 });
55 }
56 if !indices.len().is_multiple_of(rank) {
57 return Err(Error::InvalidLayout(format!(
58 "coo: indices.len()={} is not a multiple of rank={rank}",
59 indices.len()
60 )));
61 }
62 let nnz = indices.len() / rank;
63 let coords = |r: usize| &indices[r * rank..r * rank + rank];
64
65 if is_sorted {
66 let (mut lo, mut hi) = (0usize, nnz);
69 while lo < hi {
70 let mid = lo + (hi - lo) / 2;
71 match coords(mid).cmp(query) {
72 Ordering::Less => lo = mid + 1,
73 Ordering::Greater => hi = mid,
74 Ordering::Equal => return Ok(Some(mid as u64)),
75 }
76 }
77 Ok(None)
78 } else {
79 for r in 0..nnz {
80 if coords(r) == query {
81 return Ok(Some(r as u64));
82 }
83 }
84 Ok(None)
85 }
86}
87
88#[cfg(test)]
89mod tests {
90 use super::*;
91
92 const INDICES: &[u64] = &[0, 1, 2, 0, 2, 3, 3, 3];
94
95 #[test]
96 fn sorted_lookup_hits() {
97 assert_eq!(element_offset(&[0, 1], true, INDICES).unwrap(), Some(0));
98 assert_eq!(element_offset(&[2, 0], true, INDICES).unwrap(), Some(1));
99 assert_eq!(element_offset(&[2, 3], true, INDICES).unwrap(), Some(2));
100 assert_eq!(element_offset(&[3, 3], true, INDICES).unwrap(), Some(3));
101 }
102
103 #[test]
104 fn sorted_lookup_structural_zeros() {
105 assert_eq!(element_offset(&[0, 0], true, INDICES).unwrap(), None);
106 assert_eq!(element_offset(&[1, 1], true, INDICES).unwrap(), None);
107 assert_eq!(element_offset(&[2, 2], true, INDICES).unwrap(), None);
108 assert_eq!(element_offset(&[3, 0], true, INDICES).unwrap(), None);
110 }
111
112 #[test]
113 fn unsorted_lookup_matches_sorted() {
114 let unsorted: &[u64] = &[2, 3, 0, 1, 3, 3, 2, 0];
116 assert_eq!(element_offset(&[0, 1], false, unsorted).unwrap(), Some(1));
117 assert_eq!(element_offset(&[2, 3], false, unsorted).unwrap(), Some(0));
118 assert_eq!(element_offset(&[3, 3], false, unsorted).unwrap(), Some(2));
119 assert_eq!(element_offset(&[1, 1], false, unsorted).unwrap(), None);
120 }
121
122 #[test]
123 fn rank_3_lookup() {
124 let indices: &[u64] = &[0, 0, 1, 0, 2, 3, 1, 1, 0];
126 assert_eq!(element_offset(&[0, 2, 3], true, indices).unwrap(), Some(1));
127 assert_eq!(element_offset(&[1, 1, 0], true, indices).unwrap(), Some(2));
128 assert_eq!(element_offset(&[1, 1, 2], true, indices).unwrap(), None);
129 }
130
131 #[test]
132 fn empty_tensor_is_all_zeros() {
133 assert_eq!(element_offset(&[0, 0], true, &[]).unwrap(), None);
134 assert_eq!(element_offset(&[0, 0], false, &[]).unwrap(), None);
135 }
136
137 #[test]
138 fn rank_zero_query_is_rejected() {
139 assert!(matches!(
140 element_offset(&[], true, &[]),
141 Err(Error::IndexRankMismatch { .. })
142 ));
143 }
144
145 #[test]
146 fn indices_not_multiple_of_rank_is_rejected() {
147 assert!(matches!(
149 element_offset(&[0, 0], true, &[0, 1, 2]),
150 Err(Error::InvalidLayout(_))
151 ));
152 }
153}