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 @@ -63,6 +63,8 @@ import javax.ws.rs.core.UriBuilder

import java.util.Locale

import scala.collection.JavaConverters._

class VeloxSparkPlanExecApi extends SparkPlanExecApi {

/** Transform GetArrayItem to Substrait. */
Expand Down Expand Up @@ -670,26 +672,32 @@ class VeloxSparkPlanExecApi extends SparkPlanExecApi {
dataSize: SQLMetric): BuildSideRelation = {
val useOffheapBroadcastBuildRelation =
VeloxConfig.get.enableBroadcastBuildRelationInOffheap
val serialized: Array[ColumnarBatchSerializeResult] = child
val serialized: Seq[ColumnarBatchSerializeResult] = child
.executeColumnar()
.mapPartitions(itr => Iterator(BroadcastUtils.serializeStream(itr)))
.filter(_.getNumRows != 0)
.filter(_.numRows != 0)
.collect
val rawSize = serialized.flatMap(_.getSerialized.map(_.length.toLong)).sum
val rawSize = serialized.map(_.sizeInBytes()).sum
if (rawSize >= GlutenConfig.get.maxBroadcastTableSize) {
throw new SparkException(
"Cannot broadcast the table that is larger than " +
s"${SparkMemoryUtil.bytesToString(GlutenConfig.get.maxBroadcastTableSize)}: " +
s"${SparkMemoryUtil.bytesToString(rawSize)}")
}
numOutputRows += serialized.map(_.getNumRows).sum
numOutputRows += serialized.map(_.numRows).sum
dataSize += rawSize
if (useOffheapBroadcastBuildRelation) {
TaskResources.runUnsafe {
UnsafeColumnarBuildSideRelation(child.output, serialized.flatMap(_.getSerialized), mode)
UnsafeColumnarBuildSideRelation(
child.output,
serialized.flatMap(_.offHeapData().asScala),
mode)
}
} else {
ColumnarBuildSideRelation(child.output, serialized.flatMap(_.getSerialized), mode)
ColumnarBuildSideRelation(
child.output,
serialized.flatMap(_.onHeapData().asScala).toArray,
mode)
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@ import org.apache.spark.sql.types.StructType
import org.apache.spark.sql.vectorized.ColumnarBatch
import org.apache.spark.task.TaskResources

import scala.collection.JavaConverters._
import scala.collection.mutable.ArrayBuffer

// Utility methods to convert Vanilla broadcast relations from/to Velox broadcast relations.
Expand Down Expand Up @@ -93,58 +94,44 @@ object BroadcastUtils {
schema: StructType,
from: Broadcast[F],
fn: Iterator[InternalRow] => Iterator[ColumnarBatch]): Broadcast[T] = {
val useOffheapBuildRelation = VeloxConfig.get.enableBroadcastBuildRelationInOffheap

def batchIterationToRelation(batchItr: () => Iterator[ColumnarBatch]): BuildSideRelation = {
TaskResources.runUnsafe {
serializeStream(batchItr()) match {
case ColumnarBatchSerializeResult.EMPTY =>
ColumnarBuildSideRelation(
SparkShimLoader.getSparkShims.attributesFromStruct(schema),
Array[Array[Byte]](),
mode)
case result: ColumnarBatchSerializeResult =>
if (result.isOffHeap) {
UnsafeColumnarBuildSideRelation(
SparkShimLoader.getSparkShims.attributesFromStruct(schema),
result.offHeapData().asScala.toSeq,
mode)
} else {
ColumnarBuildSideRelation(
SparkShimLoader.getSparkShims.attributesFromStruct(schema),
result.onHeapData().asScala.toArray,
mode)
}
}
}
}

mode match {
case HashedRelationBroadcastMode(_, _) =>
// HashedRelation to ColumnarBuildSideRelation.
val fromBroadcast = from.asInstanceOf[Broadcast[HashedRelation]]
val fromRelation = fromBroadcast.value.asReadOnlyCopy()
val toRelation = TaskResources.runUnsafe {
val batchItr: Iterator[ColumnarBatch] = fn(reconstructRows(fromRelation))
val serialized: Array[Array[Byte]] = serializeStream(batchItr) match {
case ColumnarBatchSerializeResult.EMPTY =>
Array()
case result: ColumnarBatchSerializeResult =>
result.getSerialized
}
if (useOffheapBuildRelation) {
UnsafeColumnarBuildSideRelation(
SparkShimLoader.getSparkShims.attributesFromStruct(schema),
serialized,
mode)
} else {
ColumnarBuildSideRelation(
SparkShimLoader.getSparkShims.attributesFromStruct(schema),
serialized,
mode)
}
}
val toRelation = batchIterationToRelation(() => fn(reconstructRows(fromRelation)))
// Rebroadcast Velox relation.
context.broadcast(toRelation).asInstanceOf[Broadcast[T]]
case IdentityBroadcastMode =>
// Array[InternalRow] to ColumnarBuildSideRelation.
val fromBroadcast = from.asInstanceOf[Broadcast[Array[InternalRow]]]
val fromRelation = fromBroadcast.value
val toRelation = TaskResources.runUnsafe {
val batchItr: Iterator[ColumnarBatch] = fn(fromRelation.iterator)
val serialized: Array[Array[Byte]] = serializeStream(batchItr) match {
case ColumnarBatchSerializeResult.EMPTY =>
Array()
case result: ColumnarBatchSerializeResult =>
result.getSerialized
}
if (useOffheapBuildRelation) {
UnsafeColumnarBuildSideRelation(
SparkShimLoader.getSparkShims.attributesFromStruct(schema),
serialized,
mode)
} else {
ColumnarBuildSideRelation(
SparkShimLoader.getSparkShims.attributesFromStruct(schema),
serialized,
mode)
}
}
val toRelation = batchIterationToRelation(() => fn(fromRelation.iterator))
// Rebroadcast Velox relation.
context.broadcast(toRelation).asInstanceOf[Broadcast[T]]
case _ => throw new IllegalStateException("Unexpected broadcast mode: " + mode)
Expand Down Expand Up @@ -175,25 +162,25 @@ object BroadcastUtils {
val handle = ColumnarBatches.getNativeHandle(BackendsApiManager.getBackendName, b)
numRows += b.numRows()
try {
val unsafeBuffer = ColumnarBatchSerializerJniWrapper
ColumnarBatchSerializerJniWrapper
.create(
Runtimes
.contextInstance(
BackendsApiManager.getBackendName,
"BroadcastUtils#serializeStream"))
.serialize(handle)
try {
unsafeBuffer.toByteArray
} finally {
unsafeBuffer.close()
}
} finally {
ColumnarBatches.release(b)
}
})
.toArray
if (values.nonEmpty) {
new ColumnarBatchSerializeResult(numRows, values)
val useOffheapBroadcastBuildRelation =
VeloxConfig.get.enableBroadcastBuildRelationInOffheap
new ColumnarBatchSerializeResult(
useOffheapBroadcastBuildRelation,
numRows,
values.toSeq.asJava)
} else {
ColumnarBatchSerializeResult.EMPTY
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -178,12 +178,7 @@ class ColumnarCachedBatchSerializer extends CachedBatchSerializer with Logging {
BackendsApiManager.getBackendName,
"ColumnarCachedBatchSerializer#serialize"))
.serialize(ColumnarBatches.getNativeHandle(BackendsApiManager.getBackendName, batch))
val bytes =
try {
unsafeBuffer.toByteArray
} finally {
unsafeBuffer.close()
}
val bytes = unsafeBuffer.toByteArray
CachedColumnarBatch(batch.numRows(), bytes.length, bytes)
}
}
Expand Down

This file was deleted.

Loading