Skip to content

Commit a4215ea

Browse files
committed
WIP: Fetch parquet metadata in streaming network
1 parent cb6d5f5 commit a4215ea

6 files changed

Lines changed: 311 additions & 61 deletions

File tree

python/cudf_polars/cudf_polars/dsl/utils/io.py

Lines changed: 67 additions & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,7 @@
1313

1414
from cudf_polars.dsl.tracing import nvtx_annotate_cudf_polars
1515
from cudf_polars.dsl.traversal import traversal
16-
from cudf_polars.streaming.io import Scan, StreamingScan
16+
from cudf_polars.streaming.io import Scan
1717

1818
if TYPE_CHECKING:
1919
from cudf_polars.dsl.ir import IR
@@ -105,6 +105,70 @@ def _prefetch_parquet_footers_for_paths(paths: list[str]) -> list[CachedParquetI
105105
]
106106

107107

108+
def _cached_parquet_info_from_stats(
109+
stats: StatsCollector | None,
110+
) -> dict[str, CachedParquetInfo]:
111+
"""Return path -> cached parquet info seeded from statistics collection."""
112+
from cudf_polars.streaming.io import ParquetSourceInfo, Scan
113+
114+
cached_parquet_info: dict[str, CachedParquetInfo] = {}
115+
if stats is None:
116+
return cached_parquet_info
117+
118+
for node, datasource_info in stats.scan_stats.items():
119+
if (
120+
isinstance(node, Scan)
121+
and node.typ == "parquet"
122+
and isinstance(datasource_info, ParquetSourceInfo)
123+
and datasource_info.cached_parquet_info is not None
124+
):
125+
for info in datasource_info.cached_parquet_info:
126+
cached_parquet_info[info.path] = info
127+
return cached_parquet_info
128+
129+
130+
@nvtx_annotate_cudf_polars(message="prefetch_cached_parquet_info_for_paths")
131+
def prefetch_cached_parquet_info_for_paths(
132+
paths: list[str],
133+
*,
134+
stats: StatsCollector | None = None,
135+
py_executor: concurrent.futures.Executor | None = None,
136+
) -> list[CachedParquetInfo]:
137+
"""
138+
Prefetch parquet metadata for a path group.
139+
140+
Reuses footers already collected during statistics gathering when
141+
available and fetches any remaining paths.
142+
143+
Parameters
144+
----------
145+
paths
146+
Ordered list of parquet file paths for one scan task group.
147+
stats
148+
Optional statistics collector with already-cached footers.
149+
py_executor
150+
Thread pool used when prefetch must run on a worker thread.
151+
152+
Returns
153+
-------
154+
Cached parquet metadata ordered to match ``paths``.
155+
"""
156+
cached_by_path = _cached_parquet_info_from_stats(stats)
157+
missing_paths = [path for path in paths if path not in cached_by_path]
158+
159+
if missing_paths:
160+
if py_executor is None:
161+
fetched = _prefetch_parquet_footers_for_paths(missing_paths)
162+
else:
163+
fetched = py_executor.submit(
164+
_prefetch_parquet_footers_for_paths, missing_paths
165+
).result()
166+
for info in fetched:
167+
cached_by_path[info.path] = info
168+
169+
return [cached_by_path[path] for path in paths]
170+
171+
108172
@nvtx_annotate_cudf_polars(message="prefetch_parquet_file_metadata_for_ir")
109173
def prefetch_parquet_file_metadata_for_ir(
110174
root: IR,
@@ -130,7 +194,7 @@ def prefetch_parquet_file_metadata_for_ir(
130194
-------
131195
A dictionary mapping each individual path to its cached parquet metadata.
132196
"""
133-
from cudf_polars.streaming.io import ParquetSourceInfo, StreamingScan
197+
from cudf_polars.streaming.io import StreamingScan
134198

135199
all_paths: set[str] = set()
136200

@@ -142,18 +206,7 @@ def prefetch_parquet_file_metadata_for_ir(
142206
elif isinstance(node, Scan) and node.typ == "parquet": # pragma: no cover
143207
raise RuntimeError("Unexpected parquet 'Scan' node in lowered IR graph.")
144208

145-
cached_parquet_info: dict[str, CachedParquetInfo] = {}
146-
if stats is not None:
147-
for node, datasource_info in stats.scan_stats.items():
148-
if (
149-
isinstance(node, Scan)
150-
and node.typ == "parquet"
151-
and isinstance(datasource_info, ParquetSourceInfo)
152-
and datasource_info.cached_parquet_info is not None
153-
):
154-
for info in datasource_info.cached_parquet_info:
155-
cached_parquet_info[info.path] = info
156-
209+
cached_parquet_info = _cached_parquet_info_from_stats(stats)
157210
missing_paths = all_paths - set(cached_parquet_info.keys())
158211
cm: contextlib.AbstractContextManager[concurrent.futures.Executor | None]
159212

@@ -173,28 +226,3 @@ def prefetch_parquet_file_metadata_for_ir(
173226
for info in future.result():
174227
cached_parquet_info[info.path] = info
175228
return cached_parquet_info
176-
177-
178-
def attach_cached_parquet_metadata(
179-
root: IR,
180-
cached_parquet_info_map: dict[str, CachedParquetInfo],
181-
) -> None:
182-
"""
183-
Attach prefetched metadata to scan nodes.
184-
185-
This is an optimization only and does not affect IR identity.
186-
187-
Parameters
188-
----------
189-
root
190-
Root of the IR graph to update.
191-
cached_parquet_info_map
192-
Mapping from file paths to cached parquet metadata.
193-
"""
194-
for node in traversal([root]):
195-
if isinstance(node, StreamingScan):
196-
for scan in node.scans:
197-
cached = [cached_parquet_info_map[path] for path in scan.paths]
198-
Scan._validate_cached_parquet_info(scan.paths, cached)
199-
scan.cached_parquet_info = cached
200-
scan._non_child_args = (*scan._non_child_args[:-1], cached)

python/cudf_polars/cudf_polars/engine/core.py

Lines changed: 0 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -26,10 +26,6 @@
2626

2727
from cudf_polars.containers import DataFrame
2828
from cudf_polars.dsl.ir import IRExecutionContext
29-
from cudf_polars.dsl.utils.io import (
30-
attach_cached_parquet_metadata,
31-
prefetch_parquet_file_metadata_for_ir,
32-
)
3329
from cudf_polars.streaming.actor_graph.collectives import ReserveOpIDs
3430
from cudf_polars.streaming.actor_graph.collectives.common import reserve_op_id
3531
from cudf_polars.streaming.actor_graph.core import generate_network
@@ -696,14 +692,6 @@ def evaluate_on_rank(
696692
py_executor, get_cuda_stream=ctx.br().stream_pool.get_stream, query_id=query_id
697693
)
698694

699-
if config_options.parquet_options.prefetch_file_metadata:
700-
cached_parquet_info_map = prefetch_parquet_file_metadata_for_ir(
701-
ir,
702-
ir_context.py_executor,
703-
stats=stats,
704-
)
705-
attach_cached_parquet_metadata(ir, cached_parquet_info_map)
706-
707695
with ReserveOpIDs(ir, config_options) as collective_id_map:
708696
return execute_ir_on_rank(
709697
ctx,

python/cudf_polars/cudf_polars/streaming/actor_graph/core.py

Lines changed: 88 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,11 +19,12 @@
1919
)
2020
from cudf_polars.dsl.traversal import CachingVisitor, traversal
2121
from cudf_polars.streaming.actor_graph.dispatch import FanoutInfo
22+
from cudf_polars.streaming.actor_graph.io import parquet_metadata_prefetch_node
2223
from cudf_polars.streaming.actor_graph.nodes import (
2324
generate_ir_sub_network_wrapper,
2425
metadata_drain_node,
2526
)
26-
from cudf_polars.streaming.io import StreamingScan
27+
from cudf_polars.streaming.io import StreamingScan, can_use_native_parquet_node
2728
from cudf_polars.streaming.over import Over
2829
from cudf_polars.utils.config import SPMDContext
2930

@@ -35,6 +36,7 @@
3536
from cudf_streaming.channel_metadata import ChannelMetadata
3637
from cudf_streaming.table_chunk import TableChunk
3738
from rapidsmpf.communicator.communicator import Communicator
39+
from rapidsmpf.streaming.chunks.arbitrary import ArbitraryChunk
3840
from rapidsmpf.streaming.core.channel import Channel
3941
from rapidsmpf.streaming.core.context import Context
4042
from rapidsmpf.streaming.core.leaf_actor import DeferredMessages
@@ -44,6 +46,7 @@
4446
GenState,
4547
SubNetGenerator,
4648
)
49+
from cudf_polars.streaming.actor_graph.io import MetadataMessagePayload
4750
from cudf_polars.streaming.base import PartitionInfo, StatsCollector
4851
from cudf_polars.streaming.parallel import ConfigOptions
4952
from cudf_polars.utils.config import StreamingExecutor
@@ -200,6 +203,66 @@ def _mark_children_unbounded(node: IR) -> None:
200203
return fanout_nodes
201204

202205

206+
def _collect_scan_metadata_groups(
207+
ir: IR,
208+
*,
209+
partition_info: MutableMapping[IR, PartitionInfo],
210+
config_options: ConfigOptions,
211+
nranks: int,
212+
) -> dict[tuple[str, ...], set[StreamingScan]]:
213+
groups: defaultdict[tuple[str, ...], set[StreamingScan]] = defaultdict(set)
214+
if not config_options.parquet_options.prefetch_file_metadata:
215+
return {}
216+
217+
for node in traversal([ir]):
218+
if not isinstance(node, StreamingScan):
219+
continue
220+
if node.base_scan.typ != "parquet":
221+
continue
222+
node_partition_info = partition_info[node]
223+
assert node_partition_info.io_plan is not None, (
224+
"Scan node must have a partition plan"
225+
)
226+
use_native = can_use_native_parquet_node(
227+
node.base_scan,
228+
plan=node_partition_info.io_plan,
229+
count=node_partition_info.count,
230+
nranks=nranks,
231+
parquet_options=config_options.parquet_options,
232+
config_options=config_options,
233+
)
234+
if use_native:
235+
continue
236+
for scan in node.scans:
237+
groups[tuple(scan.paths)].add(node)
238+
return dict(groups)
239+
240+
241+
def _build_scan_metadata_channels(
242+
context: Context,
243+
metadata_scan_groups: dict[tuple[str, ...], set[StreamingScan]],
244+
) -> tuple[
245+
dict[tuple[str, ...], tuple[Channel[ArbitraryChunk[MetadataMessagePayload]], ...]],
246+
dict[
247+
StreamingScan,
248+
dict[tuple[str, ...], Channel[ArbitraryChunk[MetadataMessagePayload]]],
249+
],
250+
]:
251+
metadata_group_channels: dict[
252+
tuple[str, ...], tuple[Channel[ArbitraryChunk[MetadataMessagePayload]], ...]
253+
] = {}
254+
metadata_channels_by_scan: dict[
255+
StreamingScan,
256+
dict[tuple[str, ...], Channel[ArbitraryChunk[MetadataMessagePayload]]],
257+
] = {}
258+
for key, scans in metadata_scan_groups.items():
259+
channels = tuple(context.create_channel() for _ in scans)
260+
metadata_group_channels[key] = channels
261+
for scan, channel in zip(sorted(scans, key=id), channels, strict=True):
262+
metadata_channels_by_scan.setdefault(scan, {})[key] = channel
263+
return metadata_group_channels, metadata_channels_by_scan
264+
265+
203266
def generate_network(
204267
context: Context,
205268
comm: Communicator,
@@ -257,6 +320,15 @@ def generate_network(
257320
# Get max_io_threads from config (default: 2)
258321
max_io_threads_global = config_options.executor.max_io_threads
259322
max_io_threads_local = max(1, max_io_threads_global // max(1, num_io_nodes))
323+
metadata_scan_groups = _collect_scan_metadata_groups(
324+
ir,
325+
partition_info=partition_info,
326+
config_options=config_options,
327+
nranks=comm.nranks,
328+
)
329+
metadata_group_channels, metadata_channels_by_scan = _build_scan_metadata_channels(
330+
context, metadata_scan_groups
331+
)
260332

261333
# Generate the network
262334
state: GenState = {
@@ -269,12 +341,26 @@ def generate_network(
269341
"max_io_threads": max_io_threads_local,
270342
"stats": stats,
271343
"collective_id_map": collective_id_map,
344+
"metadata_scan_groups": metadata_scan_groups,
345+
"metadata_group_channels": metadata_group_channels,
346+
"metadata_channels_by_scan": metadata_channels_by_scan,
272347
}
273348
mapper: SubNetGenerator = CachingVisitor(
274349
generate_ir_sub_network_wrapper, state=state
275350
)
276351
nodes_dict, channels = mapper(ir)
277352
ch_out = channels[ir].reserve_output_slot()
353+
metadata_nodes = [
354+
parquet_metadata_prefetch_node(
355+
context,
356+
ir_context,
357+
key,
358+
state["metadata_group_channels"][key],
359+
stats,
360+
sorted(scans, key=id)[0],
361+
)
362+
for key, scans in state["metadata_scan_groups"].items()
363+
]
278364

279365
# Add node to drain metadata before pull_from_channel
280366
# (since pull_from_channel doesn't handle metadata messages)
@@ -294,6 +380,7 @@ def generate_network(
294380

295381
# Flatten the nodes dictionary into a list for run_actor_network
296382
nodes: list[Any] = [node for node_list in nodes_dict.values() for node in node_list]
383+
nodes.extend(metadata_nodes)
297384
nodes.extend([drain_node, output_node])
298385

299386
# Return network and output hook

python/cudf_polars/cudf_polars/streaming/actor_graph/dispatch.py

Lines changed: 22 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
# SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES.
1+
# SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
22
# SPDX-License-Identifier: Apache-2.0
33
"""Dispatching for the RapidsMPF streaming runtime."""
44

@@ -13,14 +13,15 @@
1313
from collections.abc import MutableMapping
1414

1515
from rapidsmpf.communicator.communicator import Communicator
16+
from rapidsmpf.streaming.chunks.arbitrary import ArbitraryChunk
17+
from rapidsmpf.streaming.core.channel import Channel
1618
from rapidsmpf.streaming.core.context import Context
1719

1820
from cudf_polars.dsl.ir import IR, IRExecutionContext
21+
from cudf_polars.streaming.actor_graph.io import MetadataMessagePayload
1922
from cudf_polars.streaming.actor_graph.utils import ChannelManager
20-
from cudf_polars.streaming.base import (
21-
PartitionInfo,
22-
StatsCollector,
23-
)
23+
from cudf_polars.streaming.base import PartitionInfo, StatsCollector
24+
from cudf_polars.streaming.io import StreamingScan
2425
from cudf_polars.utils.config import ConfigOptions, StreamingExecutor
2526

2627

@@ -58,6 +59,14 @@ class GenState(TypedDict):
5859
Statistics collector.
5960
collective_id_map
6061
The mapping of IR nodes to lists of collective IDs.
62+
metadata_scan_groups
63+
Mapping from parquet metadata group key to dependent StreamingScan nodes.
64+
metadata_group_channels
65+
Mapping from parquet metadata group key to output channels for each
66+
dependent StreamingScan node.
67+
metadata_channels_by_scan
68+
Mapping from each StreamingScan node to its parquet metadata input
69+
channels, keyed by parquet metadata group key.
6170
"""
6271

6372
context: Context
@@ -69,6 +78,14 @@ class GenState(TypedDict):
6978
max_io_threads: int
7079
stats: StatsCollector
7180
collective_id_map: dict[IR, list[int]]
81+
metadata_scan_groups: dict[tuple[str, ...], set[StreamingScan]]
82+
metadata_group_channels: dict[
83+
tuple[str, ...], tuple[Channel[ArbitraryChunk[MetadataMessagePayload]], ...]
84+
]
85+
metadata_channels_by_scan: dict[
86+
StreamingScan,
87+
dict[tuple[str, ...], Channel[ArbitraryChunk[MetadataMessagePayload]]],
88+
]
7289

7390

7491
SubNetGenerator: TypeAlias = GenericTransformer[

0 commit comments

Comments
 (0)