Skip to content
Merged
Show file tree
Hide file tree
Changes from 6 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
100 changes: 64 additions & 36 deletions integration_tests/src/main/python/ast_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,18 +68,22 @@
_project_ast_enabled_conf = {"spark.rapids.sql.projectAstEnabled": "true"}

def assert_gpu_ast(is_supported, func, conf={}):
exist = "GpuProjectAstExec"
non_exist = "GpuProjectExec"
ast_expression = "GpuProjectAstExpression"
exist = ast_expression
non_exist = ''
if not is_supported:
exist = "GpuProjectExec"
non_exist = "GpuProjectAstExec"
non_exist = ast_expression
ast_conf = copy_and_update(conf, _project_ast_enabled_conf)
assert_cpu_and_gpu_are_equal_collect_with_capture(
func,
exist_classes=exist,
non_exist_classes=non_exist,
conf=ast_conf)

def assert_gpu_project_without_ast(func, conf={}):
assert_gpu_ast(False, func, conf)

def assert_unary_ast(data_descr, func, conf={}):
(data_gen, is_supported) = data_descr
assert_gpu_ast(is_supported, lambda spark: func(unary_op_df(spark, data_gen)), conf=conf)
Expand All @@ -94,17 +98,17 @@ def test_literal(spark_tmp_path, data_gen):
data_path = spark_tmp_path + '/AST_TEST_DATA'
with_cpu_session(lambda spark: gen_df(spark, [("a", IntegerGen())]).write.parquet(data_path))
scalar = with_cpu_session(lambda spark: gen_scalar(data_gen, force_no_nulls=True))
assert_gpu_ast(is_supported=True,
func=lambda spark: spark.read.parquet(data_path).select(scalar))
assert_gpu_project_without_ast(
func=lambda spark: spark.read.parquet(data_path).select(scalar))

@pytest.mark.parametrize('data_gen', [boolean_gen, byte_gen, short_gen, int_gen, long_gen, float_gen, double_gen, timestamp_gen, date_gen], ids=idfn)
def test_null_literal(spark_tmp_path, data_gen):
# Write data to Parquet so Spark generates a plan using just the count of the data.
data_path = spark_tmp_path + '/AST_TEST_DATA'
with_cpu_session(lambda spark: gen_df(spark, [("a", IntegerGen())]).write.parquet(data_path))
data_type = data_gen.data_type
assert_gpu_ast(is_supported=True,
func=lambda spark: spark.read.parquet(data_path).select(f.lit(None).cast(data_type)))
assert_gpu_project_without_ast(
func=lambda spark: spark.read.parquet(data_path).select(f.lit(None).cast(data_type)))

@pytest.mark.parametrize('data_descr', ast_descrs, ids=idfn)
def test_isnull(data_descr):
Expand All @@ -118,22 +122,16 @@ def test_isnotnull(data_descr):
def test_bitwise_not(data_descr):
assert_unary_ast(data_descr, lambda df: df.selectExpr('~a'))

# This just ends up being a pass through. There is no good way to force
# a unary positive into a plan, because it gets optimized out, but this
# verifies that we can handle it.
@pytest.mark.parametrize('data_descr', [
(byte_gen, True),
(short_gen, True),
(int_gen, True),
(long_gen, True),
(float_gen, True),
(double_gen, True)], ids=idfn)
def test_unary_positive(data_descr):
assert_unary_ast(data_descr, lambda df: df.selectExpr('+a'))
# Unary positive is optimized to a pass-through, so per-expression AST has nothing to compile.
@pytest.mark.parametrize(
'data_gen', [byte_gen, short_gen, int_gen, long_gen, float_gen, double_gen], ids=idfn)
def test_unary_positive(data_gen):
assert_gpu_project_without_ast(
lambda spark: unary_op_df(spark, data_gen).selectExpr('+a'))

def test_unary_positive_for_daytime_interval():
data_descr = (DayTimeIntervalGen(), True)
assert_unary_ast(data_descr, lambda df: df.selectExpr('+a'))
assert_gpu_project_without_ast(
lambda spark: unary_op_df(spark, DayTimeIntervalGen()).selectExpr('+a'))

@pytest.mark.parametrize('data_descr', ast_arithmetic_descrs, ids=idfn)
@disable_ansi_mode
Expand Down Expand Up @@ -260,9 +258,18 @@ def test_exp(data_descr):
def test_expm1(data_descr):
assert_unary_ast(data_descr, lambda df: df.selectExpr('expm1(a)'))

@pytest.mark.parametrize('data_gen', [float_gen, double_gen], ids=idfn)
def test_folded_null_literal_stays_on_regular_project(data_gen):
assert_gpu_project_without_ast(
lambda spark: binary_op_df(spark, data_gen).select(
f.col('a') == f.lit(None).cast(data_gen.data_type),
f.col('a') == f.col('b')))

# Keep null scalars from folding unsupported comparisons into AST-compatible null literals.
@pytest.mark.parametrize('data_descr', ast_comparable_descrs, ids=idfn)
def test_eq(data_descr):
(s1, s2) = with_cpu_session(lambda spark: gen_scalars(data_descr[0], 2))
(s1, s2) = with_cpu_session(
lambda spark: gen_scalars(data_descr[0], 2, force_no_nulls=True))
Comment thread
igorpeshansky marked this conversation as resolved.
assert_binary_ast(data_descr,
lambda df: df.select(
f.col('a') == s1,
Expand All @@ -271,7 +278,8 @@ def test_eq(data_descr):

@pytest.mark.parametrize('data_descr', ast_comparable_descrs, ids=idfn)
def test_ne(data_descr):
(s1, s2) = with_cpu_session(lambda spark: gen_scalars(data_descr[0], 2))
(s1, s2) = with_cpu_session(
lambda spark: gen_scalars(data_descr[0], 2, force_no_nulls=True))
assert_binary_ast(data_descr,
lambda df: df.select(
f.col('a') != s1,
Expand All @@ -280,7 +288,8 @@ def test_ne(data_descr):

@pytest.mark.parametrize('data_descr', ast_comparable_descrs, ids=idfn)
def test_lt(data_descr):
(s1, s2) = with_cpu_session(lambda spark: gen_scalars(data_descr[0], 2))
(s1, s2) = with_cpu_session(
lambda spark: gen_scalars(data_descr[0], 2, force_no_nulls=True))
assert_binary_ast(data_descr,
lambda df: df.select(
f.col('a') < s1,
Expand All @@ -289,7 +298,8 @@ def test_lt(data_descr):

@pytest.mark.parametrize('data_descr', ast_comparable_descrs, ids=idfn)
def test_lte(data_descr):
(s1, s2) = with_cpu_session(lambda spark: gen_scalars(data_descr[0], 2))
(s1, s2) = with_cpu_session(
lambda spark: gen_scalars(data_descr[0], 2, force_no_nulls=True))
assert_binary_ast(data_descr,
lambda df: df.select(
f.col('a') <= s1,
Expand All @@ -298,7 +308,8 @@ def test_lte(data_descr):

@pytest.mark.parametrize('data_descr', ast_comparable_descrs, ids=idfn)
def test_gt(data_descr):
(s1, s2) = with_cpu_session(lambda spark: gen_scalars(data_descr[0], 2))
(s1, s2) = with_cpu_session(
lambda spark: gen_scalars(data_descr[0], 2, force_no_nulls=True))
assert_binary_ast(data_descr,
lambda df: df.select(
f.col('a') > s1,
Expand All @@ -307,7 +318,8 @@ def test_gt(data_descr):

@pytest.mark.parametrize('data_descr', ast_comparable_descrs, ids=idfn)
def test_gte(data_descr):
(s1, s2) = with_cpu_session(lambda spark: gen_scalars(data_descr[0], 2))
(s1, s2) = with_cpu_session(
lambda spark: gen_scalars(data_descr[0], 2, force_no_nulls=True))
assert_binary_ast(data_descr,
lambda df: df.select(
f.col('a') >= s1,
Expand Down Expand Up @@ -428,14 +440,30 @@ def test_multi_tier_ast():
func=lambda spark: spark.range(10).withColumn("x", f.col("id")).repartition(1)\
.selectExpr("x", "(id < x) == (id < (id + x))"))


# MUST NOT use GPU AST when project refers to string type(non-fixed-width),
# or cudf::compute_column will throw error: Invalid, non-fixed-width type
# ANSI mode is disabled here due to an overflow issue with integer multiplication on Spark 4.0.0.
Comment thread
igorpeshansky marked this conversation as resolved.
@disable_ansi_mode
@ignore_order(local=True)
def test_refer_to_non_fixed_width_column():
gens = [('col_int', int_gen), ('col_string', string_gen)]
assert_gpu_and_cpu_are_equal_collect(
lambda spark: gen_df(spark, gens).selectExpr("col_int * col_int", "col_string"),
conf=_project_ast_enabled_conf)
@pytest.mark.parametrize(
'tiered_project_enabled', ['true', 'false'], ids=['tiered', 'single_tier'])
def test_project_ast_mixed_expressions(tiered_project_enabled):
def project(spark):
df = gen_df(spark, [
('a', int_gen),
('b', int_gen),
('c', int_gen),
('d', int_gen),
('col_string', string_gen)
])
shared = f.col('a') + f.col('b')
return df.select(
(shared * f.col('c')).alias('ast_first'),
(shared * f.col('d')).alias('ast_second'),
f.greatest(shared, f.col('c')).alias('gpu_shared'),
f.length(f.col('col_string')).alias('gpu_string'),
f.col('col_string').alias('raw_string'))

assert_gpu_ast(
is_supported=True,
func=project,
conf={
'spark.rapids.sql.tiered.project.enabled': tiered_project_enabled
})
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
/*
* Copyright (c) 2019-2025, NVIDIA CORPORATION.
* Copyright (c) 2019-2026, NVIDIA CORPORATION.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
Expand Down Expand Up @@ -135,12 +135,7 @@ object GpuBindReferences extends Logging {
conf: SQLConf): GpuTieredProject = {

if (RapidsConf.ENABLE_TIERED_PROJECT.get(conf)) {
val replaced = if (RapidsConf.ENABLE_COMBINED_EXPRESSIONS.get(conf)) {
GpuEquivalentExpressions.replaceMultiExpressions(expressions, conf)
} else {
expressions
}
val exprTiers = GpuEquivalentExpressions.getExprTiers(replaced)
val exprTiers = GpuProjectAstExpression.buildExprTiers(expressions, conf)
val inputTiers = GpuEquivalentExpressions.getInputTiers(exprTiers, input)
// Update ExprTiers to include the columns that are pass through and drop unneeded columns
val newExprTiers = exprTiers.zipWithIndex.map {
Expand Down
Loading
Loading