File size: 5,607 Bytes
d91766b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 | from __future__ import annotations
import argparse
import json
import os
import re
import shutil
import sys
import tempfile
import time
from pathlib import Path
from typing import Optional
from examples.engine_lm_eval.config import BenchmarkConfig
try:
from lm_eval.__main__ import cli_evaluate
except ImportError:
cli_evaluate = None
from examples.engine_lm_eval.lm_eval_model import EngineOpenAILM # noqa: F401
def _slugify(value: str) -> str:
value = re.sub(r"[^A-Za-z0-9._-]+", "_", value.strip())
value = re.sub(r"_+", "_", value).strip("_.")
return value or "unknown"
def _resolve_include_path(config: BenchmarkConfig) -> Optional[Path]:
raw = config.eval.include_path
if raw is not None and str(raw).strip() == "":
return None
if raw:
p = Path(raw).expanduser()
if not p.is_absolute():
p = Path(os.getcwd()) / p
return p.resolve()
return (Path(__file__).resolve().parents[2] / "diffulex_bench" / "tasks").resolve()
def _task_name_to_yaml_map(include_root: Path) -> dict[str, Path]:
mapping: dict[str, Path] = {}
for yml in include_root.rglob("*.yaml"):
try:
text = yml.read_text(encoding="utf-8")
except Exception:
continue
m = re.search(r"(?m)^\s*task:\s*([^\s#]+)\s*$", text)
if m:
mapping.setdefault(m.group(1).strip(), yml)
return mapping
def _rewrite_task_data_files(task_yaml: Path, data_files: str) -> bool:
text = task_yaml.read_text(encoding="utf-8")
replacement_value = json.dumps(str(Path(data_files).expanduser().resolve()))
replaced, n = re.subn(r"(?m)^(\s*data_files:\s*).*$", rf"\1{replacement_value}", text, count=1)
if n == 0:
return False
task_yaml.write_text(replaced, encoding="utf-8")
return True
def _resolve_include_path_with_override(config: BenchmarkConfig) -> tuple[Optional[Path], Optional[Path]]:
include_path = _resolve_include_path(config)
data_files = config.eval.dataset_data_files
if not data_files:
return include_path, None
if include_path is None or not include_path.is_dir():
return include_path, None
tmp_root = Path(tempfile.mkdtemp(prefix="engine_lm_eval_tasks_")).resolve()
tmp_tasks = tmp_root / "tasks"
shutil.copytree(include_path, tmp_tasks, dirs_exist_ok=True)
task_map = _task_name_to_yaml_map(tmp_tasks)
for task_name in [name.strip() for name in config.eval.dataset_name.split(",") if name.strip()]:
task_yaml = task_map.get(task_name)
if task_yaml is not None:
_rewrite_task_data_files(task_yaml, data_files)
return tmp_tasks, tmp_root
def _run_output_dir(config: BenchmarkConfig) -> Path:
engine_name = _slugify(config.endpoint.engine_name)
root = Path(config.eval.output_dir).expanduser() / engine_name
if not config.eval.use_run_subdirectory:
root.mkdir(parents=True, exist_ok=True)
return root.resolve()
stamp = time.strftime("%Y%m%d_%H%M%S")
model_name = _slugify(Path(config.endpoint.model).name)
task_name = _slugify(config.eval.dataset_name.replace(",", "_"))
out = root / f"run_{stamp}_{engine_name}_{model_name}_{task_name}"
out.mkdir(parents=True, exist_ok=True)
return out.resolve()
def config_to_model_args(config: BenchmarkConfig, save_dir: Path) -> str:
args = {
"engine_name": config.endpoint.engine_name,
"base_url": config.endpoint.base_url,
"model": config.endpoint.model,
"api_key": config.endpoint.api_key,
"tokenizer_path": config.endpoint.tokenizer_path or config.endpoint.model,
"batch_size": 1,
"max_new_tokens": config.eval.max_tokens,
"temperature": config.eval.temperature,
"ignore_eos": config.eval.ignore_eos,
"add_bos_token": config.eval.add_bos_token if config.eval.add_bos_token is not None else False,
"apply_chat_template": config.endpoint.apply_chat_template,
"chat_completions": config.endpoint.chat_completions,
"timeout": config.endpoint.timeout,
"verify": config.endpoint.verify,
"save_dir": str(save_dir),
"trust_remote_code": config.endpoint.trust_remote_code,
}
return ",".join(f"{k}={v}" for k, v in args.items() if v is not None)
def main() -> None:
parser = argparse.ArgumentParser(description="Minimal lm-eval runner backed by an OpenAI-compatible API engine")
parser.add_argument("--config", required=True, help="Path to YAML config.")
args = parser.parse_args()
if cli_evaluate is None:
raise RuntimeError("lm-evaluation-harness is not installed.")
config = BenchmarkConfig.from_yaml(args.config)
output_dir = _run_output_dir(config)
include_path, cleanup_root = _resolve_include_path_with_override(config)
model_args = config_to_model_args(config, output_dir)
output_file = output_dir / f"{_slugify(config.endpoint.engine_name)}_lm_eval_results.json"
sys.argv = [
"lm_eval",
"--model",
"engine_oai",
"--model_args",
model_args,
"--tasks",
config.eval.dataset_name,
"--output_path",
str(output_file),
]
if include_path is not None:
sys.argv.extend(["--include_path", str(include_path)])
if config.eval.dataset_limit:
sys.argv.extend(["--limit", str(config.eval.dataset_limit)])
try:
cli_evaluate()
finally:
if cleanup_root is not None:
shutil.rmtree(cleanup_root, ignore_errors=True)
if __name__ == "__main__":
main()
|