diff --git a/examples/hf_ptq/scripts/huggingface_example.sh b/examples/hf_ptq/scripts/huggingface_example.sh index adc4cf22c..68f439ec1 100755 --- a/examples/hf_ptq/scripts/huggingface_example.sh +++ b/examples/hf_ptq/scripts/huggingface_example.sh @@ -318,9 +318,20 @@ if [[ $TASKS =~ "mmlu" ]]; then fi if [[ ! -d "$MMLU_DATA_PATH" ]] || [[ ! $(ls -A $MMLU_DATA_PATH) ]]; then echo "Preparing the MMLU test data" - wget https://people.eecs.berkeley.edu/~hendrycks/data.tar -O /tmp/mmlu.tar + MMLU_TAR="$(mktemp "${TMPDIR:-/tmp}/mmlu.XXXXXX")" || exit 1 + trap 'rm -f "$MMLU_TAR"' EXIT + # Revision-pinned HuggingFace mirror of the Berkeley tarball, which is unreachable. Bound + # the retries, which otherwise outlast the callers' timeouts, and resume rather than refetch. + wget --connect-timeout=20 --read-timeout=60 --tries=3 -c \ + https://huggingface.co/datasets/cais/mmlu/resolve/c30699e8356da336a370243923dbaf21066bb9fe/data.tar \ + -O "$MMLU_TAR" || { + echo "[ERROR] Could not download the MMLU test data. Set MMLU_DATA_PATH to a local copy." + exit 1 + } mkdir -p data - tar -xf /tmp/mmlu.tar -C data && mv data/data $MMLU_DATA_PATH + tar -xf "$MMLU_TAR" -C data && mv data/data $MMLU_DATA_PATH + rm -f "$MMLU_TAR" # 166MB; do not hold it for the rest of the eval + trap - EXIT fi mmlu_flags="" diff --git a/examples/llm_eval/README.md b/examples/llm_eval/README.md index f4c53f49d..8f761379b 100644 --- a/examples/llm_eval/README.md +++ b/examples/llm_eval/README.md @@ -155,7 +155,8 @@ Download data ```bash mkdir -p data -wget https://people.eecs.berkeley.edu/~hendrycks/data.tar -O data/mmlu.tar +wget --connect-timeout=20 --read-timeout=60 --tries=3 -c \ + https://huggingface.co/datasets/cais/mmlu/resolve/c30699e8356da336a370243923dbaf21066bb9fe/data.tar -O data/mmlu.tar tar -xf data/mmlu.tar -C data && mv data/data data/mmlu cd .. ``` diff --git a/examples/windows/accuracy_benchmark/README.md b/examples/windows/accuracy_benchmark/README.md index 8392414a7..9c379afc5 100644 --- a/examples/windows/accuracy_benchmark/README.md +++ b/examples/windows/accuracy_benchmark/README.md @@ -41,7 +41,7 @@ The table below lists the setup steps to prepare your environment for evaluating | **Install PyTorch and Related Packages** | `pip install torch==2.7.0 torchvision==0.22.0 torchaudio==2.7.0 --index-url https://download.pytorch.org/whl/cu128` | | **Install ONNX Runtime Packages** | `pip install onnxruntime-directml==1.21.1`
`pip install onnxruntime-genai-directml==0.6.0` | | **Install Benchmark Requirements** | `pip install -r requirements.txt` | -| **Download MMLU Data** | `mkdir data`
`curl -o .\data\mmlu.tar https://people.eecs.berkeley.edu/~hendrycks/data.tar`
`tar -xf .\data\mmlu.tar -C .\data`
`Move-Item .\data\data .\data\mmlu` | +| **Download MMLU Data** | `mkdir data`
`curl.exe -L -o .\data\mmlu.tar https://huggingface.co/datasets/cais/mmlu/resolve/c30699e8356da336a370243923dbaf21066bb9fe/data.tar`
`tar -xf .\data\mmlu.tar -C .\data`
`Move-Item .\data\data .\data\mmlu` | ### Evaluation Methods diff --git a/tests/_test_utils/examples/run_command.py b/tests/_test_utils/examples/run_command.py index 44d8f06f6..329aaf727 100644 --- a/tests/_test_utils/examples/run_command.py +++ b/tests/_test_utils/examples/run_command.py @@ -14,8 +14,11 @@ # limitations under the License. """Utility functions for running example commands reused in multiple example tests.""" +import contextlib import os +import signal import subprocess +import threading import time import warnings from pathlib import Path @@ -61,13 +64,76 @@ def extend_cmd_parts(cmd_parts: list[str], **kwargs): return cmd_parts +# Grace period for descendants to flush and close the inherited output pipe after the command exits, +# then a short bound on the post-SIGKILL drain (the fds close as the kernel tears the group down). +_ORPHAN_PIPE_TIMEOUT_S = 30 +_KILLED_PIPE_TIMEOUT_S = 5 + + def _run_capturing(cmd_parts: list[str], cwd: Path, env: dict[str, str]) -> tuple[int, str]: - """Run a command, capturing combined stdout/stderr to catch transient HF errors.""" - result = subprocess.run( - cmd_parts, cwd=cwd, env=env, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True + """Run a command, streaming and capturing combined stdout/stderr to catch transient HF errors. + + Drained by a reader thread rather than ``subprocess.run``, which waits for EOF: a descendant + outliving the command keeps the pipe open, hanging the read and discarding every log line. + """ + chunks: list[str] = [] + process = subprocess.Popen( + cmd_parts, + cwd=cwd, + env=env, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + start_new_session=True, ) - print(result.stdout, end="") - return result.returncode, result.stdout + # start_new_session makes the child its own group leader, so the group id is its pid (read now: + # os.getpgid() stops resolving it once wait() reaps the child). + pgid = process.pid + + def _kill_process_group(): + with contextlib.suppress(ProcessLookupError, PermissionError): + if os.name == "posix": + os.killpg(pgid, signal.SIGKILL) + else: # no process groups; the command itself is the best we can reach + process.kill() + + def _drain(): + with contextlib.suppress(ValueError): # the stream is closed under us on the escape path + for line in process.stdout: # type: ignore[union-attr] + print(line, end="") + chunks.append(line) + + reader = threading.Thread(target=_drain, daemon=True) + reader.start() + try: + if os.name == "posix": + # Wait *without* reaping: a reaped pid can be recycled, and pgid == the child's pid, so + # reaping before the kill below would risk signalling an unrelated group. + os.waitid(os.P_PID, process.pid, os.WEXITED | os.WNOWAIT) + else: # no process groups, so nothing to protect the pid for + process.wait() + except BaseException: + # Interrupted (most likely pytest-timeout) while the command runs: its own process group is + # not swept up with pytest's, so take the descendants down rather than leak them. + _kill_process_group() + with contextlib.suppress(Exception): + process.wait(timeout=_KILLED_PIPE_TIMEOUT_S) # reap: WNOWAIT left it waitable + raise + reader.join(_ORPHAN_PIPE_TIMEOUT_S) + + if reader.is_alive(): + warnings.warn( + f"{cmd_parts[0]} left descendants holding its output pipe open; killing the process " + "group. Output may be truncated." + ) + _kill_process_group() + reader.join(_KILLED_PIPE_TIMEOUT_S) + + returncode = process.wait() + + with contextlib.suppress(Exception): + process.stdout.close() # type: ignore[union-attr] + return returncode, "".join(chunks) def run_example_command( diff --git a/tests/unit/test_example_run_command.py b/tests/unit/test_example_run_command.py new file mode 100644 index 000000000..09a1f6818 --- /dev/null +++ b/tests/unit/test_example_run_command.py @@ -0,0 +1,57 @@ +# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Tests for the example-command runner shared by the example tests.""" + +import contextlib +import os +import signal +import time + +import pytest +from _test_utils.examples import run_command + + +def test_run_capturing_does_not_block_on_a_survivor_holding_the_pipe( + skip_on_windows, monkeypatch, tmp_path +): + """A killed command whose descendant inherited the output pipe must not hang the caller.""" + monkeypatch.setattr(run_command, "_ORPHAN_PIPE_TIMEOUT_S", 1) + pid_file = tmp_path / "survivor.pid" + + started = time.monotonic() + with pytest.warns(UserWarning, match="descendants holding its output pipe"): + returncode, output = run_command._run_capturing( + ["bash", "-c", f"sleep 60 & echo $! > {pid_file}; echo out; sleep 0.3; kill -9 $$"], + tmp_path, + os.environ.copy(), + ) + + survivor = int(pid_file.read_text()) + try: + assert returncode == -9 + assert "out" in output # captured despite the survivor + assert time.monotonic() - started < 10 # the patched grace period is 1s + + for _ in range(50): # the group kill is asynchronous + try: + os.kill(survivor, 0) + except ProcessLookupError: + break + time.sleep(0.1) + else: + pytest.fail(f"survivor {survivor} outlived _run_capturing") + finally: + with contextlib.suppress(ProcessLookupError): + os.kill(survivor, signal.SIGKILL)