Skip to content
Merged
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
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@

package org.apache.comet.codegen

import org.apache.arrow.memory.BufferAllocator
import org.apache.arrow.vector._
import org.apache.arrow.vector.complex.{ListVector, MapVector, StructVector}
import org.apache.arrow.vector.types.pojo.Field
Expand All @@ -28,6 +29,7 @@ import org.apache.spark.sql.catalyst.expressions.codegen._
import org.apache.spark.sql.internal.SQLConf
import org.apache.spark.sql.types._

import org.apache.comet.CometArrowAllocator
import org.apache.comet.shims.{CometExprTraitShim, CometTypeShim}

/**
Expand Down Expand Up @@ -205,8 +207,12 @@ object CometBatchKernelCodegen extends Logging with CometExprTraitShim with Come
* Allocate an Arrow output vector from a pre-built `Field`. Forwards to
* [[CometBatchKernelCodegenOutput.allocateOutput]].
*/
def allocateOutput(field: Field, numRows: Int, estimatedBytes: Int): FieldVector =
CometBatchKernelCodegenOutput.allocateOutput(field, numRows, estimatedBytes)
def allocateOutput(
field: Field,
numRows: Int,
estimatedBytes: Int,
allocator: BufferAllocator = CometArrowAllocator): FieldVector =
CometBatchKernelCodegenOutput.allocateOutput(field, numRows, estimatedBytes, allocator)

/**
* Spark `DataType` to an Arrow `Field`, resolving mismatches between Arrow Java's default field
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,6 @@ import org.apache.spark.sql.catalyst.expressions.codegen.CodegenContext
import org.apache.spark.sql.comet.util.Utils
import org.apache.spark.sql.types._

import org.apache.comet.CometArrowAllocator
import org.apache.comet.shims.CometTypeShim

/**
Expand Down Expand Up @@ -86,22 +85,26 @@ private[codegen] object CometBatchKernelCodegenOutput extends CometTypeShim {
*
* Closes the vector on any failure so a partially-initialized tree doesn't leak buffers.
*/
def allocateOutput(field: Field, numRows: Int, estimatedBytes: Int): FieldVector = {
def allocateOutput(
field: Field,
numRows: Int,
estimatedBytes: Int,
allocator: BufferAllocator): FieldVector = {
val vec: FieldVector = field.getType match {
case _: ArrowType.List | _: ArrowType.LargeList | _: ArrowType.FixedSizeList =>
val v = new RenamedListVector(field, CometArrowAllocator)
val v = new RenamedListVector(field, allocator)
v.initializeChildrenFromFields(field.getChildren)
v
case _: ArrowType.Map =>
val v = new RenamedMapVector(field, CometArrowAllocator)
val v = new RenamedMapVector(field, allocator)
v.initializeChildrenFromFields(field.getChildren)
v
case _: ArrowType.Struct =>
val v = new RenamedStructVector(field, CometArrowAllocator)
val v = new RenamedStructVector(field, allocator)
v.initializeChildrenFromFields(field.getChildren)
v
case _ =>
field.createVector(CometArrowAllocator).asInstanceOf[FieldVector]
field.createVector(allocator).asInstanceOf[FieldVector]
}
try {
vec.setInitialCapacity(numRows)
Expand Down Expand Up @@ -138,9 +141,21 @@ private[codegen] object CometBatchKernelCodegenOutput extends CometTypeShim {
override def getField: Field = exportField
}

/**
* StructVector gets a field without children, so its writer creates no children that
* initializeChildrenFromFields then drops. `getField` returns `exportField` after that call.
*/
private final class RenamedStructVector(exportField: Field, allocator: BufferAllocator)
extends StructVector(exportField, allocator, null) {
override def getField: Field = exportField
extends StructVector(exportField.getName, allocator, exportField.getFieldType, null) {
// False while the StructVector constructor runs.
private var childrenInitialized = false

override def initializeChildrenFromFields(children: java.util.List[Field]): Unit = {
super.initializeChildrenFromFields(children)
childrenInitialized = true
}

override def getField: Field = if (childrenInitialized) exportField else super.getField
}

/**
Expand Down
32 changes: 32 additions & 0 deletions spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -263,6 +263,38 @@ class CometCodegenSuite
}
}

test("a closed allocateOutput vector releases all of its memory") {
// Struct outputs leaked the children that StructVector's writer allocates in its
// constructor. The List and Map cases make sure that these outputs do not start to leak.
val pair = StructType(
Seq(StructField("name", StringType), StructField("age", IntegerType, nullable = false)))
val outputTypes = Seq(
pair,
StructType(
Seq(StructField("_1", LongType, nullable = false), StructField("_2", StringType))),
StructType(
Seq(
StructField("inner", pair),
StructField("tags", ArrayType(StringType)),
StructField("attrs", MapType(StringType, IntegerType)))),
ArrayType(pair),
MapType(StringType, pair),
StringType)
outputTypes.foreach { dataType =>
val field = CometBatchKernelCodegen.toFfiArrowField("out", dataType, nullable = true)
val allocator =
CometArrowAllocator.newChildAllocator(s"allocateOutput($dataType)", 0, Long.MaxValue)
try {
CometBatchKernelCodegen.allocateOutput(field, 4, 0, allocator).close()
assert(
allocator.getAllocatedMemory == 0,
s"the $dataType output did not release all of its memory")
} finally {
allocator.close()
}
}
}

test("ScalaUDF over concat(c1, c2) suppresses the null short-circuit") {
// Concat is not NullIntolerant. The dispatcher's short-circuit guard inspects every node in
// the bound tree and must skip the whole-tree null short-circuit because one child is
Expand Down
Loading