Skip to content

Commit 1959f14

Browse files
author
Artyom Kozhevnikov
committed
filters directly in fragment loading
1 parent c92d01f commit 1959f14

1 file changed

Lines changed: 31 additions & 13 deletions

File tree

  • src/fairseq2/data/parquet/fragment_loading

src/fairseq2/data/parquet/fragment_loading/builder.py

Lines changed: 31 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)