Skip to main content

polars_core/frame/group_by/aggregations/
boolean.rs

1use arrow::bitmap::bitmask::BitMask;
2
3use super::*;
4use crate::chunked_array::cast::CastOptions;
5use crate::chunked_array::from_iterator_par::{collect_bool_opt_par, collect_bool_par};
6use crate::chunked_array::{arg_max_bool, arg_min_bool};
7
8pub fn _agg_helper_idx_bool<F>(groups: &GroupsIdx, f: F) -> Series
9where
10    F: Fn((IdxSize, &IdxVec)) -> Option<bool> + Send + Sync,
11{
12    let ca: BooleanChunked = RAYON.install(|| {
13        let groups_len = groups.len();
14        let first = groups.first();
15        let all = groups.all();
16        collect_bool_opt_par(groups_len, |g| f((first[g], &all[g])))
17    });
18    ca.into_series()
19}
20
21pub fn _agg_helper_slice_bool<F>(groups: &[[IdxSize; 2]], f: F) -> Series
22where
23    F: Fn([IdxSize; 2]) -> Option<bool> + Send + Sync,
24{
25    let ca: BooleanChunked = RAYON.install(|| collect_bool_opt_par(groups.len(), |g| f(groups[g])));
26    ca.into_series()
27}
28
29#[cfg(feature = "bitwise")]
30impl BooleanChunked {
31    pub(crate) unsafe fn agg_and(&self, groups: &GroupsType) -> BooleanChunked {
32        self.agg_all(groups, true)
33    }
34
35    pub(crate) unsafe fn agg_or(&self, groups: &GroupsType) -> BooleanChunked {
36        self.agg_any(groups, true)
37    }
38
39    pub(crate) unsafe fn agg_xor(&self, groups: &GroupsType) -> BooleanChunked {
40        self.bool_agg(
41            groups,
42            true,
43            |values, idxs| {
44                idxs.iter()
45                    .map(|i| {
46                        <IdxSize as From<bool>>::from(unsafe {
47                            values.get_bit_unchecked(*i as usize)
48                        })
49                    })
50                    .sum::<IdxSize>()
51                    % 2
52                    == 1
53            },
54            |values, validity, idxs| {
55                idxs.iter()
56                    .map(|i| {
57                        <IdxSize as From<bool>>::from(unsafe {
58                            validity.get_bit_unchecked(*i as usize)
59                                & values.get_bit_unchecked(*i as usize)
60                        })
61                    })
62                    .sum::<IdxSize>()
63                    % 2
64                    == 1
65            },
66            |_, _, _| unreachable!(),
67            |values, start, length| {
68                unsafe { values.sliced_unchecked(start as usize, length as usize) }.set_bits() % 2
69                    == 1
70            },
71            |values, validity, start, length| {
72                let values = unsafe { values.sliced_unchecked(start as usize, length as usize) };
73                let validity =
74                    unsafe { validity.sliced_unchecked(start as usize, length as usize) };
75                values.num_intersections_with(validity) % 2 == 1
76            },
77            |_, _, _, _| unreachable!(),
78        )
79    }
80}
81
82impl BooleanChunked {
83    pub(crate) unsafe fn agg_min(&self, groups: &GroupsType) -> Series {
84        // faster paths
85        if !self.has_nulls() || matches!(groups, GroupsType::Slice { .. }) {
86            match self.is_sorted_flag() {
87                IsSorted::Ascending => {
88                    return self.clone().into_series().agg_first_non_null(groups);
89                },
90                IsSorted::Descending => {
91                    return self.clone().into_series().agg_last_non_null(groups);
92                },
93                _ => {},
94            }
95        }
96        let ca_self = self.rechunk();
97        let arr = ca_self.downcast_iter().next().unwrap();
98        let no_nulls = arr.null_count() == 0;
99        match groups {
100            GroupsType::Idx(groups) => _agg_helper_idx_bool(groups, |(first, idx)| {
101                debug_assert!(idx.len() <= self.len());
102                if idx.is_empty() {
103                    None
104                } else if idx.len() == 1 {
105                    arr.get(first as usize)
106                } else if no_nulls {
107                    take_arg_min_bool_iter_unchecked_no_nulls(arr, idx2usize(idx))
108                        .map(|p| arr.value_unchecked(idx[p] as usize))
109                } else {
110                    take_arg_min_bool_iter_unchecked_nulls(arr, idx2usize(idx))
111                        .map(|p| arr.value_unchecked(idx[p] as usize))
112                }
113            }),
114            GroupsType::Slice {
115                groups: groups_slice,
116                ..
117            } => _agg_helper_slice_bool(groups_slice, |[first, len]| {
118                debug_assert!(len <= self.len() as IdxSize);
119                match len {
120                    0 => None,
121                    1 => self.get(first as usize),
122                    _ => {
123                        let arr_group = _slice_from_offsets(self, first, len);
124                        arr_group.min()
125                    },
126                }
127            }),
128        }
129    }
130    pub(crate) unsafe fn agg_max(&self, groups: &GroupsType) -> Series {
131        // faster paths
132        if !self.has_nulls() || matches!(groups, GroupsType::Slice { .. }) {
133            match self.is_sorted_flag() {
134                IsSorted::Ascending => return self.clone().into_series().agg_last_non_null(groups),
135                IsSorted::Descending => {
136                    return self.clone().into_series().agg_first_non_null(groups);
137                },
138                _ => {},
139            }
140        }
141
142        let ca_self = self.rechunk();
143        let arr = ca_self.downcast_iter().next().unwrap();
144        let no_nulls = arr.null_count() == 0;
145        match groups {
146            GroupsType::Idx(groups) => _agg_helper_idx_bool(groups, |(first, idx)| {
147                debug_assert!(idx.len() <= self.len());
148                if idx.is_empty() {
149                    None
150                } else if idx.len() == 1 {
151                    self.get(first as usize)
152                } else if no_nulls {
153                    take_arg_max_bool_iter_unchecked_no_nulls(arr, idx2usize(idx))
154                        .map(|p| arr.value_unchecked(idx[p] as usize))
155                } else {
156                    take_arg_max_bool_iter_unchecked_nulls(arr, idx2usize(idx))
157                        .map(|p| arr.value_unchecked(idx[p] as usize))
158                }
159            }),
160            GroupsType::Slice {
161                groups: groups_slice,
162                ..
163            } => _agg_helper_slice_bool(groups_slice, |[first, len]| {
164                debug_assert!(len <= self.len() as IdxSize);
165                match len {
166                    0 => None,
167                    1 => self.get(first as usize),
168                    _ => {
169                        let arr_group = _slice_from_offsets(self, first, len);
170                        arr_group.max()
171                    },
172                }
173            }),
174        }
175    }
176
177    pub(crate) unsafe fn agg_arg_min(&self, groups: &GroupsType) -> Series {
178        // faster paths
179        if !self.has_nulls() || matches!(groups, GroupsType::Slice { .. }) {
180            match self.is_sorted_flag() {
181                IsSorted::Ascending => {
182                    return self.clone().into_series().agg_arg_first_non_null(groups);
183                },
184                IsSorted::Descending => {
185                    return self.clone().into_series().agg_arg_last_non_null(groups);
186                },
187                _ => {},
188            }
189        }
190
191        let ca_self = self.rechunk();
192        let arr = ca_self.downcast_iter().next().unwrap();
193        let no_nulls = arr.null_count() == 0;
194        match groups {
195            GroupsType::Idx(groups) => agg_helper_idx_on_all::<IdxType, _>(groups, |idx| {
196                debug_assert!(idx.len() <= ca_self.len());
197                if idx.is_empty() {
198                    None
199                } else if idx.len() == 1 {
200                    arr.get(idx[0] as usize).map(|_| 0)
201                } else if no_nulls {
202                    take_arg_min_bool_iter_unchecked_no_nulls(arr, idx2usize(idx))
203                        .map(|p| p as IdxSize)
204                } else {
205                    take_arg_min_bool_iter_unchecked_nulls(arr, idx2usize(idx))
206                        .map(|p| p as IdxSize)
207                }
208            }),
209            GroupsType::Slice {
210                groups: groups_slice,
211                ..
212            } => _agg_helper_slice::<IdxType, _>(groups_slice, |[first, len]| {
213                debug_assert!(len <= self.len() as IdxSize);
214                match len {
215                    0 => None,
216                    1 => self.get(first as usize).map(|_| 0),
217                    _ => {
218                        let group_ca = _slice_from_offsets(self, first, len);
219                        arg_min_bool(&group_ca).map(|p| p as IdxSize)
220                    },
221                }
222            }),
223        }
224    }
225
226    pub(crate) unsafe fn agg_arg_max(&self, groups: &GroupsType) -> Series {
227        // faster paths
228        if !self.has_nulls() || matches!(groups, GroupsType::Slice { .. }) {
229            match self.is_sorted_flag() {
230                IsSorted::Ascending => {
231                    return self.clone().into_series().agg_arg_last_non_null(groups);
232                },
233                IsSorted::Descending => {
234                    return self.clone().into_series().agg_arg_first_non_null(groups);
235                },
236                _ => {},
237            }
238        }
239
240        let ca_self = self.rechunk();
241        let arr = ca_self.downcast_iter().next().unwrap();
242        let no_nulls = arr.null_count() == 0;
243        match groups {
244            GroupsType::Idx(groups) => agg_helper_idx_on_all::<IdxType, _>(groups, |idx| {
245                debug_assert!(idx.len() <= ca_self.len());
246                if idx.is_empty() {
247                    None
248                } else if idx.len() == 1 {
249                    arr.get(idx[0] as usize).map(|_| 0)
250                } else if no_nulls {
251                    take_arg_max_bool_iter_unchecked_no_nulls(arr, idx2usize(idx))
252                        .map(|p| p as IdxSize)
253                } else {
254                    take_arg_max_bool_iter_unchecked_nulls(arr, idx2usize(idx))
255                        .map(|p| p as IdxSize)
256                }
257            }),
258            GroupsType::Slice {
259                groups: groups_slice,
260                ..
261            } => _agg_helper_slice::<IdxType, _>(groups_slice, |[first, len]| {
262                debug_assert!(len <= self.len() as IdxSize);
263                match len {
264                    0 => None,
265                    1 => self.get(first as usize).map(|_| 0),
266                    _ => {
267                        let group_ca = _slice_from_offsets(self, first, len);
268                        arg_max_bool(&group_ca).map(|p| p as IdxSize)
269                    },
270                }
271            }),
272        }
273    }
274
275    pub(crate) unsafe fn agg_sum(&self, groups: &GroupsType) -> Series {
276        self.cast_with_options(&IDX_DTYPE, CastOptions::Overflowing)
277            .unwrap()
278            .agg_sum(groups)
279    }
280
281    /// # Safety
282    ///
283    /// Groups should be in correct.
284    #[expect(clippy::too_many_arguments)]
285    unsafe fn bool_agg(
286        &self,
287        groups: &GroupsType,
288        ignore_nulls: bool,
289
290        idx_no_valid: impl Fn(BitMask, &[IdxSize]) -> bool + Send + Sync,
291        idx_validity: impl Fn(BitMask, BitMask, &[IdxSize]) -> bool + Send + Sync,
292        idx_kleene: impl Fn(BitMask, BitMask, &[IdxSize]) -> Option<bool> + Send + Sync,
293
294        slice_no_valid: impl Fn(BitMask, IdxSize, IdxSize) -> bool + Send + Sync,
295        slice_validity: impl Fn(BitMask, BitMask, IdxSize, IdxSize) -> bool + Send + Sync,
296        slice_kleene: impl Fn(BitMask, BitMask, IdxSize, IdxSize) -> Option<bool> + Send + Sync,
297    ) -> BooleanChunked {
298        let name = self.name().clone();
299        let values = self.rechunk();
300        let values = values.downcast_as_array();
301
302        let groups_len = groups.len();
303
304        let ca = RAYON.install(|| {
305            let validity = values
306                .validity()
307                .filter(|v| v.unset_bits() > 0)
308                .map(BitMask::from_bitmap);
309            let values = BitMask::from_bitmap(values.values());
310
311            if !ignore_nulls && let Some(validity) = validity {
312                match groups {
313                    GroupsType::Idx(idx) => {
314                        let all = idx.all();
315                        collect_bool_opt_par(groups_len, |g| idx_kleene(values, validity, &all[g]))
316                    },
317                    GroupsType::Slice { groups, .. } => collect_bool_opt_par(groups_len, |g| {
318                        let [s, l] = groups[g];
319                        slice_kleene(values, validity, s, l)
320                    }),
321                }
322            } else {
323                match groups {
324                    GroupsType::Idx(idx) => {
325                        let all = idx.all();
326                        match validity {
327                            None => collect_bool_par(groups_len, |g| idx_no_valid(values, &all[g])),
328                            Some(validity) => collect_bool_par(groups_len, |g| {
329                                idx_validity(values, validity, &all[g])
330                            }),
331                        }
332                    },
333                    GroupsType::Slice { groups, .. } => match validity {
334                        None => collect_bool_par(groups_len, |g| {
335                            let [s, l] = groups[g];
336                            slice_no_valid(values, s, l)
337                        }),
338                        Some(validity) => collect_bool_par(groups_len, |g| {
339                            let [s, l] = groups[g];
340                            slice_validity(values, validity, s, l)
341                        }),
342                    },
343                }
344            }
345        });
346        ca.with_name(name)
347    }
348
349    /// # Safety
350    ///
351    /// Groups should be in correct.
352    pub unsafe fn agg_any(&self, groups: &GroupsType, ignore_nulls: bool) -> BooleanChunked {
353        self.bool_agg(
354            groups,
355            ignore_nulls,
356            |values, idxs| {
357                idxs.iter()
358                    .any(|i| unsafe { values.get_bit_unchecked(*i as usize) })
359            },
360            |values, validity, idxs| {
361                idxs.iter().any(|i| unsafe {
362                    validity.get_bit_unchecked(*i as usize) & values.get_bit_unchecked(*i as usize)
363                })
364            },
365            |values, validity, idxs| {
366                let mut saw_null = false;
367                for i in idxs.iter() {
368                    let is_valid = unsafe { validity.get_bit_unchecked(*i as usize) };
369                    let is_true = unsafe { values.get_bit_unchecked(*i as usize) };
370
371                    if is_valid & is_true {
372                        return Some(true);
373                    }
374                    saw_null |= !is_valid;
375                }
376                (!saw_null).then_some(false)
377            },
378            |values, start, length| {
379                unsafe { values.sliced_unchecked(start as usize, length as usize) }.leading_zeros()
380                    < length as usize
381            },
382            |values, validity, start, length| {
383                let values = unsafe { values.sliced_unchecked(start as usize, length as usize) };
384                let validity =
385                    unsafe { validity.sliced_unchecked(start as usize, length as usize) };
386                values.intersects_with(validity)
387            },
388            |values, validity, start, length| {
389                let values = unsafe { values.sliced_unchecked(start as usize, length as usize) };
390                let validity =
391                    unsafe { validity.sliced_unchecked(start as usize, length as usize) };
392
393                if values.intersects_with(validity) {
394                    Some(true)
395                } else if validity.unset_bits() == 0 {
396                    Some(false)
397                } else {
398                    None
399                }
400            },
401        )
402    }
403
404    /// # Safety
405    ///
406    /// Groups should be in correct.
407    pub unsafe fn agg_all(&self, groups: &GroupsType, ignore_nulls: bool) -> BooleanChunked {
408        self.bool_agg(
409            groups,
410            ignore_nulls,
411            |values, idxs| {
412                idxs.iter()
413                    .all(|i| unsafe { values.get_bit_unchecked(*i as usize) })
414            },
415            |values, validity, idxs| {
416                idxs.iter().all(|i| unsafe {
417                    !validity.get_bit_unchecked(*i as usize) | values.get_bit_unchecked(*i as usize)
418                })
419            },
420            |values, validity, idxs| {
421                let mut saw_null = false;
422                for i in idxs.iter() {
423                    let is_valid = unsafe { validity.get_bit_unchecked(*i as usize) };
424                    let is_true = unsafe { values.get_bit_unchecked(*i as usize) };
425
426                    if is_valid & !is_true {
427                        return Some(false);
428                    }
429                    saw_null |= !is_valid;
430                }
431                (!saw_null).then_some(true)
432            },
433            |values, start, length| {
434                let values = unsafe { values.sliced_unchecked(start as usize, length as usize) };
435                values.unset_bits() == 0
436            },
437            |values, validity, start, length| {
438                let values = unsafe { values.sliced_unchecked(start as usize, length as usize) };
439                let validity =
440                    unsafe { validity.sliced_unchecked(start as usize, length as usize) };
441                values.num_intersections_with(validity) == validity.set_bits()
442            },
443            |values, validity, start, length| {
444                let values = unsafe { values.sliced_unchecked(start as usize, length as usize) };
445                let validity =
446                    unsafe { validity.sliced_unchecked(start as usize, length as usize) };
447
448                let num_non_nulls = validity.set_bits();
449
450                if values.num_intersections_with(validity) < num_non_nulls {
451                    Some(false)
452                } else if num_non_nulls < values.len() {
453                    None
454                } else {
455                    Some(true)
456                }
457            },
458        )
459    }
460}