From 98b2ed5afdd4083447f68d9083ebc06250c66b73 Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Tue, 7 Jul 2026 08:32:23 -0700 Subject: [PATCH 01/18] Fetch parquet metadata in streaming network This moves the parquet metadata prefetching that (optionally) happens between IR lowering and streaming network execution into the streaming network. This ensures that any given Scan node can begin executing as soon as *its* metadata is ready, rather than *all* the metadata in a query. At a high level, this PR adds new Actors to the streaming network that are responsible solely for fetching parquet metadata for some paths. The actors send that metadata to the Actors running the Scan nodes. One wrinkle is that we want to preserve the "no duplicate reads of parquet metadata" feature of cudf-polars' prefetching. If two Scan nodes refer to the same paths, we want a single actor that reads the metadata and sends it to both actors doing the Scans. PDS-H Query 7 exhibits this behavior. --- .../cudf_polars/cudf_polars/dsl/utils/io.py | 98 +++++++++------ python/cudf_polars/cudf_polars/engine/core.py | 12 -- .../cudf_polars/streaming/actor_graph/core.py | 91 +++++++++++++- .../streaming/actor_graph/dispatch.py | 27 ++++- .../cudf_polars/streaming/actor_graph/io.py | 112 +++++++++++++++++- .../cudf_polars/tests/streaming/test_scan.py | 67 +++++++++++ 6 files changed, 344 insertions(+), 63 deletions(-) diff --git a/python/cudf_polars/cudf_polars/dsl/utils/io.py b/python/cudf_polars/cudf_polars/dsl/utils/io.py index 288dcd79f460..a8c1ec25fa44 100644 --- a/python/cudf_polars/cudf_polars/dsl/utils/io.py +++ b/python/cudf_polars/cudf_polars/dsl/utils/io.py @@ -13,7 +13,7 @@ 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 @@ -105,6 +105,62 @@ def _prefetch_parquet_footers_for_paths(paths: list[str]) -> list[CachedParquetI ] +def _cached_parquet_info_from_stats( + stats: StatsCollector | None, +) -> dict[str, CachedParquetInfo]: + """Return path -> cached parquet info seeded from statistics collection.""" + from cudf_polars.streaming.io import ParquetSourceInfo, Scan + + cached_parquet_info: dict[str, CachedParquetInfo] = {} + if stats is None: + return cached_parquet_info + + 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 + + +@nvtx_annotate_cudf_polars(message="prefetch_cached_parquet_info_for_paths") +def prefetch_cached_parquet_info_for_paths( + paths: list[str], + *, + stats: StatsCollector | None = None, +) -> list[CachedParquetInfo]: + """ + Prefetch parquet metadata for a path group. + + Reuses footers already collected during statistics gathering when + available and fetches any remaining paths. + + Parameters + ---------- + 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``. + """ + 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] + + @nvtx_annotate_cudf_polars(message="prefetch_parquet_file_metadata_for_ir") def prefetch_parquet_file_metadata_for_ir( root: IR, @@ -130,7 +186,7 @@ def prefetch_parquet_file_metadata_for_ir( ------- A dictionary mapping each individual path to its cached parquet metadata. """ - from cudf_polars.streaming.io import ParquetSourceInfo, StreamingScan + from cudf_polars.streaming.io import StreamingScan all_paths: set[str] = set() @@ -142,18 +198,7 @@ def prefetch_parquet_file_metadata_for_ir( elif isinstance(node, Scan) and node.typ == "parquet": # pragma: no cover raise RuntimeError("Unexpected parquet 'Scan' node in lowered IR graph.") - 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 - + cached_parquet_info = _cached_parquet_info_from_stats(stats) missing_paths = all_paths - set(cached_parquet_info.keys()) cm: contextlib.AbstractContextManager[concurrent.futures.Executor | None] @@ -173,28 +218,3 @@ def prefetch_parquet_file_metadata_for_ir( for info in future.result(): 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: - """ - Attach prefetched metadata to scan nodes. - - This is an optimization only and does not affect IR identity. - - Parameters - ---------- - root - Root of the IR graph to update. - cached_parquet_info_map - Mapping from file paths to cached parquet metadata. - """ - 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) diff --git a/python/cudf_polars/cudf_polars/engine/core.py b/python/cudf_polars/cudf_polars/engine/core.py index c2aea367f683..2a649d7e950f 100644 --- a/python/cudf_polars/cudf_polars/engine/core.py +++ b/python/cudf_polars/cudf_polars/engine/core.py @@ -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 @@ -745,14 +741,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, diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py index c923d342db99..0b88813d44b6 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py @@ -7,8 +7,10 @@ import dataclasses import uuid from collections import defaultdict -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, TypeAlias +from rapidsmpf.streaming.chunks.arbitrary import ArbitraryChunk +from rapidsmpf.streaming.core.channel import Channel from rapidsmpf.streaming.core.leaf_actor import pull_from_channel import cudf_polars.dsl.tracing @@ -19,11 +21,15 @@ ) 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 ( + MetadataMessagePayload, + parquet_metadata_prefetch_node, +) from cudf_polars.streaming.actor_graph.nodes import ( generate_ir_sub_network_wrapper, metadata_drain_node, ) -from cudf_polars.streaming.io import StreamingScan +from cudf_polars.streaming.io import StreamingScan, can_use_native_parquet_node from cudf_polars.streaming.over import Over from cudf_polars.utils.config import SPMDContext @@ -35,7 +41,6 @@ from cudf_streaming.channel_metadata import ChannelMetadata from cudf_streaming.table_chunk import TableChunk from rapidsmpf.communicator.communicator import Communicator - from rapidsmpf.streaming.core.channel import Channel from rapidsmpf.streaming.core.context import Context from rapidsmpf.streaming.core.leaf_actor import DeferredMessages @@ -49,6 +54,13 @@ from cudf_polars.utils.config import StreamingExecutor +Paths: TypeAlias = tuple[str, ...] +MetadataChannel: TypeAlias = Channel[ArbitraryChunk[MetadataMessagePayload]] +MetadataScanGroups: TypeAlias = dict[Paths, set[StreamingScan]] +MetadataGroupChannels: TypeAlias = dict[Paths, tuple[MetadataChannel, ...]] +MetadataChannelsByScan: TypeAlias = dict[StreamingScan, dict[Paths, MetadataChannel]] + + def evaluate_logical_plan( ir: IR, config_options: ConfigOptions[StreamingExecutor], @@ -207,6 +219,55 @@ def _mark_children_unbounded(node: IR) -> None: return fanout_nodes +def _collect_scan_metadata_groups( + ir: IR, + *, + partition_info: MutableMapping[IR, PartitionInfo], + config_options: ConfigOptions, + nranks: int, +) -> dict[tuple[str, ...], set[StreamingScan]]: + groups: defaultdict[tuple[str, ...], set[StreamingScan]] = defaultdict(set) + if not config_options.parquet_options.prefetch_file_metadata: + return {} + + for node in traversal([ir]): + if not isinstance(node, StreamingScan): + continue + if node.base_scan.typ != "parquet": + continue + node_partition_info = partition_info[node] + assert node_partition_info.io_plan is not None, ( + "Scan node must have a partition plan" + ) + use_native = can_use_native_parquet_node( + node.base_scan, + plan=node_partition_info.io_plan, + count=node_partition_info.count, + nranks=nranks, + parquet_options=config_options.parquet_options, + config_options=config_options, + ) + if use_native: + continue + for scan in node.scans: + groups[tuple(scan.paths)].add(node) + return dict(groups) + + +def _build_scan_metadata_channels( + context: Context, + metadata_scan_groups: MetadataScanGroups, +) -> tuple[MetadataGroupChannels, MetadataChannelsByScan]: + metadata_group_channels: MetadataGroupChannels = {} + metadata_channels_by_scan: MetadataChannelsByScan = {} + for key, scans in metadata_scan_groups.items(): + channels = tuple(context.create_channel() for _ in scans) + metadata_group_channels[key] = channels + for scan, channel in zip(sorted(scans, key=id), channels, strict=True): + metadata_channels_by_scan.setdefault(scan, {})[key] = channel + return metadata_group_channels, metadata_channels_by_scan + + def generate_network( context: Context, comm: Communicator, @@ -264,6 +325,15 @@ 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_scan_groups = _collect_scan_metadata_groups( + ir, + partition_info=partition_info, + config_options=config_options, + nranks=comm.nranks, + ) + metadata_group_channels, metadata_channels_by_scan = _build_scan_metadata_channels( + context, metadata_scan_groups + ) # Generate the network state: GenState = { @@ -276,12 +346,26 @@ def generate_network( "max_io_threads": max_io_threads_local, "stats": stats, "collective_id_map": collective_id_map, + "metadata_scan_groups": metadata_scan_groups, + "metadata_group_channels": metadata_group_channels, + "metadata_channels_by_scan": metadata_channels_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, + key, + state["metadata_group_channels"][key], + stats, + sorted(scans, key=id)[0], + ) + for key, scans in state["metadata_scan_groups"].items() + ] # Add node to drain metadata before pull_from_channel # (since pull_from_channel doesn't handle metadata messages) @@ -301,6 +385,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 diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/dispatch.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/dispatch.py index 2554d95fe750..3e25f034934b 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/dispatch.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/dispatch.py @@ -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.""" @@ -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 @@ -58,6 +59,14 @@ class GenState(TypedDict): Statistics collector. collective_id_map The mapping of IR nodes to lists of collective IDs. + metadata_scan_groups + Mapping from parquet metadata group key to dependent StreamingScan nodes. + metadata_group_channels + Mapping from parquet metadata group key to output channels for each + dependent StreamingScan node. + metadata_channels_by_scan + Mapping from each StreamingScan node to its parquet metadata input + channels, keyed by parquet metadata group key. """ context: Context @@ -69,6 +78,14 @@ class GenState(TypedDict): max_io_threads: int stats: StatsCollector collective_id_map: dict[IR, list[int]] + metadata_scan_groups: dict[tuple[str, ...], set[StreamingScan]] + metadata_group_channels: dict[ + tuple[str, ...], tuple[Channel[ArbitraryChunk[MetadataMessagePayload]], ...] + ] + metadata_channels_by_scan: dict[ + StreamingScan, + dict[tuple[str, ...], Channel[ArbitraryChunk[MetadataMessagePayload]]], + ] SubNetGenerator: TypeAlias = GenericTransformer[ diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py index 4c5773874aa5..4719bc1cb450 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py @@ -8,6 +8,8 @@ import functools import io import math +import reprlib +from dataclasses import dataclass from typing import TYPE_CHECKING, Any, cast import polars as pl @@ -16,6 +18,7 @@ from cudf_streaming.channel_metadata import ChannelMetadata from cudf_streaming.table_chunk import TableChunk from rapidsmpf.memory.memory_reservation import opaque_memory_usage +from rapidsmpf.streaming.chunks.arbitrary import ArbitraryChunk from rapidsmpf.streaming.core.memory_reserve_or_wait import ( reserve_memory, ) @@ -30,6 +33,7 @@ _prepare_parquet_predicate, ) from cudf_polars.dsl.to_ast import to_parquet_filter +from cudf_polars.dsl.utils.io import prefetch_cached_parquet_info_for_paths from cudf_polars.streaming.actor_graph.dispatch import ( generate_ir_sub_network, ) @@ -58,13 +62,14 @@ from cudf_polars.streaming.rank_aware_source import RankAwareSource if TYPE_CHECKING: - from collections.abc import Callable, Sequence + from collections.abc import Callable from rapidsmpf.communicator.communicator import Communicator from rapidsmpf.streaming.core.channel import Channel from rapidsmpf.streaming.core.context import Context from cudf_polars.dsl.ir import IR, IRExecutionContext, Scan + from cudf_polars.dsl.utils.io import CachedParquetInfo from cudf_polars.streaming.actor_graph.core import SubNetGenerator from cudf_polars.streaming.actor_graph.tracing import ActorTracer from cudf_polars.streaming.base import ( @@ -76,6 +81,14 @@ from cudf_polars.utils.config import ParquetOptions +@dataclass(frozen=True) +class MetadataMessagePayload: + """Parquet metadata payload sent to scan actors.""" + + group_key: tuple[str, ...] + cached_parquet_info: list[CachedParquetInfo] + + class Lineariser: """ Linearizer that ensures ordered delivery from multiple concurrent producers. @@ -507,6 +520,7 @@ async def read_chunk( ch_out: Channel[TableChunk], ir_context: IRExecutionContext, estimated_chunk_bytes: int, + cached_parquet_info: list[CachedParquetInfo] | None = None, tracer: ActorTracer | None = None, ) -> None: """ @@ -527,9 +541,14 @@ async def read_chunk( estimated_chunk_bytes Estimated size of the chunk in bytes. Used for memory reservation with block spilling to avoid thrashing. + cached_parquet_info + Optional prefetched parquet metadata for parquet scans. tracer The actor tracer for collecting runtime statistics. """ + args = scan._non_child_args + if cached_parquet_info is not None: + args = (*args[:-1], cached_parquet_info) with opaque_memory_usage( await reserve_memory( context, size=estimated_chunk_bytes, net_memory_delta=estimated_chunk_bytes @@ -537,7 +556,7 @@ async def read_chunk( ): df = await ir_context.to_thread( scan.do_evaluate, - *scan._non_child_args, + *args, context=ir_context, ) chunk = TableChunk.from_pylibcudf_table( @@ -549,12 +568,72 @@ async def read_chunk( await send_chunk(context, ch_out, chunk, seq_num, tracer=tracer) +@define_actor() +async def parquet_metadata_prefetch_node( + context: Context, + ir_context: IRExecutionContext, + group_key: tuple[str, ...], + channels: tuple[Channel[ArbitraryChunk[MetadataMessagePayload]], ...], + stats: StatsCollector, + trace_ir: IR, +) -> None: + """Fetch parquet metadata once and fan it out to dependent scans.""" + async with shutdown_on_error( + context, *channels, trace_ir=trace_ir, ir_context=ir_context + ): + cached_parquet_info = await ir_context.to_thread( + prefetch_cached_parquet_info_for_paths, + list(group_key), + stats=stats, + ) + payload = MetadataMessagePayload( + group_key=group_key, + cached_parquet_info=cached_parquet_info, + ) + for ch in channels: + await ch.send( + context, + Message(0, ArbitraryChunk(payload)), + ) + await ch.drain(context) + + +async def recv_prefetched_parquet_metadata( + context: Context, + ch: Channel[ArbitraryChunk[MetadataMessagePayload]], +) -> Message[ArbitraryChunk[MetadataMessagePayload]] | None: + """Receive and validate one prefetched parquet metadata message.""" + return await ch.recv(context) + + +def recv_prefetched_parquet_metadata_handler( + msg: Message[ArbitraryChunk[MetadataMessagePayload]] | None, + group_key: tuple[str, ...], +) -> list[CachedParquetInfo]: + """Synchronous handler for prefetched parquet metadata messages.""" + if msg is None: # pragma: no cover; unreachable + raise AssertionError( + f"Missing parquet metadata message for paths: {reprlib.repr(group_key)}" + ) + payload = ArbitraryChunk[MetadataMessagePayload].from_message(msg).release() + if payload.group_key != group_key: # pragma: no cover; unreachable + difference = set(group_key) ^ set(payload.group_key) + raise AssertionError( + "Unexpected parquet metadata key on scan input channel. " + f"{reprlib.repr(difference)}" + ) + return payload.cached_parquet_info + + @define_actor() async def scan_node( context: Context, ir: StreamingScan, ir_context: IRExecutionContext, ch_out: Channel[TableChunk], + metadata_channels_by_key: dict[ + tuple[str, ...], Channel[ArbitraryChunk[MetadataMessagePayload]] + ], *, num_producers: int, estimated_chunk_bytes: int, @@ -572,17 +651,31 @@ async def scan_node( The execution context for the IR node. ch_out The output Channel[TableChunk]. + metadata_channels_by_key + Mapping from parquet path tuple to the corresponding metadata channel. num_producers The number of producers to use for the scan node. estimated_chunk_bytes Estimated size of each chunk in bytes. Used for memory reservation with block spilling to avoid thrashing. """ - scans: Sequence[SplitScan] | Sequence[FusedScan] = ir.scans + scans = ir.scans + + prefetched_parquet_metadata: dict[tuple[str, ...], list[CachedParquetInfo]] = {} async with shutdown_on_error( - context, ch_out, trace_ir=ir, ir_context=ir_context + context, + ch_out, + *metadata_channels_by_key.values(), + trace_ir=ir, + ir_context=ir_context, ) as tracer: + for key, ch_metadata in metadata_channels_by_key.items(): + msg = await ch_metadata.recv(context) + prefetched_parquet_metadata[key] = recv_prefetched_parquet_metadata_handler( + msg, key + ) + # Send basic metadata await send_metadata( ch_out, @@ -599,6 +692,8 @@ async def scan_node( # skip the lineariser and read the chunks directly if len(scans) == 1 or num_producers == 1: for seq_num, scan in enumerate(scans): + # mypy believes that `scan` is an `IR`, but we know that it's + # actually a `FusedScan | SplitScan` because it's `ir.scans`. await read_chunk( context, scan, @@ -606,6 +701,9 @@ async def scan_node( ch_out, ir_context, estimated_chunk_bytes, + prefetched_parquet_metadata.get( + tuple(cast("FusedScan | SplitScan", scan).paths) + ), tracer=tracer, ) await ch_out.drain(context) @@ -633,6 +731,7 @@ async def _producer(producer_id: int, ch_out: Channel) -> None: ch_out, ir_context, estimated_chunk_bytes, + prefetched_parquet_metadata.get(tuple(scan.paths)), tracer=tracer, ) await ch_out.drain(context) @@ -815,12 +914,17 @@ def _( ) nodes[ir] = [native_node, metadata_node] else: + if ir.base_scan.typ == "parquet" and parquet_options.prefetch_file_metadata: + metadata_channels_by_key = rec.state["metadata_channels_by_scan"][ir] + else: + metadata_channels_by_key = {} nodes[ir] = [ scan_node( rec.state["context"], ir, rec.state["ir_context"], ch_out, + metadata_channels_by_key, num_producers=num_producers, estimated_chunk_bytes=( plan.estimated_chunk_bytes or executor.target_partition_size diff --git a/python/cudf_polars/tests/streaming/test_scan.py b/python/cudf_polars/tests/streaming/test_scan.py index 13e88ead7731..6ba5302b7273 100644 --- a/python/cudf_polars/tests/streaming/test_scan.py +++ b/python/cudf_polars/tests/streaming/test_scan.py @@ -10,6 +10,9 @@ import polars as pl +from rapidsmpf.streaming.chunks.arbitrary import ArbitraryChunk +from rapidsmpf.streaming.core.message import Message + from cudf_polars import Translator from cudf_polars.containers import DataType from cudf_polars.dsl.ir import ( @@ -22,6 +25,10 @@ prefetch_parquet_file_metadata_for_ir, ) from cudf_polars.engine.options import StreamingOptions +from cudf_polars.streaming.actor_graph.io import ( + MetadataMessagePayload, + recv_prefetched_parquet_metadata_handler, +) from cudf_polars.streaming.base import ( DataSourceInfo, IOPartitionFlavor, @@ -124,6 +131,41 @@ def test_scan_parquet_prefetch_file_metadata( assert_gpu_result_equal(pl.scan_parquet(tmp_path), engine=streaming_engine) +def test_scan_parquet_prefetch_metadata_shared_scan_paths( + tmp_path: Path, + df: pl.DataFrame, + streaming_engine_factory: Callable[..., StreamingEngine], +): + streaming_engine = streaming_engine_factory( + StreamingOptions(parquet_options={"prefetch_file_metadata": True}), + ) + make_partitioned_source(df, tmp_path, "parquet", n_files=2) + scan = pl.scan_parquet(tmp_path) + query = pl.concat([scan.select("x"), scan.select("x")]) + assert_gpu_result_equal(query, engine=streaming_engine) + + +def test_scan_parquet_prefetch_metadata_disjoint_scan_paths( + tmp_path: Path, + streaming_engine_factory: Callable[..., StreamingEngine], +): + streaming_engine = streaming_engine_factory( + StreamingOptions(parquet_options={"prefetch_file_metadata": True}), + ) + left = pl.DataFrame({"x": [1, 2, 3]}) + right = pl.DataFrame({"x": [4, 5, 6]}) + left.write_parquet(tmp_path / "left.parquet") + right.write_parquet(tmp_path / "right.parquet") + + query = pl.concat( + [ + pl.scan_parquet(tmp_path / "left.parquet"), + pl.scan_parquet(tmp_path / "right.parquet"), + ] + ) + assert_gpu_result_equal(query, engine=streaming_engine) + + def test_prefetch_file_metadata_non_parquet_scan(df, streaming_engine_factory) -> None: streaming_engine = streaming_engine_factory( StreamingOptions(parquet_options={"prefetch_file_metadata": True}), @@ -673,3 +715,28 @@ def test_scan_partition_plan_nearest( plan = scan_partition_plan(scan, FooStats(scan, file_size), _make_config(10)) assert plan.factor == expected_factor assert plan.flavor == expected_flavor + + +def test_recv_prefetched_parquet_metadata_handler_errors() -> None: + with pytest.raises( + AssertionError, match=r"Missing parquet metadata message for paths: .*" + ): + recv_prefetched_parquet_metadata_handler(None, ("file.parquet",)) + + msg = Message( + 0, + ArbitraryChunk( + MetadataMessagePayload( + group_key=("file.parquet",), + cached_parquet_info=[ + # We don't use file_metadata, so just lie about it. + CachedParquetInfo(path="file.parquet", size=10, file_metadata=None) # type: ignore[arg-type] + ], + ) + ), + ) + with pytest.raises( + AssertionError, + match=r"Unexpected parquet metadata key on scan input channel. .*", + ): + recv_prefetched_parquet_metadata_handler(msg, ("file2.parquet",)) From 890094a3013b8d30fe7e7b899cf84f3b0c5848c5 Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Thu, 9 Jul 2026 07:03:41 -0700 Subject: [PATCH 02/18] Await prefetched metadata on demand --- .../cudf_polars/streaming/actor_graph/io.py | 46 ++++++++++++++----- 1 file changed, 34 insertions(+), 12 deletions(-) diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py index 4719bc1cb450..aa3a3e76929e 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py @@ -661,7 +661,30 @@ async def scan_node( """ scans = ir.scans - prefetched_parquet_metadata: dict[tuple[str, ...], list[CachedParquetInfo]] = {} + # A single file split across multiple SplitScans shares one paths key (and + # one metadata channel), and those splits may run concurrently across + # producers. A channel can only be received once, so memoize a shared + # receive Task per key: the first caller starts it, later callers await the + # same (possibly already-completed) Task. + recv_tasks: dict[tuple[str, ...], asyncio.Task[list[CachedParquetInfo]]] = {} + + async def cached_parquet_info_for_scan( + scan: SplitScan | FusedScan, + ) -> list[CachedParquetInfo] | None: + key = tuple[str, ...](scan.paths) + ch_metadata = metadata_channels_by_key.get(key) + if ch_metadata is None: + # TODO: Figure out what to do here. For parquet Scans with prefetching + # enabled, this would be unexpected. + return None + + async def _recv() -> list[CachedParquetInfo]: + msg = await recv_prefetched_parquet_metadata(context, ch_metadata) + return recv_prefetched_parquet_metadata_handler(msg, key) + + if key not in recv_tasks: + recv_tasks[key] = asyncio.create_task(_recv()) + return await recv_tasks[key] async with shutdown_on_error( context, @@ -670,12 +693,6 @@ async def scan_node( trace_ir=ir, ir_context=ir_context, ) as tracer: - for key, ch_metadata in metadata_channels_by_key.items(): - msg = await ch_metadata.recv(context) - prefetched_parquet_metadata[key] = recv_prefetched_parquet_metadata_handler( - msg, key - ) - # Send basic metadata await send_metadata( ch_out, @@ -694,6 +711,8 @@ async def scan_node( for seq_num, scan in enumerate(scans): # mypy believes that `scan` is an `IR`, but we know that it's # actually a `FusedScan | SplitScan` because it's `ir.scans`. + scan = cast("FusedScan | SplitScan", scan) + cached_parquet_info = await cached_parquet_info_for_scan(scan) await read_chunk( context, scan, @@ -701,9 +720,7 @@ async def scan_node( ch_out, ir_context, estimated_chunk_bytes, - prefetched_parquet_metadata.get( - tuple(cast("FusedScan | SplitScan", scan).paths) - ), + cached_parquet_info, tracer=tracer, ) await ch_out.drain(context) @@ -724,6 +741,7 @@ async def scan_node( async def _producer(producer_id: int, ch_out: Channel) -> None: for task_idx, scan in producer_tasks[producer_id]: + cached_parquet_info = await cached_parquet_info_for_scan(scan) await read_chunk( context, scan, @@ -731,7 +749,7 @@ async def _producer(producer_id: int, ch_out: Channel) -> None: ch_out, ir_context, estimated_chunk_bytes, - prefetched_parquet_metadata.get(tuple(scan.paths)), + cached_parquet_info, tracer=tracer, ) await ch_out.drain(context) @@ -915,7 +933,11 @@ def _( nodes[ir] = [native_node, metadata_node] else: if ir.base_scan.typ == "parquet" and parquet_options.prefetch_file_metadata: - metadata_channels_by_key = rec.state["metadata_channels_by_scan"][ir] + # A node with no assigned scans (e.g. a rank that received no files) + # is never registered in metadata_channels_by_scan, so default to {}. + metadata_channels_by_key = rec.state["metadata_channels_by_scan"].get( + ir, {} + ) else: metadata_channels_by_key = {} nodes[ir] = [ From f8001e8ffd14da1d9c20a6d4aa8b26b21e37b4b5 Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Thu, 9 Jul 2026 12:58:39 -0700 Subject: [PATCH 03/18] refactor Relax the requirement to read exactly once, simplify Prefetch->Scan actor relationship (1:1). --- .../cudf_polars/streaming/actor_graph/core.py | 55 +++---- .../streaming/actor_graph/dispatch.py | 23 +-- .../cudf_polars/streaming/actor_graph/io.py | 133 ++++++++-------- .../cudf_polars/tests/streaming/test_scan.py | 145 ++++++++++++++++++ 4 files changed, 245 insertions(+), 111 deletions(-) diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py index 0b88813d44b6..9199a0b68180 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py @@ -54,11 +54,8 @@ from cudf_polars.utils.config import StreamingExecutor -Paths: TypeAlias = tuple[str, ...] MetadataChannel: TypeAlias = Channel[ArbitraryChunk[MetadataMessagePayload]] -MetadataScanGroups: TypeAlias = dict[Paths, set[StreamingScan]] -MetadataGroupChannels: TypeAlias = dict[Paths, tuple[MetadataChannel, ...]] -MetadataChannelsByScan: TypeAlias = dict[StreamingScan, dict[Paths, MetadataChannel]] +MetadataChannelByScan: TypeAlias = dict[StreamingScan, MetadataChannel] def evaluate_logical_plan( @@ -219,22 +216,25 @@ def _mark_children_unbounded(node: IR) -> None: return fanout_nodes -def _collect_scan_metadata_groups( +def _collect_metadata_scans( ir: IR, *, partition_info: MutableMapping[IR, PartitionInfo], config_options: ConfigOptions, nranks: int, -) -> dict[tuple[str, ...], set[StreamingScan]]: - groups: defaultdict[tuple[str, ...], set[StreamingScan]] = defaultdict(set) +) -> tuple[StreamingScan, ...]: + """Return non-native parquet StreamingScan nodes that need metadata prefetch.""" if not config_options.parquet_options.prefetch_file_metadata: - return {} + return () + metadata_scans: list[StreamingScan] = [] for node in traversal([ir]): if not isinstance(node, StreamingScan): continue if node.base_scan.typ != "parquet": continue + if not node.scans: + continue node_partition_info = partition_info[node] assert node_partition_info.io_plan is not None, ( "Scan node must have a partition plan" @@ -249,23 +249,16 @@ def _collect_scan_metadata_groups( ) if use_native: continue - for scan in node.scans: - groups[tuple(scan.paths)].add(node) - return dict(groups) + metadata_scans.append(node) + return tuple(metadata_scans) -def _build_scan_metadata_channels( +def _build_metadata_channels_by_scan( context: Context, - metadata_scan_groups: MetadataScanGroups, -) -> tuple[MetadataGroupChannels, MetadataChannelsByScan]: - metadata_group_channels: MetadataGroupChannels = {} - metadata_channels_by_scan: MetadataChannelsByScan = {} - for key, scans in metadata_scan_groups.items(): - channels = tuple(context.create_channel() for _ in scans) - metadata_group_channels[key] = channels - for scan, channel in zip(sorted(scans, key=id), channels, strict=True): - metadata_channels_by_scan.setdefault(scan, {})[key] = channel - return metadata_group_channels, metadata_channels_by_scan + metadata_scans: tuple[StreamingScan, ...], +) -> MetadataChannelByScan: + """Create one metadata channel per eligible StreamingScan actor.""" + return {scan: context.create_channel() for scan in metadata_scans} def generate_network( @@ -325,15 +318,13 @@ 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_scan_groups = _collect_scan_metadata_groups( + metadata_scans = _collect_metadata_scans( ir, partition_info=partition_info, config_options=config_options, nranks=comm.nranks, ) - metadata_group_channels, metadata_channels_by_scan = _build_scan_metadata_channels( - context, metadata_scan_groups - ) + metadata_channel_by_scan = _build_metadata_channels_by_scan(context, metadata_scans) # Generate the network state: GenState = { @@ -346,9 +337,8 @@ def generate_network( "max_io_threads": max_io_threads_local, "stats": stats, "collective_id_map": collective_id_map, - "metadata_scan_groups": metadata_scan_groups, - "metadata_group_channels": metadata_group_channels, - "metadata_channels_by_scan": metadata_channels_by_scan, + "metadata_scans": metadata_scans, + "metadata_channel_by_scan": metadata_channel_by_scan, } mapper: SubNetGenerator = CachingVisitor( generate_ir_sub_network_wrapper, state=state @@ -359,12 +349,11 @@ def generate_network( parquet_metadata_prefetch_node( context, ir_context, - key, - state["metadata_group_channels"][key], + scan, + metadata_channel_by_scan[scan], stats, - sorted(scans, key=id)[0], ) - for key, scans in state["metadata_scan_groups"].items() + for scan in metadata_scans ] # Add node to drain metadata before pull_from_channel diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/dispatch.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/dispatch.py index 3e25f034934b..83d3b93e2f2f 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/dispatch.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/dispatch.py @@ -59,14 +59,11 @@ class GenState(TypedDict): Statistics collector. collective_id_map The mapping of IR nodes to lists of collective IDs. - metadata_scan_groups - Mapping from parquet metadata group key to dependent StreamingScan nodes. - metadata_group_channels - Mapping from parquet metadata group key to output channels for each - dependent StreamingScan node. - metadata_channels_by_scan - Mapping from each StreamingScan node to its parquet metadata input - channels, keyed by parquet metadata group key. + 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 @@ -78,13 +75,9 @@ class GenState(TypedDict): max_io_threads: int stats: StatsCollector collective_id_map: dict[IR, list[int]] - metadata_scan_groups: dict[tuple[str, ...], set[StreamingScan]] - metadata_group_channels: dict[ - tuple[str, ...], tuple[Channel[ArbitraryChunk[MetadataMessagePayload]], ...] - ] - metadata_channels_by_scan: dict[ - StreamingScan, - dict[tuple[str, ...], Channel[ArbitraryChunk[MetadataMessagePayload]]], + metadata_scans: tuple[StreamingScan, ...] + metadata_channel_by_scan: dict[ + StreamingScan, Channel[ArbitraryChunk[MetadataMessagePayload]] ] diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py index aa3a3e76929e..53ecfe17eb90 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py @@ -62,7 +62,7 @@ from cudf_polars.streaming.rank_aware_source import RankAwareSource if TYPE_CHECKING: - from collections.abc import Callable + from collections.abc import Callable, Sequence from rapidsmpf.communicator.communicator import Communicator from rapidsmpf.streaming.core.channel import Channel @@ -572,30 +572,55 @@ async def read_chunk( async def parquet_metadata_prefetch_node( context: Context, ir_context: IRExecutionContext, - group_key: tuple[str, ...], - channels: tuple[Channel[ArbitraryChunk[MetadataMessagePayload]], ...], + ir: StreamingScan, + ch_out: Channel[ArbitraryChunk[MetadataMessagePayload]], stats: StatsCollector, - trace_ir: IR, ) -> None: - """Fetch parquet metadata once and fan it out to dependent scans.""" - async with shutdown_on_error( - context, *channels, trace_ir=trace_ir, ir_context=ir_context - ): - cached_parquet_info = await ir_context.to_thread( - prefetch_cached_parquet_info_for_paths, - list(group_key), - stats=stats, - ) - payload = MetadataMessagePayload( - group_key=group_key, - cached_parquet_info=cached_parquet_info, - ) - for ch in channels: - await ch.send( + """Fetch parquet metadata for each scan task and send it to the paired scan actor.""" + async with shutdown_on_error(context, ch_out, trace_ir=ir, ir_context=ir_context): + cached_by_key: dict[tuple[str, ...], list[CachedParquetInfo]] = {} + for scan in ir.scans: + scan = cast("SplitScan | FusedScan", scan) + key = tuple[str, ...](scan.paths) + if key not in cached_by_key: + cached_by_key[key] = await ir_context.to_thread( + prefetch_cached_parquet_info_for_paths, + list(key), + stats=stats, + ) + payload = MetadataMessagePayload( + group_key=key, + cached_parquet_info=cached_by_key[key], + ) + await ch_out.send( context, Message(0, ArbitraryChunk(payload)), ) - await ch.drain(context) + await ch_out.drain(context) + + +def _start_metadata_receiver( + context: Context, + ch_metadata: Channel[ArbitraryChunk[MetadataMessagePayload]], + scans: Sequence[SplitScan | FusedScan], +) -> tuple[list[asyncio.Future[list[CachedParquetInfo]]], asyncio.Task[None]]: + """ + Receive metadata messages sequentially and expose one Future per scan task. + + A single receiver task preserves channel order while allowing concurrent + producers to await only the metadata for their assigned task index. + """ + loop = asyncio.get_running_loop() + futures = [loop.create_future() for _ in range(len(scans))] + + async def _receive() -> None: + for task_idx, scan in enumerate(scans): + msg = await recv_prefetched_parquet_metadata(context, ch_metadata) + cached = recv_prefetched_parquet_metadata_handler(msg, tuple(scan.paths)) + futures[task_idx].set_result(cached) + + receiver_task = asyncio.create_task(_receive()) + return futures, receiver_task async def recv_prefetched_parquet_metadata( @@ -631,9 +656,7 @@ async def scan_node( ir: StreamingScan, ir_context: IRExecutionContext, ch_out: Channel[TableChunk], - metadata_channels_by_key: dict[ - tuple[str, ...], Channel[ArbitraryChunk[MetadataMessagePayload]] - ], + ch_metadata: Channel[ArbitraryChunk[MetadataMessagePayload]] | None, *, num_producers: int, estimated_chunk_bytes: int, @@ -651,45 +674,37 @@ async def scan_node( The execution context for the IR node. ch_out The output Channel[TableChunk]. - metadata_channels_by_key - Mapping from parquet path tuple to the corresponding metadata channel. + ch_metadata + Optional channel carrying prefetched parquet metadata messages, one + per `SplitScan`/`FusedScan` in `ir.scans` order. num_producers The number of producers to use for the scan node. estimated_chunk_bytes Estimated size of each chunk in bytes. Used for memory reservation with block spilling to avoid thrashing. """ - scans = ir.scans - - # A single file split across multiple SplitScans shares one paths key (and - # one metadata channel), and those splits may run concurrently across - # producers. A channel can only be received once, so memoize a shared - # receive Task per key: the first caller starts it, later callers await the - # same (possibly already-completed) Task. - recv_tasks: dict[tuple[str, ...], asyncio.Task[list[CachedParquetInfo]]] = {} + scans = cast("Sequence[SplitScan | FusedScan]", ir.scans) + metadata_futures: list[asyncio.Future[list[CachedParquetInfo]]] | None = None + _metadata_receiver_task: asyncio.Task[None] | None = None + if ch_metadata is not None: + metadata_futures, _metadata_receiver_task = _start_metadata_receiver( + context, ch_metadata, scans + ) - async def cached_parquet_info_for_scan( - scan: SplitScan | FusedScan, + async def cached_parquet_info_for_task( + task_idx: int, ) -> list[CachedParquetInfo] | None: - key = tuple[str, ...](scan.paths) - ch_metadata = metadata_channels_by_key.get(key) - if ch_metadata is None: - # TODO: Figure out what to do here. For parquet Scans with prefetching - # enabled, this would be unexpected. + if metadata_futures is None: return None + return await metadata_futures[task_idx] - async def _recv() -> list[CachedParquetInfo]: - msg = await recv_prefetched_parquet_metadata(context, ch_metadata) - return recv_prefetched_parquet_metadata_handler(msg, key) - - if key not in recv_tasks: - recv_tasks[key] = asyncio.create_task(_recv()) - return await recv_tasks[key] + shutdown_channels: list[Channel[Any]] = [ch_out] + if ch_metadata is not None: + shutdown_channels.append(ch_metadata) async with shutdown_on_error( context, - ch_out, - *metadata_channels_by_key.values(), + *shutdown_channels, trace_ir=ir, ir_context=ir_context, ) as tracer: @@ -709,10 +724,7 @@ async def _recv() -> list[CachedParquetInfo]: # skip the lineariser and read the chunks directly if len(scans) == 1 or num_producers == 1: for seq_num, scan in enumerate(scans): - # mypy believes that `scan` is an `IR`, but we know that it's - # actually a `FusedScan | SplitScan` because it's `ir.scans`. - scan = cast("FusedScan | SplitScan", scan) - cached_parquet_info = await cached_parquet_info_for_scan(scan) + cached_parquet_info = await cached_parquet_info_for_task(seq_num) await read_chunk( context, scan, @@ -736,12 +748,11 @@ async def _recv() -> list[CachedParquetInfo]: ] for task_idx, scan in enumerate(scans): producer_id = task_idx % num_producers - # mypy resolves __iter__ on union-of-sequences to the common base (IR) - producer_tasks[producer_id].append((task_idx, scan)) # type: ignore[arg-type] + producer_tasks[producer_id].append((task_idx, scan)) async def _producer(producer_id: int, ch_out: Channel) -> None: for task_idx, scan in producer_tasks[producer_id]: - cached_parquet_info = await cached_parquet_info_for_scan(scan) + cached_parquet_info = await cached_parquet_info_for_task(task_idx) await read_chunk( context, scan, @@ -933,20 +944,16 @@ def _( nodes[ir] = [native_node, metadata_node] else: if ir.base_scan.typ == "parquet" and parquet_options.prefetch_file_metadata: - # A node with no assigned scans (e.g. a rank that received no files) - # is never registered in metadata_channels_by_scan, so default to {}. - metadata_channels_by_key = rec.state["metadata_channels_by_scan"].get( - ir, {} - ) + ch_metadata = rec.state["metadata_channel_by_scan"].get(ir) else: - metadata_channels_by_key = {} + ch_metadata = None nodes[ir] = [ scan_node( rec.state["context"], ir, rec.state["ir_context"], ch_out, - metadata_channels_by_key, + ch_metadata, num_producers=num_producers, estimated_chunk_bytes=( plan.estimated_chunk_bytes or executor.target_partition_size diff --git a/python/cudf_polars/tests/streaming/test_scan.py b/python/cudf_polars/tests/streaming/test_scan.py index 6ba5302b7273..e4d640002176 100644 --- a/python/cudf_polars/tests/streaming/test_scan.py +++ b/python/cudf_polars/tests/streaming/test_scan.py @@ -19,12 +19,17 @@ Empty, IRExecutionContext, Scan, + Union, ) from cudf_polars.dsl.utils.io import ( CachedParquetInfo, prefetch_parquet_file_metadata_for_ir, ) from cudf_polars.engine.options import StreamingOptions +from cudf_polars.streaming.actor_graph.core import ( + _build_metadata_channels_by_scan, + _collect_metadata_scans, +) from cudf_polars.streaming.actor_graph.io import ( MetadataMessagePayload, recv_prefetched_parquet_metadata_handler, @@ -33,6 +38,7 @@ DataSourceInfo, IOPartitionFlavor, IOPartitionPlan, + PartitionInfo, StatsCollector, ) from cudf_polars.streaming.io import ( @@ -58,6 +64,7 @@ import pylibcudf as plc import cudf_polars.engine.core + from cudf_polars.dsl.ir import IR from cudf_polars.engine.core import StreamingEngine @@ -689,6 +696,16 @@ def _make_config(target: int) -> ConfigOptions: return ConfigOptions.from_polars_engine(engine) +def _make_prefetch_config(target: int) -> ConfigOptions: + engine = pl.GPUEngine( + raise_on_fail=True, + executor="streaming", + executor_options={"target_partition_size": target}, + parquet_options={"prefetch_file_metadata": True}, + ) + return ConfigOptions.from_polars_engine(engine) + + @pytest.mark.parametrize( "file_size,n_paths,expected_factor,expected_flavor", [ @@ -717,6 +734,134 @@ def test_scan_partition_plan_nearest( assert plan.flavor == expected_flavor +def test_collect_metadata_scans_one_actor_per_streaming_scan() -> None: + parquet_options = ParquetOptions(prefetch_file_metadata=True) + paths = [f"part.{i}.parquet" for i in range(6)] + base_scan = _make_parquet_scan(paths, parquet_options) + plan = IOPartitionPlan(9, IOPartitionFlavor.SPLIT_FILES) + partition_count = plan.factor * len(paths) + streaming_scan = expand_scan_for_rank( + base_scan, + plan, + partition_count, + rank=0, + nranks=1, + parquet_options=parquet_options, + ) + assert len(streaming_scan.scans) == partition_count + assert len({tuple(scan.paths) for scan in streaming_scan.scans}) == len(paths) + + config_options = _make_prefetch_config(873_630_000) + partition_info: dict[IR, PartitionInfo] = { + streaming_scan: PartitionInfo(count=partition_count, io_plan=plan), + } + metadata_scans = _collect_metadata_scans( + streaming_scan, + partition_info=partition_info, + config_options=config_options, + nranks=1, + ) + assert metadata_scans == (streaming_scan,) + + +def test_collect_metadata_scans_union_disjoint_paths() -> None: + parquet_options = ParquetOptions(prefetch_file_metadata=True) + plan = IOPartitionPlan(1, IOPartitionFlavor.FUSED_FILES) + left = expand_scan_for_rank( + _make_parquet_scan(["left.parquet"], parquet_options), + plan, + 1, + rank=0, + nranks=1, + parquet_options=parquet_options, + ) + right = expand_scan_for_rank( + _make_parquet_scan(["right.parquet"], parquet_options), + plan, + 1, + rank=0, + nranks=1, + parquet_options=parquet_options, + ) + union = Union(left.schema, None, False, left, right) # noqa: FBT003 + config_options = _make_prefetch_config(10_000) + partition_info: dict[IR, PartitionInfo] = { + left: PartitionInfo(count=1, io_plan=plan), + right: PartitionInfo(count=1, io_plan=plan), + union: PartitionInfo(count=2), + } + metadata_scans = _collect_metadata_scans( + union, + partition_info=partition_info, + config_options=config_options, + nranks=1, + ) + assert metadata_scans == (left, right) + + +def test_collect_metadata_scans_skips_empty_rank() -> None: + parquet_options = ParquetOptions(prefetch_file_metadata=True) + plan = IOPartitionPlan(3, IOPartitionFlavor.SINGLE_READ) + paths = ["a.parquet", "b.parquet", "c.parquet"] + streaming_scan = expand_scan_for_rank( + _make_parquet_scan(paths, parquet_options), + plan, + 1, + rank=1, + nranks=2, + parquet_options=parquet_options, + ) + assert len(streaming_scan.scans) == 0 + config_options = _make_prefetch_config(10_000) + partition_info: dict[IR, PartitionInfo] = { + streaming_scan: PartitionInfo(count=0, io_plan=plan), + } + metadata_scans = _collect_metadata_scans( + streaming_scan, + partition_info=partition_info, + config_options=config_options, + nranks=2, + ) + assert metadata_scans == () + + +def test_build_metadata_channels_by_scan_one_channel_per_scan_actor() -> None: + class DummyContext: + def __init__(self) -> None: + self._channels: list[object] = [] + + def create_channel(self) -> object: + ch = object() + self._channels.append(ch) + return ch + + parquet_options = ParquetOptions(prefetch_file_metadata=True) + left = expand_scan_for_rank( + _make_parquet_scan(["left.parquet"], parquet_options), + IOPartitionPlan(1, IOPartitionFlavor.FUSED_FILES), + 1, + rank=0, + nranks=1, + parquet_options=parquet_options, + ) + right = expand_scan_for_rank( + _make_parquet_scan(["right.parquet"], parquet_options), + IOPartitionPlan(1, IOPartitionFlavor.FUSED_FILES), + 1, + rank=0, + nranks=1, + parquet_options=parquet_options, + ) + context = DummyContext() + metadata_scans = (left, right) + metadata_channel_by_scan = _build_metadata_channels_by_scan( + context, + metadata_scans, + ) + assert set(metadata_channel_by_scan) == {left, right} + assert len(set(metadata_channel_by_scan.values())) == 2 + + def test_recv_prefetched_parquet_metadata_handler_errors() -> None: with pytest.raises( AssertionError, match=r"Missing parquet metadata message for paths: .*" From ef5567ea4cf21f2300ff9746632bfe9d5c863aff Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Thu, 9 Jul 2026 13:32:48 -0700 Subject: [PATCH 04/18] refactor --- .../cudf_polars/streaming/actor_graph/core.py | 64 +++---------------- .../streaming/actor_graph/dispatch.py | 2 +- .../cudf_polars/streaming/actor_graph/io.py | 50 +++++++++++++-- .../cudf_polars/tests/streaming/test_scan.py | 48 ++------------ 4 files changed, 59 insertions(+), 105 deletions(-) diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py index 9199a0b68180..3990ee58daad 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py @@ -7,10 +7,8 @@ import dataclasses import uuid from collections import defaultdict -from typing import TYPE_CHECKING, Any, TypeAlias +from typing import TYPE_CHECKING, Any -from rapidsmpf.streaming.chunks.arbitrary import ArbitraryChunk -from rapidsmpf.streaming.core.channel import Channel from rapidsmpf.streaming.core.leaf_actor import pull_from_channel import cudf_polars.dsl.tracing @@ -22,14 +20,14 @@ 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 ( - MetadataMessagePayload, + collect_metadata_scans, parquet_metadata_prefetch_node, ) from cudf_polars.streaming.actor_graph.nodes import ( generate_ir_sub_network_wrapper, metadata_drain_node, ) -from cudf_polars.streaming.io import StreamingScan, can_use_native_parquet_node +from cudf_polars.streaming.io import StreamingScan from cudf_polars.streaming.over import Over from cudf_polars.utils.config import SPMDContext @@ -41,6 +39,7 @@ from cudf_streaming.channel_metadata import ChannelMetadata from cudf_streaming.table_chunk import TableChunk from rapidsmpf.communicator.communicator import Communicator + from rapidsmpf.streaming.core.channel import Channel from rapidsmpf.streaming.core.context import Context from rapidsmpf.streaming.core.leaf_actor import DeferredMessages @@ -54,10 +53,6 @@ from cudf_polars.utils.config import StreamingExecutor -MetadataChannel: TypeAlias = Channel[ArbitraryChunk[MetadataMessagePayload]] -MetadataChannelByScan: TypeAlias = dict[StreamingScan, MetadataChannel] - - def evaluate_logical_plan( ir: IR, config_options: ConfigOptions[StreamingExecutor], @@ -216,51 +211,6 @@ def _mark_children_unbounded(node: IR) -> None: return fanout_nodes -def _collect_metadata_scans( - ir: IR, - *, - partition_info: MutableMapping[IR, PartitionInfo], - config_options: ConfigOptions, - nranks: int, -) -> tuple[StreamingScan, ...]: - """Return non-native parquet StreamingScan nodes that need metadata prefetch.""" - if not config_options.parquet_options.prefetch_file_metadata: - return () - - metadata_scans: list[StreamingScan] = [] - for node in traversal([ir]): - if not isinstance(node, StreamingScan): - continue - if node.base_scan.typ != "parquet": - continue - if not node.scans: - continue - node_partition_info = partition_info[node] - assert node_partition_info.io_plan is not None, ( - "Scan node must have a partition plan" - ) - use_native = can_use_native_parquet_node( - node.base_scan, - plan=node_partition_info.io_plan, - count=node_partition_info.count, - nranks=nranks, - parquet_options=config_options.parquet_options, - config_options=config_options, - ) - if use_native: - continue - metadata_scans.append(node) - return tuple(metadata_scans) - - -def _build_metadata_channels_by_scan( - context: Context, - metadata_scans: tuple[StreamingScan, ...], -) -> MetadataChannelByScan: - """Create one metadata channel per eligible StreamingScan actor.""" - return {scan: context.create_channel() for scan in metadata_scans} - - def generate_network( context: Context, comm: Communicator, @@ -318,13 +268,15 @@ 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( + metadata_scans = collect_metadata_scans( ir, partition_info=partition_info, config_options=config_options, nranks=comm.nranks, ) - metadata_channel_by_scan = _build_metadata_channels_by_scan(context, metadata_scans) + metadata_channel_by_scan = { + scan: context.create_channel() for scan in metadata_scans + } # Generate the network state: GenState = { diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/dispatch.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/dispatch.py index 83d3b93e2f2f..19c3597b8552 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/dispatch.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/dispatch.py @@ -75,7 +75,7 @@ class GenState(TypedDict): max_io_threads: int stats: StatsCollector collective_id_map: dict[IR, list[int]] - metadata_scans: tuple[StreamingScan, ...] + metadata_scans: list[StreamingScan] metadata_channel_by_scan: dict[ StreamingScan, Channel[ArbitraryChunk[MetadataMessagePayload]] ] diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py index 53ecfe17eb90..d09aa35bde40 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py @@ -10,7 +10,7 @@ import math import reprlib from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, cast +from typing import TYPE_CHECKING, Any, TypeAlias, cast import polars as pl @@ -19,6 +19,7 @@ from cudf_streaming.table_chunk import TableChunk from rapidsmpf.memory.memory_reservation import opaque_memory_usage from rapidsmpf.streaming.chunks.arbitrary import ArbitraryChunk +from rapidsmpf.streaming.core.channel import Channel from rapidsmpf.streaming.core.memory_reserve_or_wait import ( reserve_memory, ) @@ -33,6 +34,7 @@ _prepare_parquet_predicate, ) from cudf_polars.dsl.to_ast import to_parquet_filter +from cudf_polars.dsl.traversal import traversal from cudf_polars.dsl.utils.io import prefetch_cached_parquet_info_for_paths from cudf_polars.streaming.actor_graph.dispatch import ( generate_ir_sub_network, @@ -62,10 +64,9 @@ from cudf_polars.streaming.rank_aware_source import RankAwareSource if TYPE_CHECKING: - from collections.abc import Callable, Sequence + from collections.abc import Callable, MutableMapping, Sequence from rapidsmpf.communicator.communicator import Communicator - from rapidsmpf.streaming.core.channel import Channel from rapidsmpf.streaming.core.context import Context from cudf_polars.dsl.ir import IR, IRExecutionContext, Scan @@ -78,7 +79,11 @@ StatsCollector, ) from cudf_polars.streaming.io import FusedScan, SplitScan - from cudf_polars.utils.config import ParquetOptions + from cudf_polars.utils.config import ConfigOptions, ParquetOptions + + +MetadataChannel: TypeAlias = Channel[ArbitraryChunk["MetadataMessagePayload"]] +MetadataChannelByScan: TypeAlias = dict[StreamingScan, MetadataChannel] @dataclass(frozen=True) @@ -1093,3 +1098,40 @@ def _( ] return nodes, channels + + +def collect_metadata_scans( + ir: IR, + *, + partition_info: MutableMapping[IR, PartitionInfo], + config_options: ConfigOptions, + nranks: int, +) -> list[StreamingScan]: + """Return non-native parquet StreamingScan nodes that need metadata prefetch.""" + if not config_options.parquet_options.prefetch_file_metadata: + return [] + + metadata_scans: list[StreamingScan] = [] + for node in traversal([ir]): + if not isinstance(node, StreamingScan): + continue + if node.base_scan.typ != "parquet": + continue + if not node.scans: + continue + node_partition_info = partition_info[node] + assert node_partition_info.io_plan is not None, ( + "Scan node must have a partition plan" + ) + use_native = can_use_native_parquet_node( + node.base_scan, + plan=node_partition_info.io_plan, + count=node_partition_info.count, + nranks=nranks, + parquet_options=config_options.parquet_options, + config_options=config_options, + ) + if use_native: + continue + metadata_scans.append(node) + return metadata_scans diff --git a/python/cudf_polars/tests/streaming/test_scan.py b/python/cudf_polars/tests/streaming/test_scan.py index e4d640002176..ecad30d6eb76 100644 --- a/python/cudf_polars/tests/streaming/test_scan.py +++ b/python/cudf_polars/tests/streaming/test_scan.py @@ -26,12 +26,9 @@ prefetch_parquet_file_metadata_for_ir, ) from cudf_polars.engine.options import StreamingOptions -from cudf_polars.streaming.actor_graph.core import ( - _build_metadata_channels_by_scan, - _collect_metadata_scans, -) from cudf_polars.streaming.actor_graph.io import ( MetadataMessagePayload, + collect_metadata_scans, recv_prefetched_parquet_metadata_handler, ) from cudf_polars.streaming.base import ( @@ -755,7 +752,7 @@ def test_collect_metadata_scans_one_actor_per_streaming_scan() -> None: partition_info: dict[IR, PartitionInfo] = { streaming_scan: PartitionInfo(count=partition_count, io_plan=plan), } - metadata_scans = _collect_metadata_scans( + metadata_scans = collect_metadata_scans( streaming_scan, partition_info=partition_info, config_options=config_options, @@ -790,7 +787,7 @@ def test_collect_metadata_scans_union_disjoint_paths() -> None: right: PartitionInfo(count=1, io_plan=plan), union: PartitionInfo(count=2), } - metadata_scans = _collect_metadata_scans( + metadata_scans = collect_metadata_scans( union, partition_info=partition_info, config_options=config_options, @@ -816,7 +813,7 @@ def test_collect_metadata_scans_skips_empty_rank() -> None: partition_info: dict[IR, PartitionInfo] = { streaming_scan: PartitionInfo(count=0, io_plan=plan), } - metadata_scans = _collect_metadata_scans( + metadata_scans = collect_metadata_scans( streaming_scan, partition_info=partition_info, config_options=config_options, @@ -825,43 +822,6 @@ def test_collect_metadata_scans_skips_empty_rank() -> None: assert metadata_scans == () -def test_build_metadata_channels_by_scan_one_channel_per_scan_actor() -> None: - class DummyContext: - def __init__(self) -> None: - self._channels: list[object] = [] - - def create_channel(self) -> object: - ch = object() - self._channels.append(ch) - return ch - - parquet_options = ParquetOptions(prefetch_file_metadata=True) - left = expand_scan_for_rank( - _make_parquet_scan(["left.parquet"], parquet_options), - IOPartitionPlan(1, IOPartitionFlavor.FUSED_FILES), - 1, - rank=0, - nranks=1, - parquet_options=parquet_options, - ) - right = expand_scan_for_rank( - _make_parquet_scan(["right.parquet"], parquet_options), - IOPartitionPlan(1, IOPartitionFlavor.FUSED_FILES), - 1, - rank=0, - nranks=1, - parquet_options=parquet_options, - ) - context = DummyContext() - metadata_scans = (left, right) - metadata_channel_by_scan = _build_metadata_channels_by_scan( - context, - metadata_scans, - ) - assert set(metadata_channel_by_scan) == {left, right} - assert len(set(metadata_channel_by_scan.values())) == 2 - - def test_recv_prefetched_parquet_metadata_handler_errors() -> None: with pytest.raises( AssertionError, match=r"Missing parquet metadata message for paths: .*" From cd7c15c7410ec95fa464cb677db930d712086170 Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Thu, 9 Jul 2026 13:47:12 -0700 Subject: [PATCH 05/18] better docs on prefetch actor --- .../cudf_polars/streaming/actor_graph/io.py | 25 ++++++++++++++++++- 1 file changed, 24 insertions(+), 1 deletion(-) diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py index d09aa35bde40..158c86acf591 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py @@ -581,7 +581,30 @@ async def parquet_metadata_prefetch_node( ch_out: Channel[ArbitraryChunk[MetadataMessagePayload]], stats: StatsCollector, ) -> None: - """Fetch parquet metadata for each scan task and send it to the paired scan actor.""" + """ + Fetch parquet metadata for each scan task and send it to the paired scan actor. + + Parameters + ---------- + context + The rapidsmpf context. + ir_context + The execution context for the IR node. Prefetching is offloaded to a thread from + its thread pool. + ir + The StreamingScan node. This actor will send one a message per scan task in this + streaming scan node. + ch_out + The output channel. The Scan actor generated for this StreamingScan node will + read messages from this channel. + stats + The statistics collector, used to populate the parquet metadata cache. + + Notes + ----- + This actor emits one message per SplitScan / FusedScan in the streaming scan. + The messages are sent in the order of the scans. + """ async with shutdown_on_error(context, ch_out, trace_ir=ir, ir_context=ir_context): cached_by_key: dict[tuple[str, ...], list[CachedParquetInfo]] = {} for scan in ir.scans: From 053d78a09a87cc3dbdce261df6f17a561def1c66 Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Thu, 9 Jul 2026 14:07:06 -0700 Subject: [PATCH 06/18] Tighten type annotation --- .../cudf_polars/streaming/actor_graph/io.py | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py index 158c86acf591..96aa14a42460 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py @@ -218,7 +218,7 @@ async def dataframescan_node( ) # Build list of IR slices to read - ir_slices = [] + ir_slices: list[DataFrameScan] = [] # Partial workaround for # https://github.com/pola-rs/polars/issues/23214 If a struct column # has nulls and is sliced then polars exports invalid validity @@ -520,7 +520,7 @@ def _( async def read_chunk( context: Context, - scan: IR, + scan: DataFrameScan | SplitScan | FusedScan, seq_num: int, ch_out: Channel[TableChunk], ir_context: IRExecutionContext, @@ -554,13 +554,18 @@ async def read_chunk( args = scan._non_child_args if cached_parquet_info is not None: args = (*args[:-1], cached_parquet_info) + + # Help mypy with the type inference of the scan.do_evaluate method. + # DataFrameScan, SplitScan, and FusedScan have different signatures for the + # do_evaluate method, but we promise that calling with with `scan.args` is fine. + do_evaluate: Callable[..., DataFrame] = scan.do_evaluate with opaque_memory_usage( await reserve_memory( context, size=estimated_chunk_bytes, net_memory_delta=estimated_chunk_bytes ) ): df = await ir_context.to_thread( - scan.do_evaluate, + do_evaluate, *args, context=ir_context, ) From 6dfbd2a59ca23c35fe8d78f74431264444ecca94 Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Thu, 9 Jul 2026 14:18:12 -0700 Subject: [PATCH 07/18] refactor args handling --- python/cudf_polars/cudf_polars/dsl/ir.py | 11 +++++++++ .../cudf_polars/streaming/actor_graph/io.py | 6 ++--- .../cudf_polars/cudf_polars/streaming/io.py | 18 ++++++++++++++ .../cudf_polars/tests/streaming/test_scan.py | 24 +++++++++++++++++++ 4 files changed, 55 insertions(+), 4 deletions(-) diff --git a/python/cudf_polars/cudf_polars/dsl/ir.py b/python/cudf_polars/cudf_polars/dsl/ir.py index f61a738d0bbd..d7e1d81f750e 100644 --- a/python/cudf_polars/cudf_polars/dsl/ir.py +++ b/python/cudf_polars/cudf_polars/dsl/ir.py @@ -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") diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py index 96aa14a42460..9159858c77a3 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py @@ -551,13 +551,11 @@ async def read_chunk( tracer The actor tracer for collecting runtime statistics. """ - args = scan._non_child_args - if cached_parquet_info is not None: - args = (*args[:-1], cached_parquet_info) + args = scan.with_prefetched_metadata(cached_parquet_info) # Help mypy with the type inference of the scan.do_evaluate method. # DataFrameScan, SplitScan, and FusedScan have different signatures for the - # do_evaluate method, but we promise that calling with with `scan.args` is fine. + # do_evaluate method, but we promise that calling with `args` is fine. do_evaluate: Callable[..., DataFrame] = scan.do_evaluate with opaque_memory_usage( await reserve_memory( diff --git a/python/cudf_polars/cudf_polars/streaming/io.py b/python/cudf_polars/cudf_polars/streaming/io.py index b3b4438812c5..de61e16c75e1 100644 --- a/python/cudf_polars/cudf_polars/streaming/io.py +++ b/python/cudf_polars/cudf_polars/streaming/io.py @@ -273,6 +273,15 @@ def get_hashable(self) -> Hashable: self.parquet_options, ) + def with_prefetched_metadata( + self, + cached_parquet_info: list[CachedParquetInfo] | None, + ) -> tuple[Any, ...]: + """Return ``do_evaluate`` args, substituting prefetched parquet metadata when provided.""" + if cached_parquet_info is None: + return self._non_child_args + return (*self._non_child_args[:-1], cached_parquet_info) + @classmethod def do_evaluate( cls, @@ -441,6 +450,15 @@ def get_hashable(self) -> Hashable: self.parquet_options, ) + def with_prefetched_metadata( + self, + cached_parquet_info: list[CachedParquetInfo] | None, + ) -> tuple[Any, ...]: + """Return ``do_evaluate`` args, substituting prefetched parquet metadata when provided.""" + if cached_parquet_info is None: + return self._non_child_args + return (*self._non_child_args[:-1], cached_parquet_info) + @classmethod def do_evaluate( cls, diff --git a/python/cudf_polars/tests/streaming/test_scan.py b/python/cudf_polars/tests/streaming/test_scan.py index ecad30d6eb76..e42c340a0d35 100644 --- a/python/cudf_polars/tests/streaming/test_scan.py +++ b/python/cudf_polars/tests/streaming/test_scan.py @@ -16,6 +16,7 @@ from cudf_polars import Translator from cudf_polars.containers import DataType from cudf_polars.dsl.ir import ( + DataFrameScan, Empty, IRExecutionContext, Scan, @@ -549,6 +550,29 @@ def test_prefetch_file_metadata_with_cached_scan_parent_nodes( assert_gpu_result_equal(q, engine=engine) +def test_with_prefetched_metadata() -> None: + base = _make_parquet_scan(["a.parquet"]) + info = _make_cached_parquet_info(base.paths) + + dfs = DataFrameScan(base.schema, pl.DataFrame({"x": [1]})._df, None) + assert dfs.with_prefetched_metadata(info) == dfs._non_child_args + assert dfs.with_prefetched_metadata(None) == dfs._non_child_args + + split = SplitScan(base.schema, base, base.paths, 0, 4, base.parquet_options, None) + assert split.with_prefetched_metadata(None) == split._non_child_args + assert split.with_prefetched_metadata(info) == ( + *split._non_child_args[:-1], + info, + ) + + fused = FusedScan(base.schema, base, base.paths, base.parquet_options, None) + assert fused.with_prefetched_metadata(None) == fused._non_child_args + assert fused.with_prefetched_metadata(info) == ( + *fused._non_child_args[:-1], + info, + ) + + def test_fused_scan_identity_equality() -> None: base = _make_parquet_scan(["a.parquet", "b.parquet"]) paths = ["a.parquet"] From 25916ea5db160ae9e7ade85fe5cdf2ed81294140 Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Thu, 9 Jul 2026 14:23:10 -0700 Subject: [PATCH 08/18] Remove dead code --- .../cudf_polars/cudf_polars/dsl/utils/io.py | 65 +------------------ .../cudf_polars/tests/streaming/test_scan.py | 9 --- 2 files changed, 1 insertion(+), 73 deletions(-) diff --git a/python/cudf_polars/cudf_polars/dsl/utils/io.py b/python/cudf_polars/cudf_polars/dsl/utils/io.py index a8c1ec25fa44..4b9126450d4d 100644 --- a/python/cudf_polars/cudf_polars/dsl/utils/io.py +++ b/python/cudf_polars/cudf_polars/dsl/utils/io.py @@ -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 if TYPE_CHECKING: - from cudf_polars.dsl.ir import IR from cudf_polars.streaming.base import StatsCollector @@ -109,7 +105,7 @@ def _cached_parquet_info_from_stats( stats: StatsCollector | None, ) -> dict[str, CachedParquetInfo]: """Return path -> cached parquet info seeded from statistics collection.""" - from cudf_polars.streaming.io import ParquetSourceInfo, Scan + from cudf_polars.streaming.io import ParquetSourceInfo cached_parquet_info: dict[str, CachedParquetInfo] = {} if stats is None: @@ -159,62 +155,3 @@ def prefetch_cached_parquet_info_for_paths( cached_by_path[info.path] = info return [cached_by_path[path] for path in paths] - - -@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, -) -> 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 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.") - - cached_parquet_info = _cached_parquet_info_from_stats(stats) - 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(): - cached_parquet_info[info.path] = info - return cached_parquet_info diff --git a/python/cudf_polars/tests/streaming/test_scan.py b/python/cudf_polars/tests/streaming/test_scan.py index e42c340a0d35..f91c8fcf564c 100644 --- a/python/cudf_polars/tests/streaming/test_scan.py +++ b/python/cudf_polars/tests/streaming/test_scan.py @@ -17,14 +17,12 @@ from cudf_polars.containers import DataType from cudf_polars.dsl.ir import ( DataFrameScan, - Empty, IRExecutionContext, Scan, Union, ) from cudf_polars.dsl.utils.io import ( CachedParquetInfo, - prefetch_parquet_file_metadata_for_ir, ) from cudf_polars.engine.options import StreamingOptions from cudf_polars.streaming.actor_graph.io import ( @@ -178,13 +176,6 @@ def test_prefetch_file_metadata_non_parquet_scan(df, streaming_engine_factory) - assert_gpu_result_equal(df.lazy().select("x"), engine=streaming_engine) -def test_prefetch_parquet_file_metadata_no_parquet_scans() -> None: - result = prefetch_parquet_file_metadata_for_ir( - Empty({}), py_executor=None, stats=None - ) - assert result == {} - - def test_prefetch_file_metadata_select_fast_count( df: pl.DataFrame, streaming_engine_factory: Callable[..., StreamingEngine], From 2ae9ed333556b0eb053c661a3aadfe2f3e7c9211 Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Thu, 9 Jul 2026 14:41:59 -0700 Subject: [PATCH 09/18] cancellation --- .../cudf_polars/streaming/actor_graph/io.py | 174 ++++++++++-------- 1 file changed, 100 insertions(+), 74 deletions(-) diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py index 9159858c77a3..2b7a164d6dab 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py @@ -5,6 +5,7 @@ from __future__ import annotations import asyncio +import contextlib import functools import io import math @@ -645,10 +646,24 @@ def _start_metadata_receiver( futures = [loop.create_future() for _ in range(len(scans))] async def _receive() -> None: - for task_idx, scan in enumerate(scans): - msg = await recv_prefetched_parquet_metadata(context, ch_metadata) - cached = recv_prefetched_parquet_metadata_handler(msg, tuple(scan.paths)) - futures[task_idx].set_result(cached) + try: + for task_idx, scan in enumerate(scans): + msg = await recv_prefetched_parquet_metadata(context, ch_metadata) + cached = recv_prefetched_parquet_metadata_handler( + msg, tuple(scan.paths) + ) + futures[task_idx].set_result(cached) + except BaseException as exc: + # Propagate the failure (including cancellation / premature stop) to + # every unfinished future so downstream awaiters fail fast instead of + # hanging on metadata that will never arrive. + for future in futures: + if not future.done(): + if isinstance(exc, asyncio.CancelledError): + future.cancel() + else: + future.set_exception(exc) + raise receiver_task = asyncio.create_task(_receive()) return futures, receiver_task @@ -733,79 +748,90 @@ async def cached_parquet_info_for_task( if ch_metadata is not None: shutdown_channels.append(ch_metadata) - async with shutdown_on_error( - context, - *shutdown_channels, - trace_ir=ir, - ir_context=ir_context, - ) as tracer: - # Send basic metadata - await send_metadata( - ch_out, + try: + async with shutdown_on_error( context, - ChannelMetadata(local_count=len(scans)), - ) - - # If there is nothing to scan, drain the channel and return - if len(scans) == 0: - await ch_out.drain(context) - return - - # If there is only one scan or one producer, we can - # skip the lineariser and read the chunks directly - if len(scans) == 1 or num_producers == 1: - for seq_num, scan in enumerate(scans): - cached_parquet_info = await cached_parquet_info_for_task(seq_num) - await read_chunk( - context, - scan, - seq_num, - ch_out, - ir_context, - estimated_chunk_bytes, - cached_parquet_info, - tracer=tracer, - ) - await ch_out.drain(context) - return - - # Use Lineariser to ensure ordered delivery - num_producers = min(num_producers, len(scans)) - lineariser = Lineariser(context, ch_out, num_producers) - - # Assign tasks to producers using round-robin - producer_tasks: list[list[tuple[int, SplitScan | FusedScan]]] = [ - [] for _ in range(num_producers) - ] - for task_idx, scan in enumerate(scans): - producer_id = task_idx % num_producers - producer_tasks[producer_id].append((task_idx, scan)) + *shutdown_channels, + trace_ir=ir, + ir_context=ir_context, + ) as tracer: + # Send basic metadata + await send_metadata( + ch_out, + context, + ChannelMetadata(local_count=len(scans)), + ) - async def _producer(producer_id: int, ch_out: Channel) -> None: - for task_idx, scan in producer_tasks[producer_id]: - cached_parquet_info = await cached_parquet_info_for_task(task_idx) - await read_chunk( - context, - scan, - task_idx, - ch_out, - ir_context, - estimated_chunk_bytes, - cached_parquet_info, - tracer=tracer, + # If there is nothing to scan, drain the channel and return + if len(scans) == 0: + await ch_out.drain(context) + return + + # If there is only one scan or one producer, we can + # skip the lineariser and read the chunks directly + if len(scans) == 1 or num_producers == 1: + for seq_num, scan in enumerate(scans): + cached_parquet_info = await cached_parquet_info_for_task(seq_num) + await read_chunk( + context, + scan, + seq_num, + ch_out, + ir_context, + estimated_chunk_bytes, + cached_parquet_info, + tracer=tracer, + ) + await ch_out.drain(context) + return + + # Use Lineariser to ensure ordered delivery + num_producers = min(num_producers, len(scans)) + lineariser = Lineariser(context, ch_out, num_producers) + + # Assign tasks to producers using round-robin + producer_tasks: list[list[tuple[int, SplitScan | FusedScan]]] = [ + [] for _ in range(num_producers) + ] + for task_idx, scan in enumerate(scans): + producer_id = task_idx % num_producers + producer_tasks[producer_id].append((task_idx, scan)) + + async def _producer(producer_id: int, ch_out: Channel) -> None: + for task_idx, scan in producer_tasks[producer_id]: + cached_parquet_info = await cached_parquet_info_for_task(task_idx) + await read_chunk( + context, + scan, + task_idx, + ch_out, + ir_context, + estimated_chunk_bytes, + cached_parquet_info, + tracer=tracer, + ) + await ch_out.drain(context) + + async with ( + shutdown_on_error(context, *lineariser.input_channels, trace_ir=ir), + ): + await gather_in_task_group( + lineariser.drain(), + *( + _producer(i, ch_in) + for i, ch_in in enumerate(lineariser.input_channels) + ), ) - await ch_out.drain(context) - - async with ( - shutdown_on_error(context, *lineariser.input_channels, trace_ir=ir), - ): - await gather_in_task_group( - lineariser.drain(), - *( - _producer(i, ch_in) - for i, ch_in in enumerate(lineariser.input_channels) - ), - ) + finally: + # Always finalize the background metadata receiver, even on early + # return or failure, so it is never left orphaned. + if _metadata_receiver_task is not None: + _metadata_receiver_task.cancel() + # Awaiting also retrieves any exception the receiver raised (already + # surfaced to producers via the per-task futures), avoiding a stray + # "Task exception was never retrieved" warning. + with contextlib.suppress(BaseException): + await _metadata_receiver_task def make_rapidsmpf_read_parquet_node( From 182a67ca0f4a5d615db341123c40c6e605ffec8a Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Fri, 10 Jul 2026 06:05:44 -0700 Subject: [PATCH 10/18] return type --- python/cudf_polars/cudf_polars/dsl/utils/io.py | 9 ++------- python/cudf_polars/tests/streaming/test_scan.py | 4 ++-- 2 files changed, 4 insertions(+), 9 deletions(-) diff --git a/python/cudf_polars/cudf_polars/dsl/utils/io.py b/python/cudf_polars/cudf_polars/dsl/utils/io.py index 4b9126450d4d..b2e65cd4a487 100644 --- a/python/cudf_polars/cudf_polars/dsl/utils/io.py +++ b/python/cudf_polars/cudf_polars/dsl/utils/io.py @@ -102,15 +102,12 @@ def _prefetch_parquet_footers_for_paths(paths: list[str]) -> list[CachedParquetI def _cached_parquet_info_from_stats( - stats: StatsCollector | None, + stats: StatsCollector, ) -> dict[str, CachedParquetInfo]: """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 None: - return cached_parquet_info - for node, datasource_info in stats.scan_stats.items(): if ( isinstance(node, Scan) @@ -125,9 +122,7 @@ def _cached_parquet_info_from_stats( @nvtx_annotate_cudf_polars(message="prefetch_cached_parquet_info_for_paths") def prefetch_cached_parquet_info_for_paths( - paths: list[str], - *, - stats: StatsCollector | None = None, + paths: list[str], *, stats: StatsCollector ) -> list[CachedParquetInfo]: """ Prefetch parquet metadata for a path group. diff --git a/python/cudf_polars/tests/streaming/test_scan.py b/python/cudf_polars/tests/streaming/test_scan.py index f91c8fcf564c..85507a907b6e 100644 --- a/python/cudf_polars/tests/streaming/test_scan.py +++ b/python/cudf_polars/tests/streaming/test_scan.py @@ -808,7 +808,7 @@ def test_collect_metadata_scans_union_disjoint_paths() -> None: config_options=config_options, nranks=1, ) - assert metadata_scans == (left, right) + assert metadata_scans == [left, right] def test_collect_metadata_scans_skips_empty_rank() -> None: @@ -834,7 +834,7 @@ def test_collect_metadata_scans_skips_empty_rank() -> None: config_options=config_options, nranks=2, ) - assert metadata_scans == () + assert metadata_scans == [] def test_recv_prefetched_parquet_metadata_handler_errors() -> None: From 9030ae0a83ad419c24e8868fbc5285fe4377535d Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Fri, 10 Jul 2026 06:22:44 -0700 Subject: [PATCH 11/18] one more --- python/cudf_polars/tests/streaming/test_scan.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/python/cudf_polars/tests/streaming/test_scan.py b/python/cudf_polars/tests/streaming/test_scan.py index 85507a907b6e..b93f0c001ddd 100644 --- a/python/cudf_polars/tests/streaming/test_scan.py +++ b/python/cudf_polars/tests/streaming/test_scan.py @@ -773,7 +773,7 @@ def test_collect_metadata_scans_one_actor_per_streaming_scan() -> None: config_options=config_options, nranks=1, ) - assert metadata_scans == (streaming_scan,) + assert metadata_scans == [streaming_scan] def test_collect_metadata_scans_union_disjoint_paths() -> None: From 999a9515689a64aea9a0022602303eade3930845 Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Fri, 10 Jul 2026 07:44:17 -0700 Subject: [PATCH 12/18] refactor fixtures --- .../cudf_polars/tests/streaming/test_scan.py | 56 ++++++++----------- 1 file changed, 24 insertions(+), 32 deletions(-) diff --git a/python/cudf_polars/tests/streaming/test_scan.py b/python/cudf_polars/tests/streaming/test_scan.py index b93f0c001ddd..9ea7383a9a20 100644 --- a/python/cudf_polars/tests/streaming/test_scan.py +++ b/python/cudf_polars/tests/streaming/test_scan.py @@ -75,6 +75,16 @@ def df(): ) +@pytest.fixture +def prefetch_file_metadata_engine( + streaming_engine_factory: Callable[..., StreamingEngine], +): + """Streaming Engine fixture with parquet metadata prefetching enabled.""" + return streaming_engine_factory( + StreamingOptions(parquet_options={"prefetch_file_metadata": True}), + ) + + @pytest.mark.parametrize( "fmt, scan_fn", [ @@ -137,24 +147,18 @@ def test_scan_parquet_prefetch_file_metadata( def test_scan_parquet_prefetch_metadata_shared_scan_paths( tmp_path: Path, df: pl.DataFrame, - streaming_engine_factory: Callable[..., StreamingEngine], + prefetch_file_metadata_engine: StreamingEngine, ): - streaming_engine = streaming_engine_factory( - StreamingOptions(parquet_options={"prefetch_file_metadata": True}), - ) make_partitioned_source(df, tmp_path, "parquet", n_files=2) scan = pl.scan_parquet(tmp_path) query = pl.concat([scan.select("x"), scan.select("x")]) - assert_gpu_result_equal(query, engine=streaming_engine) + assert_gpu_result_equal(query, engine=prefetch_file_metadata_engine) def test_scan_parquet_prefetch_metadata_disjoint_scan_paths( tmp_path: Path, - streaming_engine_factory: Callable[..., StreamingEngine], + prefetch_file_metadata_engine: StreamingEngine, ): - streaming_engine = streaming_engine_factory( - StreamingOptions(parquet_options={"prefetch_file_metadata": True}), - ) left = pl.DataFrame({"x": [1, 2, 3]}) right = pl.DataFrame({"x": [4, 5, 6]}) left.write_parquet(tmp_path / "left.parquet") @@ -166,28 +170,24 @@ def test_scan_parquet_prefetch_metadata_disjoint_scan_paths( pl.scan_parquet(tmp_path / "right.parquet"), ] ) - assert_gpu_result_equal(query, engine=streaming_engine) + assert_gpu_result_equal(query, engine=prefetch_file_metadata_engine) -def test_prefetch_file_metadata_non_parquet_scan(df, streaming_engine_factory) -> None: - streaming_engine = streaming_engine_factory( - StreamingOptions(parquet_options={"prefetch_file_metadata": True}), - ) - assert_gpu_result_equal(df.lazy().select("x"), engine=streaming_engine) +def test_prefetch_file_metadata_non_parquet_scan( + df: pl.DataFrame, prefetch_file_metadata_engine: StreamingEngine +) -> None: + assert_gpu_result_equal(df.lazy().select("x"), engine=prefetch_file_metadata_engine) def test_prefetch_file_metadata_select_fast_count( df: pl.DataFrame, - streaming_engine_factory: Callable[..., StreamingEngine], + prefetch_file_metadata_engine: StreamingEngine, tmp_path: Path, ) -> None: - streaming_engine = streaming_engine_factory( - StreamingOptions(parquet_options={"prefetch_file_metadata": True}), - ) source = tmp_path / "data.parquet" df.write_parquet(source) q = pl.scan_parquet(source).select(pl.len()) - assert_gpu_result_equal(q, engine=streaming_engine) + assert_gpu_result_equal(q, engine=prefetch_file_metadata_engine) # --------------------------------------------------------------------------- @@ -487,19 +487,15 @@ def test_split_scan_do_evaluate_missing_prefetch_metadata() -> None: def test_prefetch_file_metadata_join( - tmp_path: Path, streaming_engine_factory: Callable[..., StreamingEngine] + tmp_path: Path, prefetch_file_metadata_engine: StreamingEngine ) -> None: p1 = tmp_path / "f1.parquet" p2 = tmp_path / "f2.parquet" pl.DataFrame({"k": [1, 2, 3], "a": [4, 5, 6]}).write_parquet(p1) pl.DataFrame({"k": [1, 2, 3], "b": [7, 8, 9]}).write_parquet(p2) - engine = streaming_engine_factory( - StreamingOptions(parquet_options={"prefetch_file_metadata": True}), - ) - q = pl.scan_parquet(p1).join(pl.scan_parquet(p2), on="k") - q.collect(engine=engine) + q.collect(engine=prefetch_file_metadata_engine) def _make_cached_parquet_info( @@ -518,7 +514,7 @@ def _make_cached_parquet_info( def test_prefetch_file_metadata_with_cached_scan_parent_nodes( - tmp_path: Path, streaming_engine_factory: Callable[..., StreamingEngine] + tmp_path: Path, prefetch_file_metadata_engine: StreamingEngine ) -> None: # Regression test for replace not replacing StreamingScan nodes with their prefetched variants. source = tmp_path / "data.parquet" @@ -529,16 +525,12 @@ def test_prefetch_file_metadata_with_cached_scan_parent_nodes( } ).write_parquet(source) - engine = streaming_engine_factory( - StreamingOptions(parquet_options={"prefetch_file_metadata": True}), - ) - cached_scan = pl.scan_parquet(source).cache() left = cached_scan.group_by("k").agg(pl.col("v").sum().alias("sum_v")) right = cached_scan.group_by("k").agg(pl.len().alias("n")) q = left.join(right, on="k").sort("k") - assert_gpu_result_equal(q, engine=engine) + assert_gpu_result_equal(q, engine=prefetch_file_metadata_engine) def test_with_prefetched_metadata() -> None: From 16746c20a233184d033aaa4081580a4019c77a7d Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Fri, 10 Jul 2026 09:20:25 -0700 Subject: [PATCH 13/18] Use send_metadata, recv_metadata This lets the sender get arbitrarily far ahead of the receiever, which we want for parquet metadata footers (probably). --- python/cudf_polars/cudf_polars/streaming/actor_graph/io.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py index 2b7a164d6dab..8adc5ef4bbff 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py @@ -624,7 +624,7 @@ async def parquet_metadata_prefetch_node( group_key=key, cached_parquet_info=cached_by_key[key], ) - await ch_out.send( + await ch_out.send_metadata( context, Message(0, ArbitraryChunk(payload)), ) @@ -674,7 +674,7 @@ async def recv_prefetched_parquet_metadata( ch: Channel[ArbitraryChunk[MetadataMessagePayload]], ) -> Message[ArbitraryChunk[MetadataMessagePayload]] | None: """Receive and validate one prefetched parquet metadata message.""" - return await ch.recv(context) + return await ch.recv_metadata(context) def recv_prefetched_parquet_metadata_handler( From ce47e14cc9bdc48140537071cd118e281d4f5d33 Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Fri, 10 Jul 2026 09:22:36 -0700 Subject: [PATCH 14/18] perf --- python/cudf_polars/tests/streaming/test_scan.py | 11 +++++++++++ 1 file changed, 11 insertions(+) diff --git a/python/cudf_polars/tests/streaming/test_scan.py b/python/cudf_polars/tests/streaming/test_scan.py index 9ea7383a9a20..8d112704955c 100644 --- a/python/cudf_polars/tests/streaming/test_scan.py +++ b/python/cudf_polars/tests/streaming/test_scan.py @@ -144,11 +144,22 @@ def test_scan_parquet_prefetch_file_metadata( assert_gpu_result_equal(pl.scan_parquet(tmp_path), engine=streaming_engine) +@pytest.mark.timeout(90) def test_scan_parquet_prefetch_metadata_shared_scan_paths( tmp_path: Path, df: pl.DataFrame, prefetch_file_metadata_engine: StreamingEngine, ): + # The spmd-small case creates *many* partitions with the length-3000 df. + # A smaller dataframe gives us sufficient test coverage, and runs much faster. + if ( + prefetch_file_metadata_engine.config["executor_options"][ + "max_rows_per_partition" + ] + == SMALL_MAX_ROWS_PER_PARTITION + ): + df = df.head(40) + make_partitioned_source(df, tmp_path, "parquet", n_files=2) scan = pl.scan_parquet(tmp_path) query = pl.concat([scan.select("x"), scan.select("x")]) From e28073849decdb244ef0074116531b5558da59ba Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Fri, 10 Jul 2026 10:40:05 -0700 Subject: [PATCH 15/18] Restore at-most once reading --- .../cudf_polars/cudf_polars/dsl/utils/io.py | 6 +- .../cudf_polars/streaming/actor_graph/core.py | 4 +- .../cudf_polars/streaming/actor_graph/io.py | 88 ++++++++++++++++--- .../cudf_polars/tests/streaming/test_scan.py | 28 ++++++ 4 files changed, 111 insertions(+), 15 deletions(-) diff --git a/python/cudf_polars/cudf_polars/dsl/utils/io.py b/python/cudf_polars/cudf_polars/dsl/utils/io.py index b2e65cd4a487..403dd8358770 100644 --- a/python/cudf_polars/cudf_polars/dsl/utils/io.py +++ b/python/cudf_polars/cudf_polars/dsl/utils/io.py @@ -101,7 +101,7 @@ def _prefetch_parquet_footers_for_paths(paths: list[str]) -> list[CachedParquetI ] -def _cached_parquet_info_from_stats( +def cached_parquet_info_from_stats( stats: StatsCollector, ) -> dict[str, CachedParquetInfo]: """Return path -> cached parquet info seeded from statistics collection.""" @@ -122,7 +122,7 @@ def _cached_parquet_info_from_stats( @nvtx_annotate_cudf_polars(message="prefetch_cached_parquet_info_for_paths") def prefetch_cached_parquet_info_for_paths( - paths: list[str], *, stats: StatsCollector + paths: list[str], stats: StatsCollector ) -> list[CachedParquetInfo]: """ Prefetch parquet metadata for a path group. @@ -141,7 +141,7 @@ def prefetch_cached_parquet_info_for_paths( ------- Cached parquet metadata ordered to match ``paths``. """ - cached_by_path = _cached_parquet_info_from_stats(stats) + 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: diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py index 3990ee58daad..2c1d7977b636 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/core.py @@ -20,6 +20,7 @@ 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, ) @@ -277,6 +278,7 @@ def generate_network( metadata_channel_by_scan = { scan: context.create_channel() for scan in metadata_scans } + metadata_cache = ParquetMetadataCache(stats) # Generate the network state: GenState = { @@ -303,7 +305,7 @@ def generate_network( ir_context, scan, metadata_channel_by_scan[scan], - stats, + metadata_cache, ) for scan in metadata_scans ] diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py index 8adc5ef4bbff..35cb74fff9d2 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py @@ -95,6 +95,78 @@ class MetadataMessagePayload: cached_parquet_info: list[CachedParquetInfo] +class ParquetMetadataCache: + """ + Query-scoped cache for prefetched parquet metadata. + + Coordinates footer reads across concurrent metadata prefetch actors so each + distinct ``scan.paths`` tuple is fetched at most once per query/rank. + """ + + def __init__( + self, + stats: StatsCollector, + fetch: Callable[ + [list[str], StatsCollector], list[CachedParquetInfo] + ] = prefetch_cached_parquet_info_for_paths, + ) -> None: + self._stats = stats + self._cached_by_key: dict[tuple[str, ...], list[CachedParquetInfo]] = {} + self._pending_by_key: dict[ + tuple[str, ...], asyncio.Future[list[CachedParquetInfo]] + ] = {} + self._lock = asyncio.Lock() + self._fetch = fetch + + async def get( + self, + paths: list[str], + ir_context: IRExecutionContext, + ) -> list[CachedParquetInfo]: + """ + Return cached parquet metadata for ``paths``, fetching on first use. + + Concurrent callers with identical ``paths`` share a single in-flight fetch. + """ + key = tuple(paths) + async with self._lock: + if key in self._cached_by_key: + return self._cached_by_key[key] + if key in self._pending_by_key: + future = self._pending_by_key[key] + should_fetch = False + else: + loop = asyncio.get_running_loop() + future = loop.create_future() + self._pending_by_key[key] = future + should_fetch = True + + if should_fetch: + try: + result = await ir_context.to_thread( + self._fetch, + list(key), + self._stats, + ) + except BaseException as exc: + async with self._lock: + self._pending_by_key.pop(key, None) + if not future.done(): + if isinstance(exc, asyncio.CancelledError): + future.cancel() + else: + future.set_exception(exc) + raise + async with self._lock: + self._cached_by_key[key] = result + self._pending_by_key.pop(key) + if not future.done(): + future.set_result(result) + return result + + return await future + + class Lineariser: """ Linearizer that ensures ordered delivery from multiple concurrent producers. @@ -583,7 +655,7 @@ async def parquet_metadata_prefetch_node( ir_context: IRExecutionContext, ir: StreamingScan, ch_out: Channel[ArbitraryChunk[MetadataMessagePayload]], - stats: StatsCollector, + metadata_cache: ParquetMetadataCache, ) -> None: """ Fetch parquet metadata for each scan task and send it to the paired scan actor. @@ -601,8 +673,8 @@ async def parquet_metadata_prefetch_node( ch_out The output channel. The Scan actor generated for this StreamingScan node will read messages from this channel. - stats - The statistics collector, used to populate the parquet metadata cache. + metadata_cache + Shared query-scoped cache for prefetched parquet metadata. Notes ----- @@ -610,19 +682,13 @@ async def parquet_metadata_prefetch_node( The messages are sent in the order of the scans. """ async with shutdown_on_error(context, ch_out, trace_ir=ir, ir_context=ir_context): - cached_by_key: dict[tuple[str, ...], list[CachedParquetInfo]] = {} for scan in ir.scans: scan = cast("SplitScan | FusedScan", scan) key = tuple[str, ...](scan.paths) - if key not in cached_by_key: - cached_by_key[key] = await ir_context.to_thread( - prefetch_cached_parquet_info_for_paths, - list(key), - stats=stats, - ) + cached_parquet_info = await metadata_cache.get(list(key), ir_context) payload = MetadataMessagePayload( group_key=key, - cached_parquet_info=cached_by_key[key], + cached_parquet_info=cached_parquet_info, ) await ch_out.send_metadata( context, diff --git a/python/cudf_polars/tests/streaming/test_scan.py b/python/cudf_polars/tests/streaming/test_scan.py index 8d112704955c..213cbe45aabd 100644 --- a/python/cudf_polars/tests/streaming/test_scan.py +++ b/python/cudf_polars/tests/streaming/test_scan.py @@ -3,7 +3,9 @@ from __future__ import annotations +import asyncio import math +from concurrent.futures import ThreadPoolExecutor from typing import TYPE_CHECKING, cast import pytest @@ -27,6 +29,7 @@ from cudf_polars.engine.options import StreamingOptions from cudf_polars.streaming.actor_graph.io import ( MetadataMessagePayload, + ParquetMetadataCache, collect_metadata_scans, recv_prefetched_parquet_metadata_handler, ) @@ -863,3 +866,28 @@ def test_recv_prefetched_parquet_metadata_handler_errors() -> None: match=r"Unexpected parquet metadata key on scan input channel. .*", ): recv_prefetched_parquet_metadata_handler(msg, ("file2.parquet",)) + + +def test_parquet_metadata_cache_dedupes_identical_paths() -> None: + fetch_count = 0 + + def mock_fetch(paths: list[str], stats: StatsCollector) -> list[CachedParquetInfo]: + nonlocal fetch_count + fetch_count += 1 + return _make_cached_parquet_info(paths) + + cache = ParquetMetadataCache(StatsCollector(), fetch=mock_fetch) + paths = ["a.parquet", "b.parquet"] + + async def run() -> list[list[CachedParquetInfo]]: + with ThreadPoolExecutor(max_workers=2) as executor: + ir_context = IRExecutionContext(executor) + async with asyncio.TaskGroup() as tg: + first = tg.create_task(cache.get(paths, ir_context)) + second = tg.create_task(cache.get(paths, ir_context)) + return [first.result(), second.result()] + + results = asyncio.run(run()) + assert fetch_count == 1 + assert results[0] == results[1] + assert [info.path for info in results[0]] == paths From a14421eb329757cb6013fb9cec9ba05b1e3accb6 Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Fri, 10 Jul 2026 14:53:47 -0700 Subject: [PATCH 16/18] remove conditiona --- python/cudf_polars/cudf_polars/streaming/actor_graph/io.py | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py index 35cb74fff9d2..16ad5fb8c05e 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py @@ -1066,17 +1066,13 @@ def _( ) nodes[ir] = [native_node, metadata_node] else: - if ir.base_scan.typ == "parquet" and parquet_options.prefetch_file_metadata: - ch_metadata = rec.state["metadata_channel_by_scan"].get(ir) - else: - ch_metadata = None nodes[ir] = [ scan_node( rec.state["context"], ir, rec.state["ir_context"], ch_out, - ch_metadata, + rec.state["metadata_channel_by_scan"].get(ir), num_producers=num_producers, estimated_chunk_bytes=( plan.estimated_chunk_bytes or executor.target_partition_size From 33b6a72f8b36c8d6b9933702e83ea470c62c86d6 Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Fri, 10 Jul 2026 14:55:09 -0700 Subject: [PATCH 17/18] inline --- .../cudf_polars/streaming/actor_graph/io.py | 10 +--------- 1 file changed, 1 insertion(+), 9 deletions(-) diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py index 16ad5fb8c05e..8300be3a6ab8 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py @@ -714,7 +714,7 @@ def _start_metadata_receiver( async def _receive() -> None: try: for task_idx, scan in enumerate(scans): - msg = await recv_prefetched_parquet_metadata(context, ch_metadata) + msg = await ch_metadata.recv_metadata(context) cached = recv_prefetched_parquet_metadata_handler( msg, tuple(scan.paths) ) @@ -735,14 +735,6 @@ async def _receive() -> None: return futures, receiver_task -async def recv_prefetched_parquet_metadata( - context: Context, - ch: Channel[ArbitraryChunk[MetadataMessagePayload]], -) -> Message[ArbitraryChunk[MetadataMessagePayload]] | None: - """Receive and validate one prefetched parquet metadata message.""" - return await ch.recv_metadata(context) - - def recv_prefetched_parquet_metadata_handler( msg: Message[ArbitraryChunk[MetadataMessagePayload]] | None, group_key: tuple[str, ...], From 556ab550ae8c72056b45f5cfa06d42624cd81917 Mon Sep 17 00:00:00 2001 From: Tom Augspurger Date: Fri, 10 Jul 2026 14:56:12 -0700 Subject: [PATCH 18/18] remove pragmas --- python/cudf_polars/cudf_polars/streaming/actor_graph/io.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py index 8300be3a6ab8..82be404cecb6 100644 --- a/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py +++ b/python/cudf_polars/cudf_polars/streaming/actor_graph/io.py @@ -740,12 +740,12 @@ def recv_prefetched_parquet_metadata_handler( group_key: tuple[str, ...], ) -> list[CachedParquetInfo]: """Synchronous handler for prefetched parquet metadata messages.""" - if msg is None: # pragma: no cover; unreachable + if msg is None: raise AssertionError( f"Missing parquet metadata message for paths: {reprlib.repr(group_key)}" ) payload = ArbitraryChunk[MetadataMessagePayload].from_message(msg).release() - if payload.group_key != group_key: # pragma: no cover; unreachable + if payload.group_key != group_key: difference = set(group_key) ^ set(payload.group_key) raise AssertionError( "Unexpected parquet metadata key on scan input channel. "