Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
23 commits
Select commit Hold shift + click to select a range
98b2ed5
Fetch parquet metadata in streaming network
TomAugspurger Jul 7, 2026
890094a
Await prefetched metadata on demand
TomAugspurger Jul 9, 2026
a3a39ee
Merge remote-tracking branch 'upstream/main' into tom/cudf-polars-pre…
TomAugspurger Jul 9, 2026
f8001e8
refactor
TomAugspurger Jul 9, 2026
4f9bdba
Merge remote-tracking branch 'upstream/main' into tom/cudf-polars-pre…
TomAugspurger Jul 9, 2026
ef5567e
refactor
TomAugspurger Jul 9, 2026
cd7c15c
better docs on prefetch actor
TomAugspurger Jul 9, 2026
053d78a
Tighten type annotation
TomAugspurger Jul 9, 2026
6dfbd2a
refactor args handling
TomAugspurger Jul 9, 2026
25916ea
Remove dead code
TomAugspurger Jul 9, 2026
2ae9ed3
cancellation
TomAugspurger Jul 9, 2026
fc3827c
Merge remote-tracking branch 'upstream/main' into tom/cudf-polars-pre…
TomAugspurger Jul 10, 2026
182a67c
return type
TomAugspurger Jul 10, 2026
9030ae0
one more
TomAugspurger Jul 10, 2026
999a951
refactor fixtures
TomAugspurger Jul 10, 2026
16746c2
Use send_metadata, recv_metadata
TomAugspurger Jul 10, 2026
603469a
Merge remote-tracking branch 'upstream/main' into tom/cudf-polars-pre…
TomAugspurger Jul 10, 2026
ce47e14
perf
TomAugspurger Jul 10, 2026
e280738
Restore at-most once reading
TomAugspurger Jul 10, 2026
a14421e
remove conditiona
TomAugspurger Jul 10, 2026
33b6a72
inline
TomAugspurger Jul 10, 2026
556ab55
remove pragmas
TomAugspurger Jul 10, 2026
ef25518
Merge remote-tracking branch 'upstream/main' into tom/cudf-polars-pre…
TomAugspurger Jul 10, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 11 additions & 0 deletions python/cudf_polars/cudf_polars/dsl/ir.py
Original file line number Diff line number Diff line change
Expand Up @@ -1711,6 +1711,17 @@ def is_equal(self, other: Self) -> bool:
)
)

def with_prefetched_metadata(
self,
cached_parquet_info: list[CachedParquetInfo] | None,
) -> tuple[Any, ...]:
"""
Return ``self.args`` with cached_parquet_info inserted.

This is a noop for DataFrameScan, which doesn't use parquet metadata.
"""
return self._non_child_args

@classmethod
@log_do_evaluate
@nvtx_annotate_cudf_polars(message="DataFrameScan")
Expand Down
122 changes: 37 additions & 85 deletions python/cudf_polars/cudf_polars/dsl/utils/io.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,19 +4,15 @@

from __future__ import annotations

import concurrent.futures
import contextlib
from dataclasses import dataclass
from typing import TYPE_CHECKING

import pylibcudf as plc

from cudf_polars.dsl.tracing import nvtx_annotate_cudf_polars
from cudf_polars.dsl.traversal import traversal
from cudf_polars.streaming.io import Scan, StreamingScan
from cudf_polars.streaming.io import Scan

if TYPE_CHECKING:
from cudf_polars.dsl.ir import IR
from cudf_polars.streaming.base import StatsCollector


Expand Down Expand Up @@ -105,96 +101,52 @@ def _prefetch_parquet_footers_for_paths(paths: list[str]) -> list[CachedParquetI
]


