@@ -20,38 +20,91 @@ use arrow::array::{
2020} ;
2121use arrow:: buffer:: OffsetBuffer ;
2222use arrow:: datatypes:: { ArrowPrimitiveType , Field } ;
23- use datafusion_common:: HashSet ;
24- use datafusion_common:: hash_utils:: RandomState ;
2523use datafusion_expr_common:: groups_accumulator:: { EmitTo , GroupsAccumulator } ;
26- use std:: hash:: Hash ;
2724use std:: mem:: size_of;
2825use std:: sync:: Arc ;
2926
3027use crate :: aggregate:: groups_accumulator:: accumulate:: accumulate;
3128
29+ /// Trait for packing (group_idx, value) into a single sortable integer
30+ pub trait Packable : Copy + Send {
31+ type Packed : Ord + Copy + Send ;
32+ fn pack ( group_idx : usize , value : Self ) -> Self :: Packed ;
33+ fn unpack ( packed : Self :: Packed ) -> ( usize , Self ) ;
34+ }
35+
36+ macro_rules! impl_packable_signed {
37+ ( $native: ty, $unsigned: ty, $packed: ty, $bits: expr) => {
38+ impl Packable for $native {
39+ type Packed = $packed;
40+ #[ inline]
41+ fn pack( group_idx: usize , value: Self ) -> $packed {
42+ let val = ( value as $unsigned ^ ( 1 << ( $bits - 1 ) ) ) as $packed;
43+ ( ( group_idx as $packed) << $bits) | val
44+ }
45+ #[ inline]
46+ fn unpack( packed: $packed) -> ( usize , Self ) {
47+ let group = ( packed >> $bits) as usize ;
48+ let val = ( ( packed as $unsigned) ^ ( 1 << ( $bits - 1 ) ) ) as $native;
49+ ( group, val)
50+ }
51+ }
52+ } ;
53+ }
54+
55+ macro_rules! impl_packable_unsigned {
56+ ( $native: ty, $packed: ty, $bits: expr) => {
57+ impl Packable for $native {
58+ type Packed = $packed;
59+ #[ inline]
60+ fn pack( group_idx: usize , value: Self ) -> $packed {
61+ ( ( group_idx as $packed) << $bits) | ( value as $packed)
62+ }
63+ #[ inline]
64+ fn unpack( packed: $packed) -> ( usize , Self ) {
65+ ( ( packed >> $bits) as usize , packed as $native)
66+ }
67+ }
68+ } ;
69+ }
70+
71+ impl_packable_signed ! ( i64 , u64 , u128 , 64 ) ;
72+ impl_packable_signed ! ( i32 , u32 , u64 , 32 ) ;
73+ impl_packable_signed ! ( i16 , u16 , u64 , 16 ) ;
74+ impl_packable_signed ! ( i8 , u8 , u64 , 8 ) ;
75+
76+ impl_packable_unsigned ! ( u64 , u128 , 64 ) ;
77+ impl_packable_unsigned ! ( u32 , u64 , 32 ) ;
78+ impl_packable_unsigned ! ( u16 , u64 , 16 ) ;
79+ impl_packable_unsigned ! ( u8 , u64 , 8 ) ;
80+
81+ /// A `GroupsAccumulator` for COUNT(DISTINCT) on primitive types.
82+ ///
83+ /// Uses a flat buffer with packed integers: (group_idx, value) packed into
84+ /// a single sortable integer. Sort and dedup at evaluate time.
3285pub struct PrimitiveDistinctCountGroupsAccumulator < T : ArrowPrimitiveType >
3386where
34- T :: Native : Eq + Hash ,
87+ T :: Native : Packable ,
3588{
36- seen : HashSet < ( usize , T :: Native ) , RandomState > ,
89+ buffer : Vec < < T :: Native as Packable > :: Packed > ,
3790 num_groups : usize ,
3891}
3992
4093impl < T : ArrowPrimitiveType > PrimitiveDistinctCountGroupsAccumulator < T >
4194where
42- T :: Native : Eq + Hash ,
95+ T :: Native : Packable ,
4396{
4497 pub fn new ( ) -> Self {
4598 Self {
46- seen : HashSet :: default ( ) ,
99+ buffer : Vec :: new ( ) ,
47100 num_groups : 0 ,
48101 }
49102 }
50103}
51104
52105impl < T : ArrowPrimitiveType > Default for PrimitiveDistinctCountGroupsAccumulator < T >
53106where
54- T :: Native : Eq + Hash ,
107+ T :: Native : Packable ,
55108{
56109 fn default ( ) -> Self {
57110 Self :: new ( )
61114impl < T : ArrowPrimitiveType + Send + std:: fmt:: Debug > GroupsAccumulator
62115 for PrimitiveDistinctCountGroupsAccumulator < T >
63116where
64- T :: Native : Eq + Hash ,
117+ T :: Native : Packable ,
65118{
66119 fn update_batch (
67120 & mut self ,
@@ -73,8 +126,11 @@ where
73126 debug_assert_eq ! ( values. len( ) , 1 ) ;
74127 self . num_groups = self . num_groups . max ( total_num_groups) ;
75128 let arr = values[ 0 ] . as_primitive :: < T > ( ) ;
129+
130+ self . buffer . reserve ( arr. len ( ) ) ;
131+
76132 accumulate ( group_indices, arr, opt_filter, |group_idx, value| {
77- self . seen . insert ( ( group_idx, value) ) ;
133+ self . buffer . push ( T :: Native :: pack ( group_idx, value) ) ;
78134 } ) ;
79135 Ok ( ( ) )
80136 }
@@ -85,24 +141,31 @@ where
85141 EmitTo :: First ( n) => n,
86142 } ;
87143
144+ self . buffer . sort_unstable ( ) ;
145+
88146 let mut counts = vec ! [ 0i64 ; num_emitted] ;
89147
90148 if matches ! ( emit_to, EmitTo :: All ) {
91- for & ( group_idx, _) in self . seen . iter ( ) {
149+ self . buffer . dedup ( ) ;
150+ for & packed in & self . buffer {
151+ let ( group_idx, _) = T :: Native :: unpack ( packed) ;
92152 counts[ group_idx] += 1 ;
93153 }
94- self . seen . clear ( ) ;
154+ self . buffer . clear ( ) ;
95155 self . num_groups = 0 ;
96156 } else {
97- let mut remaining = HashSet :: default ( ) ;
98- for ( group_idx, value) in self . seen . drain ( ) {
157+ self . buffer . dedup ( ) ;
158+ let mut remaining = Vec :: new ( ) ;
159+
160+ for & packed in & self . buffer {
161+ let ( group_idx, value) = T :: Native :: unpack ( packed) ;
99162 if group_idx < num_emitted {
100163 counts[ group_idx] += 1 ;
101164 } else {
102- remaining. insert ( ( group_idx - num_emitted, value) ) ;
165+ remaining. push ( T :: Native :: pack ( group_idx - num_emitted, value) ) ;
103166 }
104167 }
105- self . seen = remaining;
168+ self . buffer = remaining;
106169 self . num_groups = self . num_groups . saturating_sub ( num_emitted) ;
107170 }
108171
@@ -115,23 +178,29 @@ where
115178 EmitTo :: First ( n) => n,
116179 } ;
117180
181+ self . buffer . sort_unstable ( ) ;
182+ self . buffer . dedup ( ) ;
183+
118184 let mut group_values: Vec < Vec < T :: Native > > = vec ! [ Vec :: new( ) ; num_emitted] ;
119185
120186 if matches ! ( emit_to, EmitTo :: All ) {
121- for ( group_idx, value) in self . seen . drain ( ) {
187+ for & packed in & self . buffer {
188+ let ( group_idx, value) = T :: Native :: unpack ( packed) ;
122189 group_values[ group_idx] . push ( value) ;
123190 }
191+ self . buffer . clear ( ) ;
124192 self . num_groups = 0 ;
125193 } else {
126- let mut remaining = HashSet :: default ( ) ;
127- for ( group_idx, value) in self . seen . drain ( ) {
194+ let mut remaining = Vec :: new ( ) ;
195+ for & packed in & self . buffer {
196+ let ( group_idx, value) = T :: Native :: unpack ( packed) ;
128197 if group_idx < num_emitted {
129198 group_values[ group_idx] . push ( value) ;
130199 } else {
131- remaining. insert ( ( group_idx - num_emitted, value) ) ;
200+ remaining. push ( T :: Native :: pack ( group_idx - num_emitted, value) ) ;
132201 }
133202 }
134- self . seen = remaining;
203+ self . buffer = remaining;
135204 self . num_groups = self . num_groups . saturating_sub ( num_emitted) ;
136205 }
137206
@@ -167,8 +236,9 @@ where
167236 for ( row_idx, group_idx) in group_indices. iter ( ) . enumerate ( ) {
168237 let inner = list_array. value ( row_idx) ;
169238 let inner_arr = inner. as_primitive :: < T > ( ) ;
239+ self . buffer . reserve ( inner_arr. len ( ) ) ;
170240 for value in inner_arr. values ( ) . iter ( ) {
171- self . seen . insert ( ( * group_idx, * value) ) ;
241+ self . buffer . push ( T :: Native :: pack ( * group_idx, * value) ) ;
172242 }
173243 }
174244
@@ -177,6 +247,6 @@ where
177247
178248 fn size ( & self ) -> usize {
179249 size_of :: < Self > ( )
180- + self . seen . capacity ( ) * ( size_of :: < ( usize , T :: Native ) > ( ) + size_of :: < u64 > ( ) )
250+ + self . buffer . capacity ( ) * size_of :: < < T :: Native as Packable > :: Packed > ( )
181251 }
182252}
0 commit comments