fix(ci): recover from corrupted MMMU parquet cache (#17256)

This commit is contained in:
Hudson Xing
2026-01-17 17:32:01 -08:00
committed by GitHub
parent 53609e5e5b
commit 90399cbc07
+76 -10
View File
@@ -1,8 +1,10 @@
import glob import glob
import json import json
import os import os
import shutil
import subprocess import subprocess
import tempfile import tempfile
from pathlib import Path
from types import SimpleNamespace from types import SimpleNamespace
from sglang.srt.environ import temp_set_env from sglang.srt.environ import temp_set_env
@@ -18,6 +20,78 @@ from sglang.test.test_utils import (
DEFAULT_MEM_FRACTION_STATIC = 0.8 DEFAULT_MEM_FRACTION_STATIC = 0.8
def _is_mmmu_parquet_corruption(error_output: str) -> bool:
"""Check if error is due to MMMU parquet file corruption."""
return (
"ArrowInvalid" in error_output
and "Parquet magic bytes not found" in error_output
and ("MMMU" in error_output or "lmms-lab--MMMU" in error_output)
)
def _cleanup_mmmu_dataset_cache():
"""Clean up corrupted MMMU dataset cache to allow fresh download."""
# Priority 1: Check CI convention path /hf_home first (used in Docker containers)
ci_hf_home = Path("/hf_home/hub/datasets--lmms-lab--MMMU")
if ci_hf_home.exists():
mmmu_cache_path = ci_hf_home
else:
# Priority 2: Use HF_HOME env var or default user cache
hf_home = os.environ.get("HF_HOME", os.path.expanduser("~/.cache/huggingface"))
mmmu_cache_path = Path(hf_home) / "hub" / "datasets--lmms-lab--MMMU"
if mmmu_cache_path.exists():
print(f"Detected corrupted MMMU parquet cache. Cleaning up: {mmmu_cache_path}")
try:
shutil.rmtree(mmmu_cache_path)
print(f"Successfully removed corrupted cache: {mmmu_cache_path}")
return True
except OSError as e:
print(f"Warning: Failed to remove cache {mmmu_cache_path}: {e}")
return False
else:
print(f"MMMU cache not found at {mmmu_cache_path}, skipping cleanup")
return False
def _run_lmms_eval_with_retry(cmd: list[str], timeout: int = 3600) -> None:
"""Run lmms_eval command with automatic retry on MMMU parquet corruption."""
try:
result = subprocess.run(
cmd,
check=True,
timeout=timeout,
capture_output=True,
text=True,
)
# Print captured output to maintain visibility of successful runs
if result.stdout:
print(result.stdout, end="")
if result.stderr:
print(result.stderr, end="")
except subprocess.CalledProcessError as e:
error_output = e.stderr + e.stdout
if _is_mmmu_parquet_corruption(error_output):
print("Detected MMMU parquet corruption error. Attempting recovery...")
if _cleanup_mmmu_dataset_cache():
print("Retrying lmms_eval with fresh download...")
with temp_set_env(
HF_HUB_OFFLINE="0",
HF_DATASETS_DOWNLOAD_MODE="force_redownload",
):
subprocess.run(cmd, check=True, timeout=timeout)
else:
print(
f"Failed to cleanup corrupted MMMU cache. Error from lmms_eval:\nStdout:\n{e.stdout}\nStderr:\n{e.stderr}"
)
raise
else:
print(
f"lmms_eval failed with an unhandled error.\nStdout:\n{e.stdout}\nStderr:\n{e.stderr}"
)
raise
class MMMUMixin: class MMMUMixin:
"""Mixin for MMMU evaluation. """Mixin for MMMU evaluation.
@@ -81,11 +155,7 @@ class MMMUMixin:
OPENAI_API_KEY=self.api_key, OPENAI_API_KEY=self.api_key,
OPENAI_API_BASE=f"{self.base_url}/v1", OPENAI_API_BASE=f"{self.base_url}/v1",
): ):
subprocess.run( _run_lmms_eval_with_retry(cmd)
cmd,
check=True,
timeout=3600,
)
def test_mmmu(self: CustomTestCase): def test_mmmu(self: CustomTestCase):
"""Run MMMU evaluation test.""" """Run MMMU evaluation test."""
@@ -209,11 +279,7 @@ class MMMUMultiModelTestBase(CustomTestCase):
*self.mmmu_args, *self.mmmu_args,
] ]
subprocess.run( _run_lmms_eval_with_retry(cmd)
cmd,
check=True,
timeout=3600,
)
def _run_vlm_mmmu_test( def _run_vlm_mmmu_test(
self, self,