Skip to content

Commit 33c903f

Browse files
authored
[Update] Add BBH cascade evaluation config (#2553)
* Add BBH cascade evaluation config * Remove BBH inferencer output limits
1 parent b7d18f4 commit 33c903f

5 files changed

Lines changed: 280 additions & 98 deletions

File tree

Lines changed: 91 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,91 @@
1+
import os
2+
from opencompass.openicl.icl_prompt_template import PromptTemplate
3+
from opencompass.openicl.icl_retriever import ZeroRetriever
4+
from opencompass.openicl.icl_inferencer import GenInferencer
5+
from opencompass.openicl.icl_evaluator import AccEvaluator
6+
from opencompass.datasets import BBHDataset, BBHEvaluator, bbh_mcq_postprocess, BBHEvaluator_mcq
7+
8+
bbh_reader_cfg = dict(input_columns=['input'], output_column='target')
9+
10+
bbh_multiple_choice_sets = [
11+
'temporal_sequences',
12+
'disambiguation_qa',
13+
'date_understanding',
14+
'tracking_shuffled_objects_three_objects',
15+
'penguins_in_a_table',
16+
'geometric_shapes',
17+
'snarks',
18+
'ruin_names',
19+
'tracking_shuffled_objects_seven_objects',
20+
'tracking_shuffled_objects_five_objects',
21+
'logical_deduction_three_objects',
22+
'hyperbaton',
23+
'logical_deduction_five_objects',
24+
'logical_deduction_seven_objects',
25+
'movie_recommendation',
26+
'salient_translation_error_detection',
27+
'reasoning_about_colored_objects',
28+
]
29+
bbh_free_form_sets = [
30+
'multistep_arithmetic_two',
31+
'navigate',
32+
'dyck_languages',
33+
'word_sorting',
34+
'sports_understanding',
35+
'boolean_expressions',
36+
'object_counting',
37+
'formal_fallacies',
38+
'causal_judgement',
39+
'web_of_lies',
40+
]
41+
42+
bbh_datasets = []
43+
for _name in bbh_multiple_choice_sets:
44+
bbh_infer_cfg = dict(prompt_template=dict(
45+
type=PromptTemplate,
46+
template=dict(round=[
47+
dict(
48+
role='HUMAN',
49+
prompt=
50+
f"Follow the given examples and answer the question.\n\nQuestion: {{input}}\n You must give your final answer by starting with 'So the answer is' "
51+
)
52+
])),
53+
retriever=dict(type=ZeroRetriever),
54+
inferencer=dict(type=GenInferencer))
55+
bbh_eval_cfg = dict(evaluator=dict(type=BBHEvaluator_mcq),
56+
pred_role='BOT',
57+
pred_postprocessor=dict(type=bbh_mcq_postprocess),
58+
dataset_postprocessor=dict(type=bbh_mcq_postprocess))
59+
60+
bbh_datasets.append(
61+
dict(type=BBHDataset,
62+
path='opencompass/bbh',
63+
name=_name,
64+
abbr='bbh-' + _name,
65+
reader_cfg=bbh_reader_cfg,
66+
infer_cfg=bbh_infer_cfg.copy(),
67+
eval_cfg=bbh_eval_cfg.copy()))
68+
69+
for _name in bbh_free_form_sets:
70+
71+
bbh_infer_cfg = dict(prompt_template=dict(
72+
type=PromptTemplate,
73+
template=dict(round=[
74+
dict(
75+
role='HUMAN',
76+
prompt=
77+
f"Follow the given examples and answer the question.\n\nQuestion: {{input}}\n You must give your final answer by starting with 'So the answer is' "
78+
)
79+
])),
80+
retriever=dict(type=ZeroRetriever),
81+
inferencer=dict(type=GenInferencer))
82+
bbh_eval_cfg = dict(evaluator=dict(type=BBHEvaluator), pred_role='BOT')
83+
84+
bbh_datasets.append(
85+
dict(type=BBHDataset,
86+
path='opencompass/bbh',
87+
name=_name,
88+
abbr='bbh-' + _name,
89+
reader_cfg=bbh_reader_cfg,
90+
infer_cfg=bbh_infer_cfg.copy(),
91+
eval_cfg=bbh_eval_cfg.copy()))

opencompass/configs/datasets/bbh/bbh_0shot_nocot_gen_9c32f6.py

Lines changed: 0 additions & 96 deletions
This file was deleted.
Lines changed: 114 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,114 @@
1+
"""
2+
Summary: A cascade evaluation config for BBH free-form tasks.
3+
Setting:
4+
Shot: 0-shot
5+
Evaluator:
6+
- CascadeEvaluator
7+
- BBHEvaluator
8+
- GenericLLMEvaluator
9+
"""
10+
11+
from opencompass.datasets import (BBHDataset, BBHEvaluator,
12+
generic_llmjudge_postprocess)
13+
from opencompass.evaluator import CascadeEvaluator, GenericLLMEvaluator
14+
from opencompass.openicl.icl_inferencer import GenInferencer
15+
from opencompass.openicl.icl_raw_prompt_template import RawPromptTemplate
16+
from opencompass.openicl.icl_retriever import ZeroRetriever
17+
18+
bbh_reader_cfg = dict(input_columns=['input'], output_column='target')
19+
20+
bbh_free_form_sets = [
21+
'multistep_arithmetic_two',
22+
'navigate',
23+
'dyck_languages',
24+
'word_sorting',
25+
'sports_understanding',
26+
'boolean_expressions',
27+
'object_counting',
28+
'formal_fallacies',
29+
'causal_judgement',
30+
'web_of_lies',
31+
]
32+
33+
bbh_infer_cfg = dict(
34+
prompt_template=dict(
35+
type=RawPromptTemplate,
36+
messages=[{
37+
'role':
38+
'user',
39+
'content':
40+
'Follow the given examples and answer the question.\n\n'
41+
'Question: {input}\n You must give your final answer by '
42+
"starting with 'So the answer is' "
43+
}],
44+
),
45+
retriever=dict(type=ZeroRetriever),
46+
inferencer=dict(type=GenInferencer),
47+
)
48+
49+
GRADER_TEMPLATE = """
50+
Please judge whether the predicted answer is consistent with the gold target
51+
for the original question. The gold target is correct. Do not solve the
52+
question yourself. Treat differences in capitalization, punctuation, or list
53+
separators as equivalent when they express the same answer. A prediction with
54+
additional explanation is correct only if its final answer agrees with the
55+
gold target.
56+
57+
Reply with exactly one letter:
58+
A: CORRECT
59+
B: INCORRECT
60+
61+
<Original Question Begin>
62+
{input}
63+
<Original Question End>
64+
65+
<Gold Target Begin>
66+
{target}
67+
<Gold Target End>
68+
69+
<Predicted Answer Begin>
70+
{prediction}
71+
<Predicted Answer End>
72+
""".strip()
73+
74+
cascade_evaluator = dict(
75+
type=CascadeEvaluator,
76+
rule_evaluator=dict(type=BBHEvaluator),
77+
llm_evaluator=dict(
78+
type=GenericLLMEvaluator,
79+
prompt_template=dict(
80+
type=RawPromptTemplate,
81+
messages=[
82+
{
83+
'role': 'system',
84+
'content': 'You are a precise answer evaluator.'
85+
},
86+
{
87+
'role': 'user',
88+
'content': GRADER_TEMPLATE
89+
},
90+
],
91+
),
92+
dataset_cfg=dict(
93+
type=BBHDataset,
94+
path='opencompass/bbh',
95+
reader_cfg=bbh_reader_cfg,
96+
),
97+
judge_cfg=dict(),
98+
dict_postprocessor=dict(type=generic_llmjudge_postprocess),
99+
),
100+
parallel=False,
101+
)
102+
103+
bbh_datasets = []
104+
for _name in bbh_free_form_sets:
105+
bbh_datasets.append(
106+
dict(
107+
type=BBHDataset,
108+
path='opencompass/bbh',
109+
name=_name,
110+
abbr='bbh-' + _name,
111+
reader_cfg=bbh_reader_cfg,
112+
infer_cfg=bbh_infer_cfg.copy(),
113+
eval_cfg=dict(evaluator=cascade_evaluator.copy()),
114+
))

opencompass/evaluator/cascade_evaluator.py

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
import inspect
12
import os
23
from typing import Any, Callable, Dict, List, Optional
34

@@ -86,8 +87,13 @@ def sample_score(self,
8687
else:
8788
# Use rule_evaluator to evaluate a single sample by calling
8889
# the score method with single-element lists
89-
result = self.rule_evaluator.score([prediction], [reference],
90-
[test_set])
90+
score_params = inspect.signature(
91+
self.rule_evaluator.score).parameters
92+
if 'test_set' in score_params:
93+
result = self.rule_evaluator.score([prediction], [reference],
94+
test_set=[test_set])
95+
else:
96+
result = self.rule_evaluator.score([prediction], [reference])
9197
if 'details' in result and len(result['details']) > 0:
9298
return result['details'][0]
9399
else:
Lines changed: 67 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,67 @@
1+
import unittest
2+
3+
from opencompass.evaluator.cascade_evaluator import CascadeEvaluator
4+
5+
6+
class RuleEvaluatorWithoutTestSet:
7+
8+
def __init__(self):
9+
self.calls = []
10+
11+
def score(self, predictions, references):
12+
self.calls.append((predictions, references))
13+
return {
14+
'details': [{
15+
'pred': predictions[0],
16+
'answer': references[0],
17+
'correct': True,
18+
}]
19+
}
20+
21+
22+
class RuleEvaluatorWithTestSet(RuleEvaluatorWithoutTestSet):
23+
24+
def score(self, predictions, references, test_set=None):
25+
self.calls.append((predictions, references, test_set))
26+
return {
27+
'details': [{
28+
'pred': predictions[0],
29+
'answer': references[0],
30+
'correct': True,
31+
}]
32+
}
33+
34+
35+
class TestCascadeEvaluator(unittest.TestCase):
36+
37+
def _make_evaluator(self, rule_evaluator):
38+
evaluator = CascadeEvaluator.__new__(CascadeEvaluator)
39+
evaluator.sample_score_fn = None
40+
evaluator.rule_evaluator = rule_evaluator
41+
return evaluator
42+
43+
def test_sample_score_without_test_set_argument(self):
44+
rule_evaluator = RuleEvaluatorWithoutTestSet()
45+
evaluator = self._make_evaluator(rule_evaluator)
46+
47+
result = evaluator.sample_score('prediction', 'reference',
48+
{'input': 'question'})
49+
50+
self.assertTrue(result['correct'])
51+
self.assertEqual(rule_evaluator.calls,
52+
[(['prediction'], ['reference'])])
53+
54+
def test_sample_score_with_test_set_argument(self):
55+
rule_evaluator = RuleEvaluatorWithTestSet()
56+
evaluator = self._make_evaluator(rule_evaluator)
57+
58+
test_item = {'input': 'question'}
59+
result = evaluator.sample_score('prediction', 'reference', test_item)
60+
61+
self.assertTrue(result['correct'])
62+
self.assertEqual(rule_evaluator.calls,
63+
[(['prediction'], ['reference'], [test_item])])
64+
65+
66+
if __name__ == '__main__':
67+
unittest.main()

0 commit comments

Comments
 (0)