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)