@nvtx_annotate_cudf_polars(message="prefetch_parquet_file_metadata_for_ir")
def prefetch_parquet_file_metadata_for_ir(
root: IR,
py_executor: concurrent.futures.Executor | None,
stats: StatsCollector | None = None,
def cached_parquet_info_from_stats(
stats: StatsCollector,
) -> dict[str, CachedParquetInfo]:
"""
Prefetch parquet metadata for all parquet scans in an IR graph.

Parameters
----------
root
The root of the IR graph, which will be traversed.
py_executor
The thread pool executor to use for fetching parquet metadata concurrently.
stats
The stats collector. The file metadata might have already been
prefetched during statistics collection, when the number of files
sampled equals the total number of files. Providing ``stats`` here will
skip rereading metadata for those files.

Returns
-------
A dictionary mapping each individual path to its cached parquet metadata.
"""
from cudf_polars.streaming.io import ParquetSourceInfo, StreamingScan

all_paths: set[str] = set()

for node in traversal([root]):
if isinstance(node, StreamingScan):
for scan in node.scans:
for path in scan.paths:
all_paths.add(path)
elif isinstance(node, Scan) and node.typ == "parquet": # pragma: no cover
raise RuntimeError("Unexpected parquet 'Scan' node in lowered IR graph.")
"""Return path -> cached parquet info seeded from statistics collection."""
from cudf_polars.streaming.io import ParquetSourceInfo

cached_parquet_info: dict[str, CachedParquetInfo] = {}
if stats is not None:
for node, datasource_info in stats.scan_stats.items():
if (
isinstance(node, Scan)
and node.typ == "parquet"
and isinstance(datasource_info, ParquetSourceInfo)
and datasource_info.cached_parquet_info is not None
):
for info in datasource_info.cached_parquet_info:
cached_parquet_info[info.path] = info

missing_paths = all_paths - set(cached_parquet_info.keys())
cm: contextlib.AbstractContextManager[concurrent.futures.Executor | None]

if py_executor is None:
cm = py_executor = concurrent.futures.ThreadPoolExecutor()
else:
# We didn't create the executor, so we don't close it.
cm = contextlib.nullcontext()

with cm:
futures = [
py_executor.submit(_prefetch_parquet_footers_for_paths, [path])
for path in missing_paths
]

for future in concurrent.futures.as_completed(futures):
for info in future.result():
for node, datasource_info in stats.scan_stats.items():
if (
isinstance(node, Scan)
and node.typ == "parquet"
and isinstance(datasource_info, ParquetSourceInfo)
and datasource_info.cached_parquet_info is not None
):
for info in datasource_info.cached_parquet_info:
cached_parquet_info[info.path] = info
return cached_parquet_info


def attach_cached_parquet_metadata(
root: IR,
cached_parquet_info_map: dict[str, CachedParquetInfo],
) -> None:
@nvtx_annotate_cudf_polars(message="prefetch_cached_parquet_info_for_paths")
def prefetch_cached_parquet_info_for_paths(
paths: list[str], stats: StatsCollector
) -> list[CachedParquetInfo]:
"""
Attach prefetched metadata to scan nodes.
Prefetch parquet metadata for a path group.

This is an optimization only and does not affect IR identity.
Reuses footers already collected during statistics gathering when
available and fetches any remaining paths.

Parameters
----------
root
Root of the IR graph to update.
cached_parquet_info_map
Mapping from file paths to cached parquet metadata.
paths
Ordered list of parquet file paths for one scan task group.
stats
Optional statistics collector with already-cached footers.

Returns
-------
Cached parquet metadata ordered to match ``paths``.
"""
for node in traversal([root]):
if isinstance(node, StreamingScan):
for scan in node.scans:
cached = [cached_parquet_info_map[path] for path in scan.paths]
Scan._validate_cached_parquet_info(scan.paths, cached)
scan.cached_parquet_info = cached
scan._non_child_args = (*scan._non_child_args[:-1], cached)
cached_by_path = cached_parquet_info_from_stats(stats)
missing_paths = [path for path in paths if path not in cached_by_path]

if missing_paths:
fetched = _prefetch_parquet_footers_for_paths(missing_paths)
for info in fetched:
cached_by_path[info.path] = info

return [cached_by_path[path] for path in paths]
12 changes: 0 additions & 12 deletions python/cudf_polars/cudf_polars/engine/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,10 +27,6 @@

from cudf_polars.containers import DataFrame
from cudf_polars.dsl.ir import IRExecutionContext
from cudf_polars.dsl.utils.io import (
attach_cached_parquet_metadata,
prefetch_parquet_file_metadata_for_ir,
)
from cudf_polars.quent._plan import build_plan
from cudf_polars.streaming.actor_graph.collectives import ReserveOpIDs
from cudf_polars.streaming.actor_graph.collectives.common import reserve_op_id
Expand Down Expand Up @@ -774,14 +770,6 @@ def evaluate_on_rank(
py_executor, get_cuda_stream=ctx.br().stream_pool.get_stream, query_id=query_id
)

if config_options.parquet_options.prefetch_file_metadata:
cached_parquet_info_map = prefetch_parquet_file_metadata_for_ir(
ir,
ir_context.py_executor,
stats=stats,
)
attach_cached_parquet_metadata(ir, cached_parquet_info_map)

with ReserveOpIDs(ir, config_options) as collective_id_map:
return execute_ir_on_rank(
ctx,
Expand Down
28 changes: 28 additions & 0 deletions python/cudf_polars/cudf_polars/streaming/actor_graph/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,11 @@
)
from cudf_polars.dsl.traversal import CachingVisitor, traversal
from cudf_polars.streaming.actor_graph.dispatch import FanoutInfo
from cudf_polars.streaming.actor_graph.io import (
ParquetMetadataCache,
collect_metadata_scans,
parquet_metadata_prefetch_node,
)
from cudf_polars.streaming.actor_graph.nodes import (
generate_ir_sub_network_wrapper,
metadata_drain_node,
Expand Down Expand Up @@ -264,6 +269,16 @@ def generate_network(
# Get max_io_threads from config (default: 2)
max_io_threads_global = config_options.executor.max_io_threads
max_io_threads_local = max(1, max_io_threads_global // max(1, num_io_nodes))
metadata_scans = collect_metadata_scans(
ir,
partition_info=partition_info,
config_options=config_options,
nranks=comm.nranks,
)
metadata_channel_by_scan = {
scan: context.create_channel() for scan in metadata_scans
}
metadata_cache = ParquetMetadataCache(stats)

# Generate the network
state: GenState = {
Expand All @@ -276,12 +291,24 @@ def generate_network(
"max_io_threads": max_io_threads_local,
"stats": stats,
"collective_id_map": collective_id_map,
"metadata_scans": metadata_scans,
"metadata_channel_by_scan": metadata_channel_by_scan,
}
mapper: SubNetGenerator = CachingVisitor(
generate_ir_sub_network_wrapper, state=state
)
nodes_dict, channels = mapper(ir)
ch_out = channels[ir].reserve_output_slot()
metadata_nodes = [
parquet_metadata_prefetch_node(
context,
ir_context,
scan,
metadata_channel_by_scan[scan],
metadata_cache,
)
for scan in metadata_scans
]

# Add node to drain metadata before pull_from_channel
# (since pull_from_channel doesn't handle metadata messages)
Expand All @@ -301,6 +328,7 @@ def generate_network(

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

# Return network and output hook
Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
# SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES.
# SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
"""Dispatching for the RapidsMPF streaming runtime."""

Expand All @@ -13,14 +13,15 @@
from collections.abc import MutableMapping

from rapidsmpf.communicator.communicator import Communicator
from rapidsmpf.streaming.chunks.arbitrary import ArbitraryChunk
from rapidsmpf.streaming.core.channel import Channel
from rapidsmpf.streaming.core.context import Context

from cudf_polars.dsl.ir import IR, IRExecutionContext
from cudf_polars.streaming.actor_graph.io import MetadataMessagePayload
from cudf_polars.streaming.actor_graph.utils import ChannelManager
from cudf_polars.streaming.base import (
PartitionInfo,
StatsCollector,
)
from cudf_polars.streaming.base import PartitionInfo, StatsCollector
from cudf_polars.streaming.io import StreamingScan
from cudf_polars.utils.config import ConfigOptions, StreamingExecutor


Expand Down Expand Up @@ -58,6 +59,11 @@ class GenState(TypedDict):
Statistics collector.
collective_id_map
The mapping of IR nodes to lists of collective IDs.
metadata_scans
Non-native parquet StreamingScan nodes that need metadata prefetch.
metadata_channel_by_scan
Mapping from each eligible StreamingScan node to its single metadata
input channel.
"""

context: Context
Expand All @@ -69,6 +75,10 @@ class GenState(TypedDict):
max_io_threads: int
stats: StatsCollector
collective_id_map: dict[IR, list[int]]
metadata_scans: list[StreamingScan]
metadata_channel_by_scan: dict[
StreamingScan, Channel[ArbitraryChunk[MetadataMessagePayload]]
]


SubNetGenerator: TypeAlias = GenericTransformer[
Expand Down
Loading
Loading