From ae7c5ab1712a02ebdf83dcc74f891db3409d7d53 Mon Sep 17 00:00:00 2001 From: jschen069 <3563624058@qq.com> Date: Wed, 5 Aug 2026 23:11:32 +0800 Subject: [PATCH] =?UTF-8?q?=E6=B7=BB=E5=8A=A0ifbench=E5=AF=B9nltk=E5=8C=85?= =?UTF-8?q?=E7=9A=84=E5=88=A4=E6=96=AD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../ifbench/ifbench_0_shot_gen_str.py | 1 + .../datasets/ifbench/instructions_util.py | 35 ++-- ais_bench/benchmark/runners/local.py | 47 ++++- .../benchmark/utils/logging/error_codes.py | 3 + tests/UT/runners/test_local.py | 163 ++++++++++++++++++ 5 files changed, 236 insertions(+), 13 deletions(-) diff --git a/ais_bench/benchmark/configs/datasets/ifbench/ifbench_0_shot_gen_str.py b/ais_bench/benchmark/configs/datasets/ifbench/ifbench_0_shot_gen_str.py index 3b54073e..7d43ba4f 100644 --- a/ais_bench/benchmark/configs/datasets/ifbench/ifbench_0_shot_gen_str.py +++ b/ais_bench/benchmark/configs/datasets/ifbench/ifbench_0_shot_gen_str.py @@ -26,6 +26,7 @@ abbr='ifbench', type=IFBenchDataset, path='ais_bench/datasets/ifbench/data/train-00000-of-00001.parquet', + nltk_path='/path/to/nltk_data', reader_cfg=ifbench_reader_cfg, infer_cfg=ifbench_infer_cfg, eval_cfg=ifbench_eval_cfg, diff --git a/ais_bench/benchmark/datasets/ifbench/instructions_util.py b/ais_bench/benchmark/datasets/ifbench/instructions_util.py index 9e621423..706941ca 100644 --- a/ais_bench/benchmark/datasets/ifbench/instructions_util.py +++ b/ais_bench/benchmark/datasets/ifbench/instructions_util.py @@ -5,18 +5,33 @@ import re +_NLTK_RESOURCE_PATHS = { + 'punkt_tab': 'tokenizers/punkt_tab', + 'averaged_perceptron_tagger_eng': + 'taggers/averaged_perceptron_tagger_eng', + 'stopwords': 'corpora/stopwords', +} + + def _ensure_nltk_data(data_name): - """Ensure NLTK data is available, downloading if necessary.""" + """Ensure the requested NLTK resource is available.""" + resource_path = _NLTK_RESOURCE_PATHS.get( + data_name, f'tokenizers/{data_name}' + ) + try: - nltk.data.find(f'tokenizers/{data_name}') + nltk.data.find(resource_path) except (LookupError, OSError): - try: - nltk.download(data_name, quiet=True) - except (LookupError, OSError): - # Fall back to downloading the full 'punkt' package if the - # specific resource (e.g. 'punkt_tab') is not available as a - # standalone download in the installed NLTK version. - nltk.download('punkt', quiet=True) + # 在线环境下仍可自动下载;离线环境预装正确后不会进入这里 + downloaded = nltk.download(data_name, quiet=True) + if not downloaded: + raise LookupError( + f"Missing NLTK resource: {resource_path}. " + f"Please install it under an NLTK_DATA directory." + ) + + # 下载后再检查,避免静默失败 + nltk.data.find(resource_path) WORD_LIST = [ @@ -248,8 +263,6 @@ def _ensure_nltk_data(data_name): 'injury', 'insect', 'surprise', 'apartment', ] # pylint: disable=line-too-long -_ensure_nltk_data('punkt_tab') - _ALPHABETS = '([A-Za-z])' _PREFIXES = '(Mr|St|Mrs|Ms|Dr)[.]' _SUFFIXES = '(Inc|Ltd|Jr|Sr|Co)' diff --git a/ais_bench/benchmark/runners/local.py b/ais_bench/benchmark/runners/local.py index 76391d89..2b8125eb 100644 --- a/ais_bench/benchmark/runners/local.py +++ b/ais_bench/benchmark/runners/local.py @@ -66,6 +66,46 @@ def __init__(self, for k, v in kwargs.items(): self.logger.warning(f'Ignored argument in {self.__module__}: {k}={v}') + def _get_subprocess_env(self, task) -> Dict[str, str]: + env = os.environ.copy() + configured_paths = [] + for dataset_cfg in getattr(task, 'dataset_cfgs', []): + if 'nltk_path' not in dataset_cfg: + continue + raw_path = dataset_cfg['nltk_path'] + if not isinstance(raw_path, str) or not raw_path.strip(): + message = ( + f"Task {task.name} has invalid nltk_path: {raw_path!r}; " + "expected a non-empty string") + self.logger.error(message) + raise ParameterValueError( + RUNNER_CODES.INVALID_NLTK_PATH, message) + path = osp.abspath(osp.expandvars(osp.expanduser(raw_path.strip()))) + configured_paths.append(path) + + distinct_paths = {osp.normcase(path) for path in configured_paths} + if len(distinct_paths) > 1: + message = ( + f"Task {task.name} has conflicting nltk_path values: " + f"{configured_paths}") + self.logger.error(message) + raise ParameterValueError(RUNNER_CODES.INVALID_NLTK_PATH, message) + if not configured_paths: + return env + + nltk_path = configured_paths[0] + if not osp.isdir(nltk_path) or not os.access(nltk_path, os.R_OK): + message = ( + f"Task {task.name} has unusable nltk_path: {nltk_path}; " + "the path must be an existing readable directory") + self.logger.error(message) + raise ParameterValueError(RUNNER_CODES.INVALID_NLTK_PATH, message) + + env['NLTK_DATA'] = nltk_path + self.logger.info( + f"Task {task.name} sets NLTK_DATA={nltk_path} for its subprocess") + return env + def launch(self, tasks: List[Dict[str, Any]]) -> List[Tuple[str, int]]: """Launch multiple tasks. @@ -151,7 +191,8 @@ def _run_debug(self, tasks: List[Dict[str, Any]], all_gpu_ids: List[int], monito tmpl = get_command_template(all_gpu_ids[:num_gpus]) cmd = task.get_command(cfg_path=param_file, template=tmpl) - proc = subprocess.Popen(cmd, shell=True, text=True) + env = self._get_subprocess_env(task) + proc = subprocess.Popen(cmd, shell=True, text=True, env=env) try: proc.wait() except KeyboardInterrupt: @@ -253,12 +294,14 @@ def _launch(self, task, gpu_ids, index): # Run command out_path = task.get_log_path(file_extension='out') mmengine.mkdir_or_exist(osp.split(out_path)[0]) + env = self._get_subprocess_env(task) with open(out_path, 'w', encoding='utf-8') as stdout: result = subprocess.run(cmd, shell=True, text=True, stdout=stdout, - stderr=stdout) + stderr=stdout, + env=env) if result.returncode != 0: self.logger.error(RUNNER_CODES.TASK_FAILED, f"{task_name} failed with code {result.returncode}, see\n{out_path}") finally: diff --git a/ais_bench/benchmark/utils/logging/error_codes.py b/ais_bench/benchmark/utils/logging/error_codes.py index a1361de2..8e5cec91 100644 --- a/ais_bench/benchmark/utils/logging/error_codes.py +++ b/ais_bench/benchmark/utils/logging/error_codes.py @@ -116,6 +116,9 @@ class SUMM_CODES: class RUNNER_CODES: UNKNOWN_ERROR = BaseErrorCode("RUNNER-UNK-001", ErrorModule.RUNNER, ErrorType.UNKNOWN, 1, "unknown error of runner") TASK_FAILED = BaseErrorCode("RUNNER-TASK-001", ErrorModule.RUNNER, ErrorType.TASK, 1, "task failed") # docs coverd + INVALID_NLTK_PATH = BaseErrorCode( + 'RUNNER-PARAM-001', ErrorModule.RUNNER, ErrorType.PARAM, 1, + 'invalid NLTK data path') class TMON_CODES: diff --git a/tests/UT/runners/test_local.py b/tests/UT/runners/test_local.py index 5c9a5b6c..8e5813be 100644 --- a/tests/UT/runners/test_local.py +++ b/tests/UT/runners/test_local.py @@ -6,6 +6,14 @@ from mmengine.config import ConfigDict from ais_bench.benchmark.runners.local import LocalRunner, get_command_template +from ais_bench.benchmark.utils.logging.exceptions import ParameterValueError + + +def _task_with_datasets(*datasets): + task = MagicMock() + task.name = 'test_task' + task.dataset_cfgs = [ConfigDict(dataset) for dataset in datasets] + return task class TestGetCommandTemplate(unittest.TestCase): @@ -64,6 +72,104 @@ def setUp(self): self.max_num_workers = 4 self.max_workers_per_gpu = 1 + @patch('ais_bench.benchmark.runners.base.AISLogger') + def test_get_subprocess_env_sets_nltk_data_without_mutating_parent( + self, mock_logger_class): + """A valid dataset path is private to the child environment.""" + runner = LocalRunner(task=self.task_cfg) + with tempfile.TemporaryDirectory() as nltk_dir: + parent_before = os.environ.get('NLTK_DATA') + env = runner._get_subprocess_env( + _task_with_datasets({'nltk_path': nltk_dir})) + + self.assertEqual(env['NLTK_DATA'], os.path.abspath(nltk_dir)) + self.assertEqual(os.environ.get('NLTK_DATA'), parent_before) + mock_logger_class.return_value.info.assert_called_once() + mock_logger_class.return_value.error.assert_not_called() + + @patch('ais_bench.benchmark.runners.base.AISLogger') + def test_get_subprocess_env_preserves_inherited_value_when_unconfigured( + self, mock_logger_class): + """Tasks without the setting retain the parent process value.""" + runner = LocalRunner(task=self.task_cfg) + with patch.dict(os.environ, {'NLTK_DATA': 'inherited-value'}): + env = runner._get_subprocess_env(_task_with_datasets({})) + self.assertEqual(os.environ.get('NLTK_DATA'), 'inherited-value') + + self.assertEqual(env['NLTK_DATA'], 'inherited-value') + mock_logger_class.return_value.info.assert_not_called() + mock_logger_class.return_value.error.assert_not_called() + + @patch('ais_bench.benchmark.runners.base.AISLogger') + def test_get_subprocess_env_rejects_conflicting_paths( + self, mock_logger_class): + """Distinct dataset paths cannot be selected for one child task.""" + runner = LocalRunner(task=self.task_cfg) + with tempfile.TemporaryDirectory() as first, \ + tempfile.TemporaryDirectory() as second: + with self.assertRaises(ParameterValueError) as context: + runner._get_subprocess_env(_task_with_datasets( + {'nltk_path': first}, {'nltk_path': second})) + + self.assertEqual(context.exception.error_code_str, 'RUNNER-PARAM-001') + mock_logger_class.return_value.error.assert_called_once() + mock_logger_class.return_value.info.assert_not_called() + + @patch('ais_bench.benchmark.runners.base.AISLogger') + def test_get_subprocess_env_expands_environment_path(self, + mock_logger_class): + """Configured paths support environment-variable expansion.""" + runner = LocalRunner(task=self.task_cfg) + with tempfile.TemporaryDirectory() as nltk_dir: + with patch.dict( + os.environ, + {'AISBENCH_NLTK_TEST_PATH': nltk_dir}): + env = runner._get_subprocess_env(_task_with_datasets( + {'nltk_path': '%AISBENCH_NLTK_TEST_PATH%'})) + + self.assertEqual(env['NLTK_DATA'], os.path.abspath(nltk_dir)) + mock_logger_class.return_value.info.assert_called_once() + mock_logger_class.return_value.error.assert_not_called() + + @patch('ais_bench.benchmark.runners.base.AISLogger') + def test_get_subprocess_env_rejects_invalid_configured_path( + self, mock_logger_class): + """Malformed and unusable configured paths are rejected and logged.""" + runner = LocalRunner(task=self.task_cfg) + missing = os.path.join( + tempfile.gettempdir(), 'aisbench-missing-nltk-data') + invalid_values = [None, '', 123, missing] + + for value in invalid_values: + with self.subTest(value=value): + with self.assertRaises(ParameterValueError) as context: + runner._get_subprocess_env( + _task_with_datasets({'nltk_path': value})) + self.assertEqual( + context.exception.error_code_str, 'RUNNER-PARAM-001') + + with tempfile.NamedTemporaryFile() as regular_file: + with self.assertRaises(ParameterValueError) as context: + runner._get_subprocess_env( + _task_with_datasets({'nltk_path': regular_file.name})) + self.assertEqual(context.exception.error_code_str, 'RUNNER-PARAM-001') + + with tempfile.TemporaryDirectory() as unreadable: + with patch( + 'ais_bench.benchmark.runners.local.os.access', + return_value=False): + with self.assertRaises(ParameterValueError) as context: + runner._get_subprocess_env( + _task_with_datasets({'nltk_path': unreadable})) + self.assertEqual( + context.exception.error_code_str, 'RUNNER-PARAM-001') + + self.assertEqual( + mock_logger_class.return_value.error.call_count, + len(invalid_values) + 2, + ) + mock_logger_class.return_value.info.assert_not_called() + @patch('ais_bench.benchmark.runners.base.AISLogger') def test_init(self, mock_logger_class): """测试LocalRunner初始化""" @@ -212,6 +318,36 @@ def test_launch_with_visible_devices(self, mock_environ, mock_logger_class): call_args = str(mock_logger.debug.call_args_list) self.assertIn("Available devices", call_args) + @patch('ais_bench.benchmark.runners.base.AISLogger') + @patch('ais_bench.benchmark.runners.local.TASKS') + def test_run_debug_passes_nltk_environment(self, mock_tasks, + mock_logger_class): + """Debug subprocess receives the configured NLTK data path.""" + runner = LocalRunner(task=self.task_cfg, debug=True) + mock_task = MagicMock() + mock_task.name = 'test_task' + mock_task.num_gpus = 0 + mock_task.get_command.return_value = 'python test.py' + mock_task.cfg.dump = MagicMock() + mock_tasks.build.return_value = mock_task + + with tempfile.TemporaryDirectory() as nltk_dir: + mock_task.dataset_cfgs = [ConfigDict({'nltk_path': nltk_dir})] + with patch( + 'ais_bench.benchmark.runners.local.subprocess.Popen' + ) as subprocess_mock, patch( + 'ais_bench.benchmark.runners.local.os.remove' + ), patch( + 'ais_bench.benchmark.runners.local.mmengine.mkdir_or_exist' + ), patch('uuid.uuid4', return_value=MagicMock(hex='test')): + subprocess_mock.return_value.wait.return_value = None + runner._run_debug([{'work_dir': '/tmp/test', 'cli_args': {}}], + [], MagicMock()) + + self.assertEqual( + subprocess_mock.call_args.kwargs['env']['NLTK_DATA'], + os.path.abspath(nltk_dir)) + @patch('ais_bench.benchmark.runners.base.AISLogger') @patch('ais_bench.benchmark.runners.local.TASKS') def test_run_debug(self, mock_tasks, mock_logger_class): @@ -257,6 +393,33 @@ def test_run_debug(self, mock_tasks, mock_logger_class): self.assertEqual(status[0][0], "test_task") self.assertEqual(status[0][1], 0) + @patch('ais_bench.benchmark.runners.base.AISLogger') + def test_launch_passes_nltk_environment(self, mock_logger_class): + """Normal subprocess receives the configured NLTK data path.""" + runner = LocalRunner(task=self.task_cfg, debug=False) + mock_task = MagicMock() + mock_task.name = 'test_task' + mock_task.get_command.return_value = 'python test.py' + mock_task.get_log_path.return_value = '/tmp/test.out' + mock_task.cfg.dump = MagicMock() + + with tempfile.TemporaryDirectory() as nltk_dir: + mock_task.dataset_cfgs = [ConfigDict({'nltk_path': nltk_dir})] + with patch( + 'ais_bench.benchmark.runners.local.subprocess.run' + ) as subprocess_mock, patch( + 'ais_bench.benchmark.runners.local.os.remove' + ), patch( + 'ais_bench.benchmark.runners.local.mmengine.mkdir_or_exist' + ), patch('uuid.uuid4', return_value=MagicMock(hex='test')), patch( + 'builtins.open', create=True): + subprocess_mock.return_value.returncode = 0 + runner._launch(mock_task, [0], 0) + + self.assertEqual( + subprocess_mock.call_args.kwargs['env']['NLTK_DATA'], + os.path.abspath(nltk_dir)) + @patch('ais_bench.benchmark.runners.base.AISLogger') def test_launch_method(self, mock_logger_class): """测试_launch方法"""