Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
74 changes: 43 additions & 31 deletions python/cudf_polars/cudf_polars/engine/options.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,14 +18,18 @@
from cudf_polars.engine.hardware_binding import (
HardwareBindingPolicy,
)
from cudf_polars.utils.config import UNSPECIFIED, MemoryResourceConfig, Unspecified
from cudf_polars.utils.config import (
UNSPECIFIED,
DynamicPlanningOptions,
MemoryResourceConfig,
Unspecified,
)

if TYPE_CHECKING:
from collections.abc import Callable

from cudf_polars.quent import QuentContext
from cudf_polars.utils.config import (
DynamicPlanningOptions,
JoinFilterPushdownOptions,
ParquetOptions,
)
Expand Down Expand Up @@ -467,7 +471,10 @@ def _from_argparse(cls, args: argparse.Namespace) -> StreamingOptions:

Designed to work with namespaces produced by :func:`argparse.ArgumentParser`
parsers that have been augmented with :meth:`_add_cli_args`. Fields not present
in the namespace (or set to ``None``) remain :data:`UNSPECIFIED`.
in the namespace (or set to ``None``) let their default factories read the
environment variable. Fields with no explicit value or environment variable
remain :data:`UNSPECIFIED`; their downstream consumer applies its built-in
default.

Parameters
----------
Expand All @@ -490,12 +497,19 @@ def _get(attr: str) -> Any:
v = getattr(args, attr, None)
return UNSPECIFIED if v is None else v

# Special: dynamic_planning bool → None (disabled) or UNSPECIFIED
# True (the build_parser default) → UNSPECIFIED (use library default)
# False → explicitly disable (None)
# absent / None → UNSPECIFIED
# dynamic_planning is a BooleanOptionalAction with default=None, so
# dyn is True only when --dynamic-planning was passed explicitly,
# False only for --no-dynamic-planning, and None when absent.
# An explicit True must be preserved (not UNSPECIFIED), otherwise an
# env var disabling dynamic planning would silently override it.
dyn = getattr(args, "dynamic_planning", None)
dynamic_planning: Any = None if dyn is False else UNSPECIFIED
dynamic_planning: Any = (
UNSPECIFIED
if dyn is None
else None
if dyn is False
else DynamicPlanningOptions()
)

# target_partition_size: canonical dest from _add_cli_args; fall back to
# "blocksize" for legacy benchmark scripts that predate this module.
Expand All @@ -505,30 +519,28 @@ def _get(attr: str) -> Any:
else _get("blocksize")
)

# A field's CLI dest is usually its own name, except where a bare name
# would be ambiguous on the command line (e.g. "log" -> "--log" reads as
# a generic logging flag, so its dest is namespaced to "rapidsmpf_log").
cli_dest_overrides = {
"log": "rapidsmpf_log",
"statistics": "rapidsmpf_statistics",
}
special_cased = {
"target_partition_size": target_partition_size,
"dynamic_planning": dynamic_planning,
}
kwargs: dict[str, Any] = {
f.name: special_cased.get(
f.name, _get(cli_dest_overrides.get(f.name, f.name))
)
for f in dataclasses.fields(cls)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
}
# Omit UNSPECIFIED entries rather than passing them explicitly, so
# each field's own default_factory (env var, then built-in default)
# runs instead of being short-circuited.
return cls(
num_streaming_threads=_get("num_streaming_threads"),
num_streams=_get("num_streams"),
log=_get("rapidsmpf_log"), # renamed: dest rapidsmpf_log → log
statistics=_get("rapidsmpf_statistics"), # renamed
memory_reserve_timeout=_get("memory_reserve_timeout"),
allow_overbooking_by_default=_get("allow_overbooking_by_default"),
pinned_memory=_get("pinned_memory"),
pinned_initial_pool_size=_get("pinned_initial_pool_size"),
pinned_max_pool_size=_get("pinned_max_pool_size"),
spill_device_limit=_get("spill_device_limit"),
periodic_spill_check=_get("periodic_spill_check"),
unbounded_file_read_cache=_get("unbounded_file_read_cache"),
hardware_binding=_get("hardware_binding"),
num_py_executors=_get("num_py_executors"),
max_concurrent_io_tasks=_get("max_concurrent_io_tasks"),
fallback_mode=_get("fallback_mode"),
max_rows_per_partition=_get("max_rows_per_partition"),
broadcast_limit=_get("broadcast_limit"),
target_partition_size=target_partition_size,
dynamic_planning=dynamic_planning,
raise_on_fail=_get("raise_on_fail"),
parquet_options=_get("parquet_options"),
memory_resource_config=_get("memory_resource_config"),
**{k: v for k, v in kwargs.items() if not isinstance(v, Unspecified)}
)

@staticmethod
Expand Down
26 changes: 24 additions & 2 deletions python/cudf_polars/tests/streaming/test_options.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@
Unspecified,
_parse_memory_resource_config,
)
from cudf_polars.utils.config import MemoryResourceConfig
from cudf_polars.utils.config import DynamicPlanningOptions, MemoryResourceConfig

# ---------------------------------------------------------------------------
# Sentinel
Expand Down Expand Up @@ -312,10 +312,16 @@ def test_from_argparse_renames() -> None:
def test_from_argparse_dynamic_planning() -> None:
assert isinstance(
StreamingOptions._from_argparse(
argparse.Namespace(dynamic_planning=True)
argparse.Namespace(dynamic_planning=None)
).dynamic_planning,
Unspecified,
)
assert (
StreamingOptions._from_argparse(
argparse.Namespace(dynamic_planning=True)
).dynamic_planning
== DynamicPlanningOptions()
)
assert (
StreamingOptions._from_argparse(
argparse.Namespace(dynamic_planning=False)
Expand Down Expand Up @@ -358,6 +364,22 @@ def test_add_cli_args_then_from_argparse_roundtrip() -> None:
assert isinstance(opts.fallback_mode, Unspecified)


def test_from_argparse_omitted_flag_still_picks_up_env_var(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("CUDF_POLARS__EXECUTOR__NUM_PY_EXECUTORS", "16")
monkeypatch.setenv("RAPIDSMPF_NUM_STREAMING_THREADS", "16")

parser = argparse.ArgumentParser()
StreamingOptions._add_cli_args(parser)
args = parser.parse_args([])
opts = StreamingOptions._from_argparse(args)

assert opts.num_py_executors == 16
assert opts.to_executor_options()["num_py_executors"] == 16
assert opts.num_streaming_threads == 16


# ---------------------------------------------------------------------------
# to_dict / roundtrip
# ---------------------------------------------------------------------------
Expand Down
Loading