From 32425a473e8a0e9d446508bdaef474bfdb9052ae Mon Sep 17 00:00:00 2001 From: Matthew Roeschke <10647082+mroeschke@users.noreply.github.com> Date: Tue, 7 Jul 2026 21:58:45 +0000 Subject: [PATCH] Support pl.Expr.dt.quarter in cudf_polars --- .../cudf_polars/dsl/expressions/datetime.py | 11 ++++++++++ .../tests/expressions/test_datetime_basic.py | 21 +++++++++++++++++++ 2 files changed, 32 insertions(+) diff --git a/python/cudf_polars/cudf_polars/dsl/expressions/datetime.py b/python/cudf_polars/cudf_polars/dsl/expressions/datetime.py index b6b582738188..b7452411e5e7 100644 --- a/python/cudf_polars/cudf_polars/dsl/expressions/datetime.py +++ b/python/cudf_polars/cudf_polars/dsl/expressions/datetime.py @@ -145,6 +145,7 @@ def from_polars(cls, obj: polars._expr_nodes.TemporalFunction) -> Self: Name.TimeStamp, Name.CastTimeUnit, Name.Truncate, + Name.Quarter, *_TOTAL_COMPONENT_NANOSECONDS.keys(), } @@ -225,6 +226,16 @@ def do_evaluate( ), dtype=self.dtype, ) + elif self.name is TemporalFunction.Name.Quarter: + (column,) = columns + return Column( + plc.unary.cast( + plc.datetime.extract_quarter(column.obj, stream=df.stream), + self.dtype.plc_type, + stream=df.stream, + ), + dtype=self.dtype, + ) elif self.name is TemporalFunction.Name.CastTimeUnit: (column,) = columns return Column( diff --git a/python/cudf_polars/tests/expressions/test_datetime_basic.py b/python/cudf_polars/tests/expressions/test_datetime_basic.py index 26d4c01a4766..bdc473ed4d32 100644 --- a/python/cudf_polars/tests/expressions/test_datetime_basic.py +++ b/python/cudf_polars/tests/expressions/test_datetime_basic.py @@ -490,6 +490,27 @@ def test_datetime_truncate_unsupported(engine: pl.GPUEngine, every: str): assert_ir_translation_raises(q, engine, NotImplementedError) +@pytest.mark.parametrize( + "dtype", [pl.Date(), pl.Datetime("ms"), pl.Datetime("us"), pl.Datetime("ns")] +) +def test_datetime_quarter(engine: pl.GPUEngine, dtype): + data = pl.Series( + [ + datetime.date(2001, 1, 1), + datetime.date(2001, 3, 31), + datetime.date(2001, 4, 1), + datetime.date(2001, 6, 30), + datetime.date(2001, 9, 15), + datetime.date(2001, 12, 27), + None, + ], + dtype=pl.Date(), + ).cast(dtype) + ldf = pl.LazyFrame({"dates": data}) + q = ldf.select(pl.col("dates").dt.quarter()) + assert_gpu_result_equal(q, engine=engine) + + @pytest.mark.parametrize( "datetime_dtype", [