From e51ff84c0359ed543f367a62d8c0cfc892dc4089 Mon Sep 17 00:00:00 2001 From: Ranadeep Singh Date: Tue, 29 Sep 2026 09:10:06 +0000 Subject: [PATCH 1/3] fix: expand LightGBM partitions to numTasks and explain missing-task timeouts Without barrier execution mode, the LightGBM driver waits for exactly numTasks training tasks to report. Two cases left it waiting until the timeout, after which tasks failed with a misleading "Connection refused" error: - An explicit numTasks larger than the input partition count. coalesce can only merge partitions, so fewer tasks ran than the driver expected. The input is now repartitioned up to numTasks. The ranker repartitions by its grouping column so query groups stay whole. The partition count is only read for an explicit numTasks, so automatic sizing adds no extra work. - numTasks larger than the tasks Spark can run at once. This can't be fixed automatically, so the driver now fails with LightGBMMissingTasksException, which lists how many tasks reported and which partitions are missing. The job failure is attached as a suppressed exception. Adds regression tests that fail on master, and a troubleshooting note in the LightGBM overview. Addresses part of #2699. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- docs/Explore Algorithms/LightGBM/Overview.md | 14 ++ .../synapse/ml/lightgbm/LightGBMBase.scala | 39 ++++- .../LightGBMMissingTasksException.scala | 38 +++++ .../synapse/ml/lightgbm/LightGBMRanker.scala | 40 ++--- .../synapse/ml/lightgbm/NetworkManager.scala | 36 +++- .../split1/DriverSocketRetrySuite.scala | 45 ++++- .../split1/LightGBMNumTasksSuite.scala | 155 ++++++++++++++++++ 7 files changed, 341 insertions(+), 26 deletions(-) create mode 100644 lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMMissingTasksException.scala create mode 100644 lightgbm/src/test/scala/com/microsoft/azure/synapse/ml/lightgbm/split1/LightGBMNumTasksSuite.scala diff --git a/docs/Explore Algorithms/LightGBM/Overview.md b/docs/Explore Algorithms/LightGBM/Overview.md index 97a6ed4184..22e56b97cc 100644 --- a/docs/Explore Algorithms/LightGBM/Overview.md +++ b/docs/Explore Algorithms/LightGBM/Overview.md @@ -404,6 +404,20 @@ When this happens, the reported error explains that it's a retry that could not names the partition to investigate. Look for the **first** failed attempt of that partition in the executor logs — that attempt holds the real cause. +#### When numTasks is larger than the tasks Spark can run at once + +Without barrier execution mode, the driver waits for exactly *numTasks* training tasks to report, and they +must all run at the same time. SynapseML repartitions the input when an explicit *numTasks* is larger than its +partition count, but Spark can only run as many tasks at once as the cluster has task slots (each executor's +cores divided by `spark.task.cpus`, summed across executors). If *numTasks* is larger than that, the tasks +that started wait for tasks that can't start until they finish. + +After *timeout* seconds, the driver stops waiting and training fails with an error that lists how many tasks +reported and which partitions are missing. Other tasks may then also report "could not reach the driver" or +"Connection refused"; those errors are a result of the timeout. To fix it, lower *numTasks* to the number of +task slots, or leave *numTasks* unset so SynapseML chooses it. Also check that no executor was lost or still +starting when training began. + ### IPv6 clusters Distributed training works on clusters whose executors only have IPv6 addresses. diff --git a/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMBase.scala b/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMBase.scala index 8a888039cf..e2e6986b90 100644 --- a/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMBase.scala +++ b/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMBase.scala @@ -22,6 +22,7 @@ import org.apache.spark.sql.types._ import scala.collection.immutable.HashSet import scala.language.existentials import scala.math.min +import scala.util.control.NonFatal import scala.util.matching.Regex // scalastyle:off file.size.limit @@ -181,11 +182,37 @@ trait LightGBMBase[TrainedModel <: Model[TrainedModel] with LightGBMModelParams] } else { df } + } else { + fitNonBarrierPartitions(df, numTasks) + } + } + + /** Gives non-barrier training exactly numTasks partitions. + * + * Without barrier execution the driver waits for exactly numTasks workers, so fewer partitions leave + * it waiting for tasks that never start. coalesce can only merge partitions, so an explicit numTasks + * above the input partition count needs a shuffle. An automatic numTasks is already capped at the input + * partition count. The count is only read for an explicit numTasks because reading it can run the + * input's adaptive shuffle stages early. + */ + private def fitNonBarrierPartitions(df: DataFrame, numTasks: Int): DataFrame = { + if (getNumTasks > 0) { + val numPartitions = df.rdd.getNumPartitions + if (numPartitions < numTasks) { + log.info(s"Repartitioning $numPartitions input partitions to numTasks=$numTasks, because training " + + "without barrier execution mode waits for exactly numTasks workers") + expandPartitions(df, numTasks) + } else { + df.coalesce(numTasks) + } } else { df.coalesce(numTasks) } } + /** Splits the training data into numTasks partitions when the input has fewer. */ + protected def expandPartitions(df: DataFrame, numTasks: Int): DataFrame = df.repartition(numTasks) + protected def getTrainingCols: Array[(String, Seq[DataType])] = { val colsToCheck: Array[(Option[String], Seq[DataType])] = Array( (Some(getLabelCol), Seq(DoubleType)), @@ -758,7 +785,17 @@ trait LightGBMBase[TrainedModel <: Model[TrainedModel] with LightGBMModelParams] // Execute the Tasks on workers val lightGBMBooster = try { - val booster = executePartitionTasks(ctx, dataframe, measures) + val booster = try { + executePartitionTasks(ctx, dataframe, measures) + } catch { + case NonFatal(jobFailure) => + // Tasks that lose the driver report a misleading connection error, so prefer the driver's + // own explanation when it stopped waiting for tasks that never started. + throw networkManager.missingTasksFailure(LightGBMMissingTasksException.MaxDriverWait).map { missingTasks => + missingTasks.addSuppressed(jobFailure) + missingTasks + }.getOrElse(jobFailure) + } // Wait for network to complete (should be done by now) networkManager.waitForNetworkCommunicationsDone() diff --git a/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMMissingTasksException.scala b/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMMissingTasksException.scala new file mode 100644 index 0000000000..12fe8f5d44 --- /dev/null +++ b/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMMissingTasksException.scala @@ -0,0 +1,38 @@ +// Copyright (C) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. See LICENSE in project root for information. + +package com.microsoft.azure.synapse.ml.lightgbm + +import scala.concurrent.duration.{Duration, SECONDS} + +/** Thrown when training without barrier execution mode stops waiting for tasks that never reported. */ +class LightGBMMissingTasksException private[lightgbm] (message: String, cause: Throwable) + extends Exception(message, cause) + +private[lightgbm] object LightGBMMissingTasksException { + private val MaxListedPartitions = 20 + private val MaxDriverWaitSeconds = 30 + + /** How long a failed training job waits to learn whether the driver timed out first. */ + val MaxDriverWait: Duration = Duration(MaxDriverWaitSeconds, SECONDS) + + def apply(numTasks: Int, + missingPartitions: Seq[Int], + timeoutSeconds: Double, + cause: Throwable): LightGBMMissingTasksException = { + val listed = missingPartitions.take(MaxListedPartitions).mkString(", ") + val unlisted = missingPartitions.size - MaxListedPartitions + val missingList = if (unlisted > 0) s"$listed, and $unlisted more" else listed + val timeoutText = if (timeoutSeconds.isWhole) timeoutSeconds.toLong.toString else timeoutSeconds.toString + val message = + s"The LightGBM driver received network reports from ${numTasks - missingPartitions.size} of $numTasks " + + s"training tasks, then stopped waiting because no task connected for $timeoutText seconds. Missing " + + s"partitions: $missingList. Without barrier execution mode, all numTasks tasks must run at the same " + + "time. Check that numTasks is no larger than the number of tasks Spark can run at once (each " + + "executor's cores divided by spark.task.cpus, summed across executors), and that no executor was " + + "lost or still starting. If a task failed before reporting, its own error is the root cause. Later " + + "\"could not reach the driver\" or \"connection refused\" errors from other tasks are a result of " + + "this timeout." + new LightGBMMissingTasksException(message, cause) + } +} diff --git a/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMRanker.scala b/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMRanker.scala index 4b8865e849..f4c1f5d56d 100644 --- a/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMRanker.scala +++ b/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMRanker.scala @@ -92,27 +92,27 @@ class LightGBMRanker(override val uid: String) override def copy(extra: ParamMap): LightGBMRanker = defaultCopy(extra) override def prepareDataframe(dataset: Dataset[_], numTasks: Int): DataFrame = { - if (getRepartitionByGroupingColumn) { - val repartitionedDataset = getOptGroupCol match { - case None => dataset - case Some(groupingCol) => - val numPartitions = dataset.rdd.getNumPartitions - val groupingPartitions = if (getUseBarrierExecutionMode) { - math.min(numPartitions, numTasks) - } else { - numTasks - } - - // Use an explicit partition count so adaptive execution preserves the - // grouping topology. Barrier mode preserves its existing no-expansion - // behavior, while non-barrier mode must create the numTasks workers that - // NetworkManager waits for. - dataset.repartition(groupingPartitions, new Column(groupingCol)) - } + getOptGroupCol.filter(_ => getRepartitionByGroupingColumn) match { + case Some(groupingCol) if !getUseBarrierExecutionMode => + // An explicit partition count keeps adaptive execution from coalescing the grouping shuffle, + // and gives NetworkManager the numTasks workers it waits for. The result already has exactly + // numTasks partitions, so the base class's partition fitting is skipped. + castColumns(dataset.repartition(numTasks, new Column(groupingCol)), getTrainingCols) + case Some(groupingCol) => + // Barrier mode never expands the input, and waits only for the tasks the stage actually runs. + val numPartitions = dataset.rdd.getNumPartitions + super.prepareDataframe( + dataset.repartition(math.min(numPartitions, numTasks), new Column(groupingCol)), numTasks) + case None => + super.prepareDataframe(dataset, numTasks) + } + } - super.prepareDataframe(repartitionedDataset, numTasks) - } else { - super.prepareDataframe(dataset, numTasks) + /** Expands by the grouping column, so every query group stays within one partition. */ + override protected def expandPartitions(df: DataFrame, numTasks: Int): DataFrame = { + getOptGroupCol match { + case Some(groupingCol) => df.repartition(numTasks, new Column(groupingCol)) + case None => super.expandPartitions(df, numTasks) } } } diff --git a/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/NetworkManager.scala b/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/NetworkManager.scala index 9fe0905049..5269bef861 100644 --- a/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/NetworkManager.scala +++ b/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/NetworkManager.scala @@ -14,7 +14,7 @@ import org.slf4j.Logger import java.io.{BufferedReader, BufferedWriter, IOException, InputStreamReader, OutputStreamWriter} import java.net.{ConnectException, ServerSocket, Socket, SocketException, SocketTimeoutException} -import java.util.concurrent.{ExecutorService, Executors} +import java.util.concurrent.{ExecutorService, Executors, TimeoutException} import scala.annotation.tailrec import scala.collection.mutable import scala.concurrent.{Await, ExecutionContext, ExecutionContextExecutor, Future} @@ -147,8 +147,10 @@ object NetworkManager { "partial task retry." } else { s"LightGBM task $taskId (partition $partitionId) could not reach the driver network topology endpoint " + - s"$endpoint on its first attempt. Verify that executors are allowed to open connections to the driver " + - "on that port, and that the driver was not shut down before training started." + s"$endpoint on its first attempt. Either executors cannot open connections to the driver on that " + + "port, or the driver stopped waiting before this task reported, for example because numTasks is " + + "larger than the number of tasks Spark can run at once. Check the driver log for a LightGBM " + + "missing-tasks error." } log.error(message, cause) new Exception(message, cause) @@ -578,7 +580,33 @@ case class NetworkManager(numTasks: Int, if (reportedTaskCount < numTasks) connectToWorkers() } - connectToWorkers() + try { + connectToWorkers() + } catch { + case acceptTimeout: SocketTimeoutException => + // Tasks that already reported are disconnected next and fail with "could not reach the driver", + // which hides the fact that the driver gave up waiting for the rest. + val missing = synchronized((0 until numTasks).filterNot(taskConnectionsByPartition.contains)) + val failure = LightGBMMissingTasksException(numTasks, missing, timeout, acceptTimeout) + log.error(failure.getMessage) + throw failure + } + } + } + + /** Returns the driver's missing-task timeout, if that is what ended the topology round. + * + * Closing the connections first releases a driver still blocked in accept(), so the wait is short. + */ + private[lightgbm] def missingTasksFailure(maxWait: Duration): Option[LightGBMMissingTasksException] = { + closeConnections() + try { + Await.ready(networkCommunicationThread, maxWait) + } catch { + case _: TimeoutException => () + } + networkCommunicationThread.value.flatMap(_.failed.toOption).collect { + case failure: LightGBMMissingTasksException => failure } } diff --git a/lightgbm/src/test/scala/com/microsoft/azure/synapse/ml/lightgbm/split1/DriverSocketRetrySuite.scala b/lightgbm/src/test/scala/com/microsoft/azure/synapse/ml/lightgbm/split1/DriverSocketRetrySuite.scala index 87315c9e6c..a3b21f3be0 100644 --- a/lightgbm/src/test/scala/com/microsoft/azure/synapse/ml/lightgbm/split1/DriverSocketRetrySuite.scala +++ b/lightgbm/src/test/scala/com/microsoft/azure/synapse/ml/lightgbm/split1/DriverSocketRetrySuite.scala @@ -3,7 +3,8 @@ package com.microsoft.azure.synapse.ml.lightgbm.split1 -import com.microsoft.azure.synapse.ml.lightgbm.{LightGBMConstants, NetworkManager, TaskMessageInfo, WorkerMessage} +import com.microsoft.azure.synapse.ml.lightgbm.{LightGBMConstants, LightGBMMissingTasksException, NetworkManager, + TaskMessageInfo, WorkerMessage} import org.scalatest.funsuite.AnyFunSuite import java.io.{BufferedReader, BufferedWriter, IOException, InputStreamReader, OutputStreamWriter} @@ -11,6 +12,7 @@ import java.net.{ConnectException, InetSocketAddress, ServerSocket, Socket, Sock import java.util.concurrent.{CountDownLatch, TimeUnit} import java.util.concurrent.atomic.AtomicInteger import scala.collection.mutable.ListBuffer +import scala.concurrent.duration.{Duration, SECONDS} /** Covers the driver topology socket lifecycle behind repeated * "java.net.ConnectException: Connection refused" failures in distributed LightGBM training. @@ -404,6 +406,47 @@ class DriverSocketRetrySuite extends AnyFunSuite { assert(failure.getMessage.contains("closed the connection before sending a status message")) } + test("Non-barrier topology names the missing partitions when the driver stops waiting") { + val serverSocket = new ServerSocket(0) + val port = serverSocket.getLocalPort + // The accept timeout ends the round quickly; the manager timeout only bounds the wait below. + serverSocket.setSoTimeout(500) + val manager = NetworkManager(3, serverSocket, host, port, timeout, useBarrierExecutionMode = false) + var task = Option.empty[FakeTask] + try { + task = Some(new FakeTask(host, port, partitionId = 1)) + task.get.report() + + val failure = intercept[LightGBMMissingTasksException] { + manager.waitForNetworkCommunicationsDone() + } + assert(failure.getMessage.contains("from 1 of 3 training tasks")) + assert(failure.getMessage.contains("Missing partitions: 0, 2.")) + assert(failure.getMessage.contains("numTasks")) + assert(failure.getCause.isInstanceOf[SocketTimeoutException]) + assert(task.get.isClosedByDriver, "The driver kept the reported task waiting after giving up") + assert(manager.missingTasksFailure(Duration(5, SECONDS)).contains(failure)) + } finally { + manager.closeConnections() + closeTasks(task) + } + } + + test("Other topology failures are not reported as missing tasks") { + val (manager, _, port) = newManager(numTasks = 2) + var task = Option.empty[FakeTask] + try { + task = Some(new FakeTask(host, port, partitionId = 0)) + task.get.report() + + // The training job failing first closes the driver socket, which is not a missing-task timeout. + assert(manager.missingTasksFailure(Duration(5, SECONDS)).isEmpty) + } finally { + manager.closeConnections() + closeTasks(task) + } + } + test("The driver server socket is released when a training job fails before the round completes") { // Only one of the two expected tasks reports, so the network thread stays blocked in accept(). val (manager, _, port) = newManager(numTasks = 2, useBarrierExecutionMode = true) diff --git a/lightgbm/src/test/scala/com/microsoft/azure/synapse/ml/lightgbm/split1/LightGBMNumTasksSuite.scala b/lightgbm/src/test/scala/com/microsoft/azure/synapse/ml/lightgbm/split1/LightGBMNumTasksSuite.scala new file mode 100644 index 0000000000..5ef324e310 --- /dev/null +++ b/lightgbm/src/test/scala/com/microsoft/azure/synapse/ml/lightgbm/split1/LightGBMNumTasksSuite.scala @@ -0,0 +1,155 @@ +// Copyright (C) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. See LICENSE in project root for information. + +package com.microsoft.azure.synapse.ml.lightgbm.split1 + +import com.microsoft.azure.synapse.ml.lightgbm.{LightGBMClassifier, LightGBMMissingTasksException, LightGBMRanker} +import org.apache.spark.TaskContext +import org.apache.spark.ml.feature.VectorAssembler +import org.apache.spark.sql.DataFrame +import org.apache.spark.sql.functions.{col, floor} + +/** Covers non-barrier training when the requested numTasks does not match the tasks that can run. */ +// scalastyle:off magic.number +class LightGBMNumTasksSuite extends LightGBMTestUtils { + + private val queryCol = "query" + private val numRows = 512 + + private def trainingData(inputPartitions: Int): DataFrame = { + val input = spark.range(0L, numRows.toLong, 1L, inputPartitions) + .withColumn(queryCol, floor(col("id") / 8L).cast("long")) + .withColumn(labelCol, (col("id") % 2L).cast("double")) + .withColumn("signal", col(labelCol)) + .withColumn("other", (col("id") % 11L).cast("double")) + + new VectorAssembler() + .setInputCols(Array("signal", "other")) + .setOutputCol(featuresCol) + .transform(input) + .select(queryCol, labelCol, featuresCol) + } + + /** Exposes the protected partition preparation that fit uses. */ + private class PreparingClassifier extends LightGBMClassifier { + def prepare(data: DataFrame, numTasks: Int): DataFrame = prepareDataframe(data, numTasks) + } + + private def classifier(numTasks: Option[Int], timeoutSeconds: Double = 120): PreparingClassifier = { + val estimator = new PreparingClassifier() + estimator + .setFeaturesCol(featuresCol) + .setLabelCol(labelCol) + .setUseBarrierExecutionMode(false) + .setNumThreads(1) + .setNumLeaves(4) + .setNumIterations(10) + // Bound a worker-count mismatch so it fails the test instead of waiting for the 1200s default. + .setTimeout(timeoutSeconds) + .setDefaultListenPort(getAndIncrementPort()) + numTasks.foreach(estimator.setNumTasks) + estimator + } + + private def ranker(numTasks: Int): LightGBMRanker = { + new LightGBMRanker() + .setFeaturesCol(featuresCol) + .setLabelCol(labelCol) + .setGroupCol(queryCol) + .setRepartitionByGroupingColumn(false) + .setUseBarrierExecutionMode(false) + .setNumTasks(numTasks) + .setNumThreads(1) + .setNumLeaves(4) + .setNumIterations(10) + .setTimeout(120) + .setDefaultListenPort(getAndIncrementPort()) + } + + private def partitionGroups(df: DataFrame): Array[(Int, Set[Long])] = { + import df.sparkSession.implicits._ + // mapPartitions runs once per partition, so empty partitions are still counted. + df.select(queryCol).as[Long].mapPartitions { groups => + Iterator(TaskContext.getPartitionId() -> groups.toSet.toSeq) + }.collect().map { case (partitionIndex, groups) => partitionIndex -> groups.toSet } + } + + private def causes(failure: Throwable): Seq[Throwable] = { + Iterator.iterate(failure)(_.getCause).takeWhile(_ != null).take(20).toSeq //scalastyle:ignore null + } + + test("an explicit numTasks above the input partition count repartitions to numTasks") { + val prepared = classifier(Some(4)).prepare(trainingData(inputPartitions = 1), numTasks = 4) + + assert(prepared.rdd.getNumPartitions === 4) + assert(prepared.count() === numRows) + } + + test("an explicit numTasks at or below the input partition count still coalesces") { + val prepared = classifier(Some(2)).prepare(trainingData(inputPartitions = 4), numTasks = 2) + + assert(prepared.rdd.getNumPartitions === 2) + } + + test("an automatic numTasks does not add a shuffle") { + // determineNumTasks never picks more tasks than input partitions, so this only guards the check. + val prepared = classifier(None).prepare(trainingData(inputPartitions = 1), numTasks = 4) + + assert(prepared.rdd.getNumPartitions === 1) + } + + // Non-barrier training needs every task running at once, and CI agents have two cores, + // so the fit tests below request two tasks. + test("non-barrier classifier fits when numTasks exceeds the input partitions") { + val data = trainingData(inputPartitions = 1) + val model = classifier(Some(2)).fit(data) + try { + val scored = model.transform(data).select(labelCol, predCol).collect() + assert(scored.length === numRows) + val accuracy = scored.count(row => row.getDouble(0) == row.getDouble(1)).toDouble / scored.length + assert(accuracy > 0.95, s"The signal feature equals the label, but accuracy was $accuracy") + } finally { + model.getModel.freeNativeMemory() + } + } + + test("ranker expansion keeps each query group in one partition without grouping repartition") { + val partitions = partitionGroups(ranker(4).prepareDataframe(trainingData(inputPartitions = 1), numTasks = 4)) + + assert(partitions.length === 4) + val groupLocations = partitions.flatMap { case (partitionIndex, groups) => + groups.map(_ -> partitionIndex) + }.groupBy(_._1).map { case (group, locations) => group -> locations.map(_._2).toSet } + assert(groupLocations.size === numRows / 8) + assert(groupLocations.values.forall(_.size === 1)) + } + + test("non-barrier ranker fits when numTasks exceeds the input partitions without grouping repartition") { + val data = trainingData(inputPartitions = 1) + val model = ranker(2).fit(data) + try { + val predictions = model.transform(data).select(predCol).collect().map(_.getDouble(0)) + assert(predictions.length === numRows) + assert(predictions.forall(p => !p.isNaN && !p.isInfinite)) + } finally { + model.getModel.freeNativeMemory() + } + } + + test("numTasks above the concurrent task slots fails with the missing-task explanation") { + // Only defaultParallelism tasks can run at once in local mode, so one task never starts and the + // driver stops waiting after the timeout. + val numTasks = spark.sparkContext.defaultParallelism + 1 + val data = trainingData(inputPartitions = numTasks) + val failure = intercept[Exception] { + classifier(Some(numTasks), timeoutSeconds = 10).fit(data) + } + + val missingTasks = causes(failure).collectFirst { case e: LightGBMMissingTasksException => e } + assert(missingTasks.isDefined, s"Expected a missing-task explanation, got: $failure") + val message = missingTasks.get.getMessage + assert(message.contains(s"of $numTasks training tasks"), message) + assert(message.contains("Missing partitions:"), message) + assert(message.contains("numTasks is no larger than the number of tasks Spark can run at once"), message) + } +} From 028c5aa1a0d0c3eb581c7bfa670abed697146072 Mon Sep 17 00:00:00 2001 From: Ranadeep Singh Date: Wed, 30 Sep 2026 08:31:44 +0000 Subject: [PATCH 2/3] fix: warn when LightGBM adds a shuffle for numTasks and document the extra pass Log the numTasks expansion at warn level, because it adds a shuffle the user did not ask for, and say how to avoid it. Document that an explicit numTasks now reads the input partition count, like automatic numTasks and barrier mode already do, which runs a pending adaptive shuffle once more on uncached input. Fabric E2E (runtime Spark 3.5.5, 1 executor x 8 slots, AQE on) measured explicit numTasks fits at 10.2s median versus 7.9s on master for an uncached 20M-row aggregate, matching one extra upstream pass. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- docs/Explore Algorithms/LightGBM/Overview.md | 6 ++++++ .../microsoft/azure/synapse/ml/lightgbm/LightGBMBase.scala | 5 +++-- 2 files changed, 9 insertions(+), 2 deletions(-) diff --git a/docs/Explore Algorithms/LightGBM/Overview.md b/docs/Explore Algorithms/LightGBM/Overview.md index 22e56b97cc..a7a2b6afea 100644 --- a/docs/Explore Algorithms/LightGBM/Overview.md +++ b/docs/Explore Algorithms/LightGBM/Overview.md @@ -418,6 +418,12 @@ reported and which partitions are missing. Other tasks may then also report "cou task slots, or leave *numTasks* unset so SynapseML chooses it. Also check that no executor was lost or still starting when training began. +To decide whether to repartition, SynapseML reads the input's partition count when *numTasks* is set. It +already does this when *numTasks* is unset or barrier execution mode is on. If the input is an uncached +DataFrame with a shuffle that hasn't run yet, such as a join or aggregation with adaptive query execution +on, reading the count runs that shuffle one extra time. Training reads the input several times anyway, so +if the input is expensive to compute, cache or persist it before calling `fit`. + ### IPv6 clusters Distributed training works on clusters whose executors only have IPv6 addresses. diff --git a/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMBase.scala b/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMBase.scala index e2e6986b90..5022c0eee3 100644 --- a/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMBase.scala +++ b/lightgbm/src/main/scala/com/microsoft/azure/synapse/ml/lightgbm/LightGBMBase.scala @@ -199,8 +199,9 @@ trait LightGBMBase[TrainedModel <: Model[TrainedModel] with LightGBMModelParams] if (getNumTasks > 0) { val numPartitions = df.rdd.getNumPartitions if (numPartitions < numTasks) { - log.info(s"Repartitioning $numPartitions input partitions to numTasks=$numTasks, because training " + - "without barrier execution mode waits for exactly numTasks workers") + log.warn(s"Repartitioning $numPartitions input partitions to numTasks=$numTasks, because training " + + "without barrier execution mode waits for exactly numTasks workers. This adds a shuffle; give the " + + "input at least numTasks partitions to avoid it") expandPartitions(df, numTasks) } else { df.coalesce(numTasks) From 1e1687e55dcffa928dfaa62289b80127f171bbb2 Mon Sep 17 00:00:00 2001 From: Ranadeep Singh Date: Wed, 30 Sep 2026 12:06:16 +0000 Subject: [PATCH 3/3] test: start the short driver accept timeout after the first report The missing-partition test set a 500ms accept timeout before the fake worker connected, so a slow agent could time out with zero reports and fail the assertion on partition 1. The driver accepts one connection at a time, so the helper now applies the short timeout on the second accept, after the first report is recorded, and the test waits for that point. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../split1/DriverSocketRetrySuite.scala | 19 ++++++++++++++----- 1 file changed, 14 insertions(+), 5 deletions(-) diff --git a/lightgbm/src/test/scala/com/microsoft/azure/synapse/ml/lightgbm/split1/DriverSocketRetrySuite.scala b/lightgbm/src/test/scala/com/microsoft/azure/synapse/ml/lightgbm/split1/DriverSocketRetrySuite.scala index a3b21f3be0..bfeae5fd38 100644 --- a/lightgbm/src/test/scala/com/microsoft/azure/synapse/ml/lightgbm/split1/DriverSocketRetrySuite.scala +++ b/lightgbm/src/test/scala/com/microsoft/azure/synapse/ml/lightgbm/split1/DriverSocketRetrySuite.scala @@ -72,13 +72,21 @@ class DriverSocketRetrySuite extends AnyFunSuite { override def close(): Unit = socket.close() } - private class SignalSecondAcceptServerSocket extends ServerSocket(0) { + /** @param laterAcceptTimeoutMillis accept timeout from the second accept on. The driver accepts one + * connection at a time, so this starts only after the first report + * was recorded. + */ + private class SignalSecondAcceptServerSocket(laterAcceptTimeoutMillis: Int = socketTimeoutMillis) + extends ServerSocket(0) { private val acceptCount = new AtomicInteger() private val secondAcceptStarted = new CountDownLatch(1) setSoTimeout(socketTimeoutMillis) override def accept(): Socket = { - if (acceptCount.incrementAndGet() == 2) secondAcceptStarted.countDown() + if (acceptCount.incrementAndGet() == 2) { + setSoTimeout(laterAcceptTimeoutMillis) + secondAcceptStarted.countDown() + } super.accept() } @@ -407,15 +415,16 @@ class DriverSocketRetrySuite extends AnyFunSuite { } test("Non-barrier topology names the missing partitions when the driver stops waiting") { - val serverSocket = new ServerSocket(0) + // The short accept timeout ends the round quickly once the first report is recorded; the manager + // timeout only bounds the wait below. + val serverSocket = new SignalSecondAcceptServerSocket(laterAcceptTimeoutMillis = 500) val port = serverSocket.getLocalPort - // The accept timeout ends the round quickly; the manager timeout only bounds the wait below. - serverSocket.setSoTimeout(500) val manager = NetworkManager(3, serverSocket, host, port, timeout, useBarrierExecutionMode = false) var task = Option.empty[FakeTask] try { task = Some(new FakeTask(host, port, partitionId = 1)) task.get.report() + serverSocket.awaitSecondAccept() val failure = intercept[LightGBMMissingTasksException] { manager.waitForNetworkCommunicationsDone()