comphead commented on code in PR #6361:
URL: https://github.com/apache/datafusion-comet/pull/6361#discussion_r4146854810
##########
spark/src/test/scala/org/apache/comet/exec/CometColumnarShuffleSuite.scala:
##########
@@ -959,6 +959,34 @@ class CometShuffleEncryptionSuite extends CometTestBase {
}
}
+class CometShuffleChecksumDisabledSuite extends CometTestBase {
+
+ override protected def sparkConf: SparkConf = {
+ val conf = super.sparkConf
+ conf.set("spark.shuffle.checksum.enabled", "false")
+ }
+
+ test("comet columnar shuffle with shuffle checksums disabled") {
+ // 10 partitions use the bypass merge sort writer, while 300 partitions
exceed
+ // spark.shuffle.sort.bypassMergeThreshold and use the sort-based writer.
The lower spill
+ // threshold makes every map task spill, so the sort-based writer also
merges spill files.
+ Seq(10, 300).foreach { numPartitions =>
+ Seq(Int.MaxValue, 2000).foreach { spillThreshold =>
+ withSQLConf(
+ CometConf.COMET_SHUFFLE_MODE.key -> "jvm",
+ CometConf.COMET_SHUFFLE_JVM_SPILL_THRESHOLD.key ->
spillThreshold.toString) {
+ val df = spark
+ .range(0, 100000, 1, 4)
+ .selectExpr("id", "cast(id as string) as s")
+ .repartition(numPartitions, col("id"))
+ checkCometExchange(df, 1, false)
Review Comment:
The comment says 300 partitions use the sort-based writer. Would it be worth
asserting that, so the test can't silently stop covering it if a threshold
changes? `checkCometExchange` returns the exchanges, so something like
`assert(exchanges.head.shuffleDependency.shuffleHandle.getClass.getName.contains("CometSerializedShuffleHandle")
== (numPartitions > 200))` should work (I haven't run it). The similar checks
in `Columnar shuffle for large shuffle partition number` discard the boolean,
so as written they don't assert anything.
##########
spark/src/main/java/org/apache/spark/shuffle/sort/SpillSorter.java:
##########
@@ -287,7 +287,9 @@ public void writeSortedFileNative(boolean isLastFile,
boolean tracingEnabled) th
spillInfo.partitionLengths[currentPartition] = written;
// Store the checksum for the current partition.
- partitionChecksums[currentPartition] = getChecksum();
+ if (partitionChecksums.length > 0) {
Review Comment:
Nit: the partition-switch block here and the final-partition block below are
near copies, and the final one already had this guard, so they drifted apart.
Would a small private helper that sets the checksum, calls `doSpilling`,
records the length and stores the checksum keep the guard in one place? Happy
to leave that for a follow-up if you'd rather keep this fix small for
backporting.
##########
spark/src/test/scala/org/apache/spark/sql/comet/execution/shuffle/CometDiskBlockWriterSuite.scala:
##########
@@ -212,6 +214,32 @@ class CometDiskBlockWriterSuite extends AnyFunSuite {
new JLinkedList[CometDiskBlockWriter]())
}
+ test("a writer computes no checksum when shuffle checksums are disabled") {
+ // With spark.shuffle.checksum.enabled=false the bypass merge sort writer
never calls
+ // setChecksum, so every write, including the ones after the first spill,
must skip it.
+ val conf = new SparkConf()
+ .set("spark.memory.offHeap.enabled", "true")
+ .set("spark.memory.offHeap.size", "1g")
+ val tmm = new TaskMemoryManager(new TestMemoryManager(conf), 0L)
+ val allocator = CometShuffleMemoryAllocator.getInstance(tmm, pageSize)
+ val tempDir = Utils.createTempDir()
+ // Spill every two rows so that the writer writes to its file several
times.
+ SQLConf.get.setConfString(CometConf.COMET_SHUFFLE_JVM_SPILL_THRESHOLD.key,
"2")
+ try {
+ val writer =
+ newWriter(new File(tempDir, "partition0"), allocator,
newTaskContext(tmm, 0L), conf)
+ val toUnsafe = UnsafeProjection.create(schema)
+ (0 until 5).foreach(_ => writer.insertRow(toUnsafe(InternalRow(new
Array[Byte](16))), 0))
+ assert(writer.close().length > 0)
+ assert(writer.getOutputRecords == 5)
+ assert(writer.getChecksum == -1)
Review Comment:
Would it make sense to also run this scenario with checksums enabled? That
is the default, and this PR changes the assignment in `doSpilling` that the
enabled path relies on, but I couldn't find a test that checks a checksum
value. With `writer.setChecksum(Long.MIN_VALUE)` and
`writer.setChecksumAlgo("adler32")` before the inserts, I'd expect
`getChecksum` to equal a one-shot `java.util.zip.Adler32` over the partition
file (I haven't run this). Parameterizing this test over both settings would
keep the test count flat.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]