diff --git a/python/cudf_polars/cudf_polars/dsl/expressions/datetime.py b/python/cudf_polars/cudf_polars/dsl/expressions/datetime.py index c02f40a88dcb..1707603b7741 100644 --- a/python/cudf_polars/cudf_polars/dsl/expressions/datetime.py +++ b/python/cudf_polars/cudf_polars/dsl/expressions/datetime.py @@ -151,6 +151,7 @@ def from_polars(cls, obj: polars._expr_nodes.TemporalFunction) -> Self: Name.TimeStamp, Name.CastTimeUnit, Name.Truncate, + Name.Date, Name.DaysInMonth, *_CENTURY_MILLENNIUM_DIVISOR.keys(), *_TOTAL_COMPONENT_NANOSECONDS.keys(), @@ -233,6 +234,14 @@ def do_evaluate( ), dtype=self.dtype, ) + elif self.name is TemporalFunction.Name.Date: + (column,) = columns + # Casting the timestamp to TIMESTAMP_DAYS (the storage of ``pl.Date``) + # drops the sub-day component. + return Column( + plc.unary.cast(column.obj, self.dtype.plc_type, stream=df.stream), + dtype=self.dtype, + ) elif self.name is TemporalFunction.Name.DaysInMonth: (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 b34d4693b552..a0a5d3a1a49e 100644 --- a/python/cudf_polars/tests/expressions/test_datetime_basic.py +++ b/python/cudf_polars/tests/expressions/test_datetime_basic.py @@ -252,6 +252,25 @@ def test_century_millennium_date_extreme_years(engine: pl.GPUEngine, method, day assert_gpu_result_equal(q, engine=engine) +@pytest.mark.parametrize( + "dtype", [pl.Datetime("ms"), pl.Datetime("us"), pl.Datetime("ns")] +) +def test_datetime_date(engine: pl.GPUEngine, dtype): + data = pl.Series( + [ + datetime.datetime(1978, 1, 1, 1, 1, 1), + datetime.datetime(1969, 12, 31, 23, 59, 59), # pre-epoch (floors down) + datetime.datetime(2024, 10, 13, 5, 30, 14, 500_000), + datetime.datetime(2065, 1, 1, 10, 20, 30, 60_000), + None, + ], + dtype=dtype, + ) + ldf = pl.LazyFrame({"datetimes": data}) + q = ldf.select(pl.col("datetimes").dt.date()) + assert_gpu_result_equal(q, engine=engine) + + @pytest.mark.parametrize( "dtype", [pl.Date(), pl.Datetime("ms"), pl.Datetime("us"), pl.Datetime("ns")] )