Creates a dictionary of task objects from either a name of task, config, or prepared Task object. :param task_name_list: List[Union[str, Dict, Task]] Name of model or LM object, see lm_eval.models.get_model :param task_manager: TaskManager = None A TaskManager object that st
(task_name_list: List[Union[str, Dict, Task]], task_manager: TaskManager = None)
| 383 | ) |
| 384 | |
| 385 | def get_task_dict(task_name_list: List[Union[str, Dict, Task]], task_manager: TaskManager = None): |
| 386 | """Creates a dictionary of task objects from either a name of task, config, or prepared Task object. |
| 387 | |
| 388 | :param task_name_list: List[Union[str, Dict, Task]] |
| 389 | Name of model or LM object, see lm_eval.models.get_model |
| 390 | :param task_manager: TaskManager = None |
| 391 | A TaskManager object that stores indexed tasks. If not set, |
| 392 | task_manager will load one. This should be set by the user |
| 393 | if there are additional paths that want to be included |
| 394 | via `include_path` |
| 395 | |
| 396 | :return |
| 397 | Dictionary of task objects |
| 398 | """ |
| 399 | task_name_from_string_dict = {} |
| 400 | task_name_from_config_dict = {} |
| 401 | task_name_from_object_dict = {} |
| 402 | |
| 403 | if isinstance(task_name_list, str): |
| 404 | task_name_list = [task_name_list] |
| 405 | |
| 406 | string_task_name_list = [task for task in task_name_list if isinstance(task, str)] |
| 407 | others_task_name_list = [task for task in task_name_list if ~isinstance(task, str)] |
| 408 | if len(string_task_name_list) > 0: |
| 409 | if task_manager is None: |
| 410 | task_manager = TaskManager() |
| 411 | |
| 412 | task_name_from_string_dict = task_manager.load_task_or_group(string_task_name_list) |
| 413 | |
| 414 | for task_element in others_task_name_list: |
| 415 | if isinstance(task_element, dict): |
| 416 | task_name_from_config_dict = { |
| 417 | **task_name_from_config_dict, |
| 418 | **task_manager.load_config(config=task_element), |
| 419 | } |
| 420 | |
| 421 | elif isinstance(task_element, Task): |
| 422 | task_name_from_object_dict = { |
| 423 | **task_name_from_object_dict, |
| 424 | get_task_name_from_object(task_element): task_element, |
| 425 | } |
| 426 | |
| 427 | assert set(task_name_from_string_dict.keys()).isdisjoint( |
| 428 | set(task_name_from_object_dict.keys()) |
| 429 | ) |
| 430 | return { |
| 431 | **task_name_from_string_dict, |
| 432 | **task_name_from_config_dict, |
| 433 | **task_name_from_object_dict, |
| 434 | } |
no test coverage detected