Skip to content

Commit b628a94

Browse files
firasbouzaziviirya
authored andcommitted
[SPARK-59985][ML][PYTHON] Fix malformed warning in TorchDistributor w…
### What changes were proposed in this pull request? Fix the warning that `TorchDistributor._run_distributed_training` emits when the log streaming server fails to start. The message had a stray comma, so the error text was passed to `logger.warning` as an extra positional argument. Formatting the record then raised `TypeError: not all arguments converted during string formatting`. The warning now uses a single format string with the exception as its argument (`"... error: %r.", e`). ### Why are the changes needed? The fallback added in SPARK-44909 still works (training continues with the port set to `-1`), but the warning that tells users why worker logs are missing is broken: - With default logging settings, users see a `--- Logging error ---` traceback that looks like a crash instead of a one-line warning. - With `logging.raiseExceptions = False`, the warning is silently dropped. - `record.getMessage()` raises, so the warning cannot be asserted on. ### Does this PR introduce _any_ user-facing change? Yes. When the log streaming server fails to start, users now get the intended warning instead of a logging error traceback. Before: ``` --- Logging error --- Traceback (most recent call last): ... TypeError: not all arguments converted during string formatting Call stack: File ".../pyspark/ml/torch/distributor.py", line 815, in _run_distributed_training self.logger.warning( Message: 'Start torch distributor log streaming server failed, You cannot receive logs sent from distributor workers, ' Arguments: ("error: OSError('Address already in use').",) ``` After: ``` Start torch distributor log streaming server failed, You cannot receive logs sent from distributor workers, error: OSError('Address already in use'). ``` ### How was this patch tested? Added `test_log_streaming_server_start_failure_warns` to `TorchDistributorBaselineUnitTestsMixin`. It makes `LogStreamingServer.start` raise and checks that execution reaches task-function creation with log streaming disabled and that exactly one warning is logged containing the error. The test fails with the `TypeError` above before the fix and passes after it. The rest of `TorchDistributorBaselineUnitTests` still passes. ``` python/run-tests --testnames 'pyspark.ml.torch.tests.test_distributor TorchDistributorBaselineUnitTests' ``` ### Was this patch authored or co-authored using generative AI tooling? Co Authered with claude. Closes #59242 from firasbouzazi/SPARK-59985-fix-torch-distributor-warning. Authored-by: Firas Bouzazi <115628206+firasbouzazi@users.noreply.github.com> Signed-off-by: Liang-Chi Hsieh <viirya@gmail.com> (cherry picked from commit d7635d0) Signed-off-by: Liang-Chi Hsieh <viirya@gmail.com>
1 parent 72ec2bf commit b628a94

2 files changed

Lines changed: 28 additions & 2 deletions

File tree

‎python/pyspark/ml/torch/distributor.py‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -809,8 +809,9 @@ def _run_distributed_training(
809809
self.log_streaming_server_port = -1
810810
self.logger.warning(
811811
"Start torch distributor log streaming server failed, "
812-
"You cannot receive logs sent from distributor workers, ",
813-
f"error: {repr(e)}.",
812+
"You cannot receive logs sent from distributor workers, "
813+
"error: %r.",
814+
e,
814815
)
815816

816817
try:

‎python/pyspark/ml/torch/tests/test_distributor.py‎

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -275,6 +275,31 @@ def test_execute_command_survives_log_socket_drop(self) -> None:
275275
)
276276
self.assertIn("hello_after_socket_drop", output.getvalue().strip())
277277

278+
def test_log_streaming_server_start_failure_warns(self) -> None:
279+
"""A failure to start the log streaming server must only emit a warning."""
280+
281+
class StopTraining(Exception):
282+
pass
283+
284+
dist = TorchDistributor(num_processes=2, local_mode=False, use_gpu=False)
285+
with (
286+
patch(
287+
"pyspark.ml.torch.distributor.LogStreamingServer.start",
288+
side_effect=OSError("Address already in use"),
289+
),
290+
patch.object(dist, "_get_spark_task_function", side_effect=StopTraining),
291+
self.assertLogs(dist.logger, level="WARNING") as logs,
292+
self.assertRaises(StopTraining),
293+
):
294+
dist._run_distributed_training(MagicMock(), MagicMock(), None, None)
295+
296+
self.assertEqual(dist.log_streaming_server_port, -1)
297+
self.assertIsNone(dist.log_streaming_auth_secret)
298+
self.assertEqual(len(logs.records), 1)
299+
message = logs.records[0].getMessage()
300+
self.assertIn("Start torch distributor log streaming server failed", message)
301+
self.assertIn("Address already in use", message)
302+
278303
def test_create_torchrun_command(self) -> None:
279304
train_path = "train.py"
280305
args_string = ["1", "3"]

0 commit comments

Comments
 (0)