-
Notifications
You must be signed in to change notification settings - Fork 301
Add experimental Project AST JIT for integral add and multiply [fast-ut] [databricks] #15312
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from 22 commits
a3190f1
9bb4ff2
4f6972a
f75b0c1
b4a1182
6cea0cf
6ed9b58
f3ebe60
b242c73
8ffede9
8fc3fb5
ce448ef
a98a812
c20b430
3cc28a8
555dc68
8adcdcd
9772c07
fabcb01
93b73a8
116528e
59e26f2
1058db8
084c57f
c3a1f6d
3372a30
7dc0992
5b45ce1
9df526a
a0a891c
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,144 @@ | ||
| /* | ||
| * Copyright (c) 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. | ||
| * You may obtain a copy of the License at | ||
| * | ||
| * http://www.apache.org/licenses/LICENSE-2.0 | ||
| * | ||
| * Unless required by applicable law or agreed to in writing, software | ||
| * distributed under the License is distributed on an "AS IS" BASIS, | ||
| * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| * See the License for the specific language governing permissions and | ||
| * limitations under the License. | ||
| */ | ||
|
|
||
| package com.nvidia.spark.rapids | ||
|
|
||
| import ai.rapids.cudf.Table | ||
| import ai.rapids.cudf.ast.CompiledExpression | ||
| import com.nvidia.spark.Retryable | ||
| import com.nvidia.spark.rapids.Arm.{closeOnExcept, withResource} | ||
| import com.nvidia.spark.rapids.GpuMetric.OP_TIME_LEGACY | ||
| import com.nvidia.spark.rapids.RapidsPluginImplicits._ | ||
| import com.nvidia.spark.rapids.ScalableTaskCompletion.onTaskCompletion | ||
| import com.nvidia.spark.rapids.shims.ShimUnaryExpression | ||
|
|
||
| import org.apache.spark.TaskContext | ||
| import org.apache.spark.sql.catalyst.expressions.{Expression, NamedExpression} | ||
| import org.apache.spark.sql.types.DataType | ||
| import org.apache.spark.sql.vectorized.ColumnarBatch | ||
|
|
||
| object GpuAstJitExpression { | ||
| private def canUseAstJit(expression: GpuExpression): Boolean = | ||
| GpuBatchUtils.isFixedWidth(expression.dataType) && | ||
| expression.supportsAstJit && expression.containsAstJitOperator | ||
|
|
||
| private[rapids] def wrapTierExpression(expression: Expression): Expression = expression match { | ||
| case alias @ GpuAlias(_: GpuAstJitExpression, _) => alias | ||
| case alias @ GpuAlias(astExpression: GpuProjectAstExpression, _) | ||
| if canUseAstJit(astExpression.child) => | ||
| GpuProjectAstExpression.replaceChild(alias, GpuAstJitExpression(astExpression.child)) | ||
| case alias @ GpuAlias(child: GpuExpression, _) if canUseAstJit(child) => | ||
| GpuProjectAstExpression.replaceChild(alias, GpuAstJitExpression(child)) | ||
| case other => other | ||
| } | ||
|
igorpeshansky marked this conversation as resolved.
|
||
|
|
||
| private[rapids] def wrapProjectExpressions( | ||
| expressions: List[NamedExpression]): List[NamedExpression] = { | ||
| expressions.map(wrapTierExpression(_).asInstanceOf[NamedExpression]) | ||
| } | ||
|
|
||
| private[rapids] def contains(expression: Expression): Boolean = | ||
| GpuProjectAstExpressionBase.extractTopLevel(expression) | ||
| .exists(_.isInstanceOf[GpuAstJitExpression]) | ||
| } | ||
|
|
||
| case class GpuAstJitExpression(child: GpuExpression) | ||
|
igorpeshansky marked this conversation as resolved.
Outdated
igorpeshansky marked this conversation as resolved.
Outdated
|
||
| extends ShimUnaryExpression with GpuProjectAstExpressionBase | ||
| with GpuMetricsInjectable with Retryable with AutoCloseable { | ||
|
|
||
| @transient private[this] var compiledExpression: CompiledExpression = _ | ||
| @transient private[this] var completionRegistered = false | ||
| private[this] var opTime: GpuMetric = NoopMetric | ||
|
|
||
| override def dataType: DataType = child.dataType | ||
|
|
||
| override def nullable: Boolean = child.nullable | ||
|
|
||
| override def toString: String = s"AST_JIT($child)" | ||
|
|
||
| override def injectMetrics(metrics: Map[String, GpuMetric]): Unit = { | ||
| opTime = metrics.getOrElse(OP_TIME_LEGACY, NoopMetric) | ||
| } | ||
|
|
||
| override def checkpoint(): Unit = { | ||
| getCompiledExpression | ||
| } | ||
|
|
||
| // Compiled ASTs are immutable and remain valid across retry attempts. | ||
| override def restore(): Unit = () | ||
|
|
||
| override def close(): Unit = closeCompiledExpression() | ||
|
|
||
| override def columnarEval(batch: ColumnarBatch): GpuColumnVector = { | ||
| withResource(GpuProjectAstExpression.tableFromBatch(batch)) { table => | ||
| computeColumn(table) | ||
| } | ||
| } | ||
|
|
||
| private[rapids] override def computeColumn(table: Table): GpuColumnVector = { | ||
| val compiled = getCompiledExpression | ||
| NvtxIdWithMetrics(NvtxRegistry.PROJECT_AST, opTime) { | ||
| closeOnExcept(compiled.computeColumnJit(table)) { result => | ||
| GpuColumnVector.from(result, dataType) | ||
| } | ||
| } | ||
| } | ||
|
|
||
| private def getCompiledExpression: CompiledExpression = synchronized { | ||
| if (compiledExpression == null) { | ||
| compiledExpression = NvtxIdWithMetrics(NvtxRegistry.COMPILE_ASTS, opTime) { | ||
| // Force every bound reference to the left table; Project AST has one input table. | ||
| child.convertToAst(Int.MaxValue).compile() | ||
| } | ||
| } | ||
| if (!completionRegistered) { | ||
|
igorpeshansky marked this conversation as resolved.
Outdated
|
||
| Option(TaskContext.get()).foreach { taskContext => | ||
| completionRegistered = true | ||
| try { | ||
| onTaskCompletion(taskContext) { | ||
| clearRegistrationAndClose() | ||
| } | ||
| if (!completionRegistered) { | ||
| throw new IllegalStateException( | ||
| "Task completed while registering the AST JIT cleanup callback") | ||
| } | ||
| } catch { | ||
| case t: Throwable => | ||
| clearRegistrationAndClose(t) | ||
| throw t | ||
| } | ||
| } | ||
| } | ||
| compiledExpression | ||
| } | ||
|
|
||
| private def clearRegistrationAndClose(error: Throwable = null): Unit = | ||
| closeCompiledExpression(error, clearRegistration = true) | ||
|
|
||
| private def closeCompiledExpression( | ||
| error: Throwable = null, | ||
| clearRegistration: Boolean = false): Unit = { | ||
| val toClose = synchronized { | ||
| if (clearRegistration) { | ||
| completionRegistered = false | ||
| } | ||
| val current = compiledExpression | ||
| compiledExpression = null | ||
| current | ||
| } | ||
| Option(toClose).foreach(_.safeClose(error)) | ||
|
igorpeshansky marked this conversation as resolved.
Outdated
|
||
| } | ||
| } | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -124,18 +124,19 @@ object GpuBindReferences extends Logging { | |
| } | ||
|
|
||
| /** | ||
| * Binding method for tiered expressions without metric injection. | ||
| * This is for use by GpuBind implementations and should not be called directly | ||
| * from SparkPlan nodes. Use the public API that requires metrics instead, except | ||
| * when absolutely needed. | ||
| * Shared implementation for generic and Project-specific tiered binding. | ||
| * | ||
| * @param enableProjectAstJit whether eligible tiers may use Project AST JIT | ||
| */ | ||
| def bindGpuReferencesTieredNoMetrics[A <: Expression]( | ||
| private def bindGpuReferencesTieredNoMetricsInternal[A <: Expression]( | ||
| expressions: Seq[A], | ||
| input: AttributeSeq, | ||
| conf: SQLConf): GpuTieredProject = { | ||
| conf: SQLConf, | ||
| enableProjectAstJit: Boolean): GpuTieredProject = { | ||
|
|
||
| if (RapidsConf.ENABLE_TIERED_PROJECT.get(conf)) { | ||
| val exprTiers = GpuProjectAstExpression.buildExprTiers(expressions, conf) | ||
| val exprTiers = GpuProjectAstExpression.buildExprTiers( | ||
| expressions, conf, enableProjectAstJit) | ||
| 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 { | ||
|
|
@@ -174,10 +175,42 @@ object GpuBindReferences extends Logging { | |
| } | ||
| GpuTieredProject(tiered) | ||
| } else { | ||
| GpuTieredProject(Seq(GpuBindReferences.bindGpuReferencesNoMetrics(expressions, input))) | ||
| val projectExpressions = if (enableProjectAstJit) { | ||
| expressions.map(GpuAstJitExpression.wrapTierExpression) | ||
| } else { | ||
| expressions | ||
| } | ||
| GpuTieredProject(Seq( | ||
| GpuBindReferences.bindGpuReferencesNoMetrics(projectExpressions, input))) | ||
| } | ||
| } | ||
|
|
||
| /** | ||
| * Binding method for tiered expressions without metric injection. | ||
| * This is for use by GpuBind implementations and should not be called directly | ||
| * from SparkPlan nodes. Use the public API that requires metrics instead, except | ||
| * when absolutely needed. | ||
| */ | ||
| def bindGpuReferencesTieredNoMetrics[A <: Expression]( | ||
|
igorpeshansky marked this conversation as resolved.
|
||
| expressions: Seq[A], | ||
| input: AttributeSeq, | ||
| conf: SQLConf): GpuTieredProject = { | ||
| bindGpuReferencesTieredNoMetricsInternal( | ||
| expressions, input, conf, enableProjectAstJit = false) | ||
| } | ||
|
|
||
| /** | ||
| * Project-specific tiered binding without metric injection. Unlike the generic binder, | ||
| * this path allows configured Project AST JIT selection. | ||
| */ | ||
| private[rapids] def bindGpuProjectReferencesTieredNoMetrics[A <: Expression]( | ||
|
igorpeshansky marked this conversation as resolved.
Outdated
|
||
| expressions: Seq[A], | ||
| input: AttributeSeq, | ||
| conf: SQLConf): GpuTieredProject = { | ||
| bindGpuReferencesTieredNoMetricsInternal( | ||
| expressions, input, conf, RapidsConf.ENABLE_PROJECT_AST_JIT.get(conf)) | ||
| } | ||
|
|
||
| // ========== Public "Front Door" APIs (for use by SparkPlan nodes) ========== | ||
| // These methods require metrics and inject them after binding | ||
|
|
||
|
|
@@ -257,12 +290,32 @@ object GpuBindReferences extends Logging { | |
| bound.injectMetrics(metrics) | ||
| bound | ||
| } | ||
|
|
||
| /** | ||
| * Bind Project expressions in a tiered manner and inject metrics. Project AST JIT selection is | ||
| * confined to this entry point so generic tiered binders do not enable it for other operators. | ||
| * @param expressions The expressions to bind | ||
| * @param input The input schema | ||
| * @param conf SQL configuration | ||
| * @param metrics Metrics to inject into the bound expressions | ||
| */ | ||
| def bindGpuProjectReferencesTiered[A <: Expression]( | ||
|
igorpeshansky marked this conversation as resolved.
Outdated
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I personally don't like the name. This just like If there are good reasons to keep them separate, can we rename this or modify the original API to take in the JIT/AST enable param? To me that is much cleaner and less confusing. |
||
| expressions: Seq[A], | ||
| input: AttributeSeq, | ||
| conf: SQLConf, | ||
| metrics: Map[String, GpuMetric]): GpuTieredProject = { | ||
| val bound = bindGpuProjectReferencesTieredNoMetrics(expressions, input, conf) | ||
| bound.injectMetrics(metrics) | ||
| bound | ||
| } | ||
| } | ||
|
|
||
| case class GpuBoundReference(ordinal: Int, dataType: DataType, nullable: Boolean) | ||
| (val exprId: ExprId, val name: String) | ||
| extends GpuLeafExpression with ShimExpression { | ||
|
|
||
| override def selfSupportsAstJit: Boolean = true | ||
|
|
||
| override def toString: String = | ||
| s"input[$ordinal, ${dataType.simpleString}, $nullable]($name#${exprId.id})" | ||
|
|
||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -197,6 +197,31 @@ trait GpuExpression extends Expression { | |
| def convertToAst(numFirstTableColumns: Int): ast.AstExpression = | ||
| throw new IllegalStateException(s"Cannot convert ${this.getClass.getSimpleName} to AST") | ||
|
|
||
| /** | ||
| * Whether this node supports AST JIT for its current semantics and types, excluding its | ||
| * children. Operator overrides must validate their execution modes and local input/output types. | ||
|
igorpeshansky marked this conversation as resolved.
Outdated
|
||
| */ | ||
| def selfSupportsAstJit: Boolean = false | ||
|
|
||
| /** | ||
| * Whether this node is an operation, rather than an AST-compatible leaf. Literals and references | ||
| * leave this false so they do not trigger compilation without useful work to JIT. | ||
| */ | ||
| def selfIsAstJitOperator: Boolean = false | ||
|
igorpeshansky marked this conversation as resolved.
|
||
|
|
||
| /** Whether this node and its complete expression subtree support AST JIT. */ | ||
| final def supportsAstJit: Boolean = selfSupportsAstJit && children.forall { | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I get that you are being conservative right now. But I am concerned that this is adding in a lot of code that we are going to have to rip out when we actually do it right. Currently If I have an expression tree like
Collaborator
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Totally agreed. If we could partially enable the AST JIT inside the expression, that would be better. Also, since the multi-output and CSE NVIDIA/cudf#23621 have been merged, we might need to adjust some design decisions here. I'm converting this to a draft now to test more solutions... |
||
| case child: GpuExpression => child.supportsAstJit | ||
| case _: AttributeReference => true | ||
| case _ => false | ||
| } | ||
|
|
||
| /** Whether this expression subtree contains an operation that makes AST JIT useful. */ | ||
| final def containsAstJitOperator: Boolean = selfIsAstJitOperator || children.exists { | ||
| case child: GpuExpression => child.containsAstJitOperator | ||
| case _ => false | ||
| } | ||
|
|
||
| /** Could evaluating this expression cause side-effects, such as throwing an exception? */ | ||
| def hasSideEffects: Boolean = | ||
| children.exists { | ||
|
|
@@ -391,6 +416,8 @@ trait CudfBinaryExpression extends GpuBinaryExpression { | |
| def castOutputAtEnd: Boolean = false | ||
| def astOperator: Option[ast.BinaryOperator] = None | ||
|
|
||
| override def selfIsAstJitOperator: Boolean = selfSupportsAstJit | ||
|
|
||
| def outputType(l: BinaryOperable, r: BinaryOperable): DType = { | ||
| val over = outputTypeOverride | ||
| if (over == null) { | ||
|
|
||
Uh oh!
There was an error while loading. Please reload this page.