@@ -64,35 +64,47 @@ def stable_hash(self, seed=None) -> int:
6464 def load (
6565 self ,
6666 columns : Optional [List [str ]] = None ,
67+ filters : Optional [pa .dataset .Expression ] = None ,
6768 use_threads : bool = False ,
6869 add_fragment_traces : bool = True ,
6970 add_partitioning_columns : bool = True ,
7071 ) -> pa .Table :
72+ physical_schema = self .fragment .physical_schema
7173 if columns is not None :
72- fragment_columns = [
73- col for col in columns if col in self .fragment .physical_schema .names
74- ]
74+ fragment_columns = [col for col in columns if col in physical_schema .names ]
7575 else :
76- fragment_columns = list (self . fragment . physical_schema .names )
76+ fragment_columns = list (physical_schema .names )
7777 # adding technical columns for tracking
7878 if add_fragment_traces :
7979 fragment_columns = list (fragment_columns ) + [
8080 "__batch_index" ,
8181 "__fragment_index" ,
8282 "__filename" ,
8383 ]
84+
85+ can_apply_on_phyiscal_schema = False
86+ if filters is not None :
87+ try :
88+ _ = physical_schema .empty_table ().filter (filters )
89+ can_apply_on_phyiscal_schema = True
90+ except pa .ArrowInvalid as e :
91+ pass
92+
8493 try :
8594 fragment_table = self .fragment .to_table (
86- columns = fragment_columns , use_threads = use_threads
95+ columns = fragment_columns ,
96+ use_threads = use_threads ,
97+ filters = filters if can_apply_on_phyiscal_schema else None ,
8798 )
88-
8999 except OSError as e :
90100 log .info (
91101 "could not load fragment, reinit the fragment state. Error: " , str (e )
92102 )
93103 self .fragment = loads (dumps (self .fragment ))
94104 fragment_table = self .fragment .to_table (
95- columns = fragment_columns , use_threads = use_threads
105+ columns = fragment_columns ,
106+ use_threads = use_threads ,
107+ filters = filters if can_apply_on_phyiscal_schema else None ,
96108 )
97109
98110 if add_partitioning_columns :
@@ -101,6 +113,10 @@ def load(
101113 )
102114 if add_fragment_traces :
103115 fragment_table = add_fragments_trace (fragment_table , self .fragment )
116+
117+ # otherwise, apply filters on full schema
118+ if filters is not None and not can_apply_on_phyiscal_schema :
119+ fragment_table = fragment_table .filter (filters )
104120 return fragment_table
105121
106122
@@ -128,6 +144,7 @@ def load_fn(fragment: pa.dataset.ParquetFileFragment) -> pa.Table | None:
128144 columns = self .columns ,
129145 add_fragment_traces = self .config .add_fragment_traces ,
130146 use_threads = self .config .use_threads ,
147+ filters = self .filters ,
131148 add_partitioning_columns = True ,
132149 )
133150
@@ -141,13 +158,14 @@ def load_fn(fragment: pa.dataset.ParquetFileFragment) -> pa.Table | None:
141158 lambda table : isinstance (table , pa .Table )
142159 )
143160
144- loading_pipeline = loading_pipeline .map (
145- partial (
146- apply_filter ,
147- filters = self .filters ,
148- drop_null = self .config .drop_null ,
161+ if self .config .drop_null :
162+ loading_pipeline = loading_pipeline .map (
163+ partial (
164+ apply_filter ,
165+ filters = None ,
166+ drop_null = self .config .drop_null ,
167+ )
149168 )
150- )
151169
152170 loading_pipeline = loading_pipeline .filter (
153171 lambda table : bool (len (table ) >= self .config .min_batch_size )
0 commit comments