Skip to content
Closed
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
20 changes: 20 additions & 0 deletions docs/Explore Algorithms/LightGBM/Overview.md
Original file line number Diff line number Diff line change
Expand Up @@ -404,6 +404,26 @@ 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.

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.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -181,11 +182,38 @@ 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.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)
}
} 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)),
Expand Down Expand Up @@ -758,7 +786,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()
Expand Down
Original file line number Diff line number Diff line change
@@ -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)
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,14 +3,16 @@

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}
import java.net.{ConnectException, InetSocketAddress, ServerSocket, Socket, SocketException, SocketTimeoutException}
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.
Expand Down Expand Up @@ -70,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()
}

Expand Down Expand Up @@ -404,6 +414,48 @@ 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") {
// 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
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()
}
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)
Expand Down
Loading
Loading