Initial commit
This commit is contained in:
@@ -0,0 +1,258 @@
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import re
|
||||
|
||||
from .grammar_definition import flatten, _one_text_field
|
||||
from .parsing_supernatural_instructions_tasks import SUPERNATURAL_INSTRUCTIONS_TASKS_WITH_NO_FORMAT, \
|
||||
create_initial_structured_prompt_format
|
||||
|
||||
DEFAULT_SUPERNATURAL_INSTRUCTIONS_DIRECTORY = '../natural-instructions/tasks'
|
||||
DEFAULT_INSTRUCTION_INDUCTION_DIRECTORY = '../instruction-induction'
|
||||
|
||||
STRING_ALL_CHARACTERS_FOR_REGEX_MATCHING = r"""([A-Za-z0-9α-ωΑ-Ω“”‘’′`,.…'-–—−:∶()\[\]{}/%?!\" ;$≤≥≠†€₹→≡~∨⊃·°•∃∀ʻ&⁄_#\n𝑆𝑚√𝑠𝑁𝐴𝑒𝑅𝑇ι⟩⟨›‹ου‖♥‰�龍►➥™,‚∼⋅]+)"""
|
||||
random.seed(0)
|
||||
|
||||
|
||||
def extract_regex(prompt_format):
|
||||
prompt_format_original = prompt_format.replace('<|text|>', '<text>') # pipe cannot be used for regex
|
||||
regex_sentence_extractor_str = re.escape(prompt_format_original).replace(
|
||||
'<text>', STRING_ALL_CHARACTERS_FOR_REGEX_MATCHING)
|
||||
regex_sentence_extractor_str = '^' + regex_sentence_extractor_str + '$'
|
||||
regex_sentence_extractor = re.compile(regex_sentence_extractor_str)
|
||||
return regex_sentence_extractor
|
||||
|
||||
|
||||
def _extract_fields_from_dataset(regex_sentence_extractor_dict, dataset, num_samples):
|
||||
input_fields_list = []
|
||||
outputs_list = []
|
||||
|
||||
# tells us which key in regex_sentence_extractor_dict matched, useful for knowing
|
||||
# which format version (with number of enumerations) to apply later
|
||||
regex_key_idx_list = []
|
||||
selected_ids = []
|
||||
|
||||
for i, entry in enumerate(dataset):
|
||||
if len(input_fields_list) == num_samples:
|
||||
break
|
||||
|
||||
# we skip data points that we could not parse:
|
||||
# sometimes even in the same task, the spacing is not respected (probably due to manual errors)
|
||||
# note: we process possible regexes from longest to shortest, because often a template with two fields would
|
||||
# match a string that actually has five fields
|
||||
input_fields, regex_key_idx = None, None
|
||||
for regex_key_idx, regex_sentence_extractor in sorted(regex_sentence_extractor_dict.items(), reverse=True):
|
||||
input_fields = re.search(regex_sentence_extractor, entry['input'])
|
||||
if input_fields:
|
||||
break
|
||||
if not input_fields:
|
||||
print(f"WARNING: data point {i} ({entry['input']}) was not able to be processed.")
|
||||
print('CHARACTERS USED:', [e for e in set(entry['input']) if not re.match(STRING_ALL_CHARACTERS_FOR_REGEX_MATCHING, e)])
|
||||
continue
|
||||
|
||||
input_fields = input_fields.groups()
|
||||
input_fields_list.append(input_fields)
|
||||
regex_key_idx_list.append(regex_key_idx)
|
||||
|
||||
outputs_list.append(entry['output'])
|
||||
selected_ids.append(i)
|
||||
|
||||
return input_fields_list, outputs_list, regex_key_idx_list, selected_ids
|
||||
|
||||
|
||||
def _load_raw_dataset_supernatural_instructions(args):
|
||||
# find filename based on task_filename
|
||||
dataset_directory = args.natural_instructions_dir
|
||||
if not os.path.isdir(dataset_directory):
|
||||
raise FileNotFoundError(
|
||||
f'Natural Instructions tasks directory not found: {dataset_directory}. '
|
||||
'Clone https://github.com/allenai/natural-instructions beside this project, '
|
||||
'or pass --natural_instructions_dir /path/to/natural-instructions/tasks.')
|
||||
task_filenames = [f for f in os.listdir(dataset_directory) if args.task_filename in f]
|
||||
assert len(task_filenames) == 1, f"Expected exactly one task matching {args.task_filename!r}; found {task_filenames}"
|
||||
task_filename = task_filenames[0]
|
||||
|
||||
filepath = os.path.join(dataset_directory, task_filename)
|
||||
raw_dataset = json.load(open(filepath, 'r'))
|
||||
return raw_dataset
|
||||
|
||||
|
||||
def set_up_prompt_variation_exploration_without_extra_files(
|
||||
args,
|
||||
structured_prompt_format,
|
||||
extra_params_structured_prompt_format,
|
||||
instruction=None
|
||||
):
|
||||
"""
|
||||
Mel notes: currently
|
||||
choosing demonstrations;
|
||||
loading dataset;
|
||||
potentially adding "answer" field; create
|
||||
regex extracting fields
|
||||
"""
|
||||
|
||||
raw_dataset = _load_raw_dataset_supernatural_instructions(args)
|
||||
demonstration_definition = raw_dataset['Definition'][0] if instruction is None else instruction
|
||||
|
||||
raw_dataset = raw_dataset['Instances']
|
||||
if hasattr(args, 'dataset_ordered_ids') and args.dataset_ordered_ids:
|
||||
assert len(args.dataset_ordered_ids) == len(raw_dataset)
|
||||
raw_dataset = [raw_dataset[i] for i in args.dataset_ordered_ids]
|
||||
else:
|
||||
random.shuffle(raw_dataset)
|
||||
|
||||
demonstrations = raw_dataset[:10]
|
||||
dataset = [entry for entry in raw_dataset[10:]]
|
||||
|
||||
if extra_params_structured_prompt_format and extra_params_structured_prompt_format.get('enumeration_length_range'):
|
||||
regex_sentence_extractor_dict = {}
|
||||
for e in range(*extra_params_structured_prompt_format.get('enumeration_length_range')):
|
||||
prompt_format_original = flatten(structured_prompt_format.solve({'enumeration_length': e}))
|
||||
regex_sentence_extractor_dict[e] = extract_regex(prompt_format_original)
|
||||
else:
|
||||
regex_sentence_extractor = extract_regex(flatten(structured_prompt_format.solve()))
|
||||
regex_sentence_extractor_dict = {None: regex_sentence_extractor} # None because there is no length
|
||||
|
||||
return demonstration_definition, dataset, regex_sentence_extractor_dict, demonstrations, len(raw_dataset)
|
||||
|
||||
|
||||
def setup_demonstrations(args, regex_sentence_extractor_dict, demonstrations):
|
||||
|
||||
demos_fields_list, demonstrations_outputs, demos_regex_key_idx_list, _ = _extract_fields_from_dataset(
|
||||
regex_sentence_extractor_dict, demonstrations, num_samples=args.n_shot)
|
||||
|
||||
if len(demos_fields_list) != args.n_shot:
|
||||
print("Insufficient n-shot demos.")
|
||||
print(len(demos_fields_list))
|
||||
assert False, f"{len(demos_fields_list)} != {args.n_shot}"
|
||||
exit(1)
|
||||
|
||||
file_suffix = ''
|
||||
return demos_fields_list, demonstrations_outputs, demos_regex_key_idx_list, file_suffix
|
||||
|
||||
|
||||
def load_supernatural_instructions_task(args):
|
||||
"""
|
||||
All logic for loading the dataset, extracting the original formatting from the text.
|
||||
|
||||
PRECOMPUTE
|
||||
1. Load model and tokenizer (OK)
|
||||
2. Detect regex to extract fields from dataset (currently from external file, but it could be from the initial structure)
|
||||
3. Extract formatting from dataset (keep a set of fields)
|
||||
4. Extract desired few shot examples and extract their formatting (keep a set of fields)
|
||||
|
||||
Args params needed:
|
||||
|
||||
args.task_filename
|
||||
args.num_samples
|
||||
args.n_shot
|
||||
Plus the ones needed for uses of args_compute_node_score
|
||||
"""
|
||||
|
||||
# SuperNaturalInstructions Tasks without a defined format
|
||||
if any(t in args.task_filename for t in SUPERNATURAL_INSTRUCTIONS_TASKS_WITH_NO_FORMAT):
|
||||
raw_dataset = _load_raw_dataset_supernatural_instructions(args)
|
||||
demonstration_definition = raw_dataset['Definition'][0]
|
||||
raw_dataset = raw_dataset['Instances']
|
||||
return _setup_non_formatted_dataset_with_one_field_only(args, raw_dataset, demonstration_definition)
|
||||
|
||||
# Parse Formatted SuperNaturalInstructions Tasks
|
||||
structured_prompt_format, global_constraints, extra_params_structured_prompt_format, \
|
||||
instruction, original_multiple_choice_output_format = create_initial_structured_prompt_format(args)
|
||||
demonstration_definition, dataset, regex_sentence_extractor_dict, demonstrations, raw_dataset_size = \
|
||||
set_up_prompt_variation_exploration_without_extra_files(
|
||||
args, structured_prompt_format, extra_params_structured_prompt_format, instruction)
|
||||
demonstration_definition = demonstration_definition if instruction is None else instruction
|
||||
|
||||
input_fields_list, _, regex_key_idx_list, selected_dataset_ids = _extract_fields_from_dataset(
|
||||
regex_sentence_extractor_dict, dataset, num_samples=args.num_samples)
|
||||
|
||||
demos_fields_list, demonstrations_outputs, demos_regex_key_idx_list, demonstrations_filename_suffix = \
|
||||
setup_demonstrations(args, regex_sentence_extractor_dict, demonstrations)
|
||||
|
||||
args_compute_node_score = {
|
||||
'args': args,
|
||||
'dataset': dataset,
|
||||
'input_fields_list': input_fields_list,
|
||||
'regex_key_idx_list': regex_key_idx_list, # tells us which of the options of enumeration quantities applies
|
||||
'selected_dataset_ids': selected_dataset_ids,
|
||||
'demos_fields_list': demos_fields_list,
|
||||
'demonstrations_outputs': demonstrations_outputs,
|
||||
'demos_regex_key_idx_list': demos_regex_key_idx_list,
|
||||
# tells us which of the options of enumeration quantities applies
|
||||
'demonstration_definition': demonstration_definition,
|
||||
}
|
||||
|
||||
return structured_prompt_format, global_constraints, extra_params_structured_prompt_format, \
|
||||
original_multiple_choice_output_format, args_compute_node_score, raw_dataset_size
|
||||
|
||||
|
||||
def _setup_non_formatted_dataset_with_one_field_only(args, raw_dataset, demonstration_definition):
|
||||
# set up initial formatting
|
||||
structured_prompt_format, global_constraints = _one_text_field('Input', answer_field_text='Output', chosen_space='\n')
|
||||
extra_params_structured_prompt_format = None
|
||||
original_multiple_choice_output_format = None
|
||||
|
||||
if hasattr(args, 'dataset_ordered_ids') and args.dataset_ordered_ids:
|
||||
assert len(args.dataset_ordered_ids) == len(raw_dataset)
|
||||
raw_dataset = [raw_dataset[i] for i in args.dataset_ordered_ids]
|
||||
else:
|
||||
random.shuffle(raw_dataset)
|
||||
|
||||
# set up dataset & demonstrations with the same fields and formatting as SuperNatural Instructions
|
||||
demonstrations = raw_dataset[:10]
|
||||
dataset = [entry for entry in raw_dataset[10:]]
|
||||
|
||||
demos_fields_list = [tuple([example['input']]) for example in demonstrations][:args.n_shot]
|
||||
demonstrations_outputs = [example['output'] for example in demonstrations][:args.n_shot]
|
||||
|
||||
input_fields_list = [tuple([example['input']]) for example in dataset][:args.num_samples]
|
||||
selected_dataset_ids = list(range(len(input_fields_list)))
|
||||
|
||||
args_compute_node_score = {
|
||||
'args': args,
|
||||
'dataset': dataset,
|
||||
'input_fields_list': input_fields_list,
|
||||
'regex_key_idx_list': [None] * len(input_fields_list), # setting to None because there is only one format option (no enumeration length variation)
|
||||
'selected_dataset_ids': selected_dataset_ids,
|
||||
'demos_fields_list': demos_fields_list,
|
||||
'demonstrations_outputs': demonstrations_outputs,
|
||||
'demos_regex_key_idx_list': [None] * len(demonstrations_outputs), # setting to None because there is only one format option (no enumeration length variation)
|
||||
'demonstration_definition': demonstration_definition,
|
||||
}
|
||||
|
||||
return structured_prompt_format, global_constraints, extra_params_structured_prompt_format, \
|
||||
original_multiple_choice_output_format, args_compute_node_score, len(raw_dataset)
|
||||
|
||||
|
||||
def load_instruction_induction_task(args):
|
||||
"""
|
||||
This dataset doesn't have a pre-defined format to extract like SuperNatural Instructions.
|
||||
We will use the formatting that APE has used as a starting point.
|
||||
|
||||
We'll generate the equivalent structures as the ones generated in SuperNaturalInstructions.
|
||||
|
||||
Instructions: https://github.com/orhonovich/instruction-induction/blob/main/data/annotations/antonyms.json
|
||||
I-O: https://github.com/orhonovich/instruction-induction/tree/main/data/raw/induce
|
||||
"""
|
||||
|
||||
# load datasets
|
||||
# task_filename = f"{task_name}.json"
|
||||
dataset_directory = args.instruction_induction_dir
|
||||
if not os.path.isdir(dataset_directory):
|
||||
raise FileNotFoundError(
|
||||
f'Instruction Induction directory not found: {dataset_directory}. '
|
||||
'Clone https://github.com/orhonovich/instruction-induction beside this project, '
|
||||
'or pass --instruction_induction_dir /path/to/instruction-induction.')
|
||||
instructions = json.load(open(os.path.join(dataset_directory, 'data', 'annotations', args.task_filename), 'r'))
|
||||
instructions = instructions['annotations']
|
||||
print('instructions', instructions)
|
||||
|
||||
raw_dataset = json.load(open(os.path.join(dataset_directory, 'data', 'raw', 'induce', args.task_filename), 'r'))
|
||||
raw_dataset = list(raw_dataset['examples'].values())
|
||||
raw_dataset = [{'input': entry['input'], 'output': [entry['output']]} for entry in raw_dataset]
|
||||
|
||||
# chose best instruction with some criterion (long, is properly cased to begin with)
|
||||
demonstration_definition = sorted([inst for inst in instructions if inst[0].isupper()], reverse=True, key=len)[0]
|
||||
|
||||
return _setup_non_formatted_dataset_with_one_field_only(args, raw_dataset, demonstration_definition)
|
||||
@@ -0,0 +1,505 @@
|
||||
import copy
|
||||
import random
|
||||
from typing import List
|
||||
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from .grammar_definition import pointers_to_all_objects, create_pointer_action_type_pairs, \
|
||||
flatten, MAPPING_ALL_CATEGORIES, holistic_node_format_sanity_checks
|
||||
from .utils import evaluate_prompt_format
|
||||
|
||||
random.seed(0)
|
||||
|
||||
|
||||
def value_assignment_str_to_indices(value_assignments, pointer_action_pairs):
|
||||
value_assignments_ids = []
|
||||
for assignment in value_assignments:
|
||||
assert len(pointer_action_pairs) == len(assignment), f"{len(pointer_action_pairs)} != {len(assignment)}"
|
||||
assignment_ids = []
|
||||
for (_, _, action_type), assignment_value in zip(pointer_action_pairs, assignment):
|
||||
idx = [i for i, (_, v) in enumerate(MAPPING_ALL_CATEGORIES[action_type]) if v == assignment_value][0]
|
||||
assignment_ids.append(idx)
|
||||
value_assignments_ids.append(assignment_ids)
|
||||
return value_assignments_ids
|
||||
|
||||
|
||||
class GeneticAlgorithmAmongPrompts:
|
||||
|
||||
def __init__(self,
|
||||
structured_prompt_format,
|
||||
global_constraints,
|
||||
extra_params_structured_prompt_format,
|
||||
args_compute_node_score,
|
||||
objective,
|
||||
allow_text_action_type=True,
|
||||
original_multiple_choice_output_format=None):
|
||||
self.args_compute_node_score = args_compute_node_score
|
||||
self.metadata = {}
|
||||
self.all_structured_prompt_formats_last_id_evaluated = {}
|
||||
self.all_structured_prompt_formats_accuracies = {} # actually has the accuracies computed
|
||||
self.objective = objective
|
||||
self.extra_params_structured_prompt_format = extra_params_structured_prompt_format
|
||||
self.original_multiple_choice_output_format = original_multiple_choice_output_format
|
||||
|
||||
# nodes (prompt formats) are represented by their solved_format
|
||||
solved_format = self._get_node_from_format(structured_prompt_format)
|
||||
self.all_structured_prompt_formats = {
|
||||
solved_format: [structured_prompt_format, global_constraints] # nodes
|
||||
}
|
||||
|
||||
# all multiple choice classes in the original format, important to know how to update them when format changes
|
||||
original_multiple_choice_classes = self.find_all_multiple_choice_output_classes(
|
||||
solved_format, original_multiple_choice_output_format)
|
||||
self.original_multiple_choice_classes = original_multiple_choice_classes
|
||||
|
||||
self.generation_order = {solved_format: 0}
|
||||
self.edges = []
|
||||
self.allow_text_action_type = allow_text_action_type
|
||||
|
||||
self.metadata = {}
|
||||
self.metadata['extra_params'] = {'allow_text_action_type': self.allow_text_action_type}
|
||||
self.metadata['nodes'] = {} # used in some extensions of this class
|
||||
self.metadata['bit_representations'] = {} # used in some extensions of this class
|
||||
|
||||
self.all_structured_prompt_formats_accuracies = {
|
||||
solved_format: self._compute_node_score(structured_prompt_format, num_samples_to_test=-1)
|
||||
}
|
||||
self.metadata['bit_representations'][solved_format] = [None] # None = no actions have been done yet
|
||||
self.metadata['extra_params']['objective'] = self.objective
|
||||
|
||||
all_pointers = pointers_to_all_objects(structured_prompt_format) + global_constraints
|
||||
all_pointers_enumerated = [(e, i) for i, e in enumerate(all_pointers)]
|
||||
pointer_action_pairs = create_pointer_action_type_pairs(
|
||||
all_pointers_enumerated, allow_text_action_type=self.allow_text_action_type)
|
||||
self.initial_structured_prompt_format = structured_prompt_format
|
||||
self.initial_global_constraints = global_constraints
|
||||
self.pointer_action_pairs = pointer_action_pairs
|
||||
|
||||
action_value_options = []
|
||||
for a, b, action_type in pointer_action_pairs:
|
||||
action_value_options.append(range(len(MAPPING_ALL_CATEGORIES[action_type])))
|
||||
self.action_value_options = action_value_options
|
||||
|
||||
def find_all_multiple_choice_output_classes(self, resolved_node_format, output_format):
|
||||
if not output_format:
|
||||
return []
|
||||
|
||||
# output_format = "Option {enum1}", where "enum1" is the object name
|
||||
object_name = output_format.split('{')[1].split('}')[0]
|
||||
|
||||
structured_prompt_format, global_constraints = self.all_structured_prompt_formats[resolved_node_format]
|
||||
all_pointers = pointers_to_all_objects(structured_prompt_format) + global_constraints
|
||||
pointer_to_object_list = [pointer
|
||||
for pointer in all_pointers
|
||||
if 'object_name' in pointer.__dict__ and pointer.object_name == object_name]
|
||||
assert len(pointer_to_object_list) == 1
|
||||
pointer_to_object = pointer_to_object_list[0]
|
||||
return [output_format.format(**{object_name: pointer_to_object.chosen_number_format(idx)})
|
||||
for idx in pointer_to_object.enumeration_item_id_list]
|
||||
|
||||
def _get_node_from_format(self, prompt_format):
|
||||
extra_params = {'print_output_fields': True, 'exclude_text_field_for_output_fields': False}
|
||||
return flatten(prompt_format.solve(extra_params)).replace('<|text|>', '{}')
|
||||
|
||||
def _copy_objects_before_expanding_node(self, solved_format):
|
||||
# this function creates a copy of the passed format node (solved formats)
|
||||
# this prevents accidentally modifying the previous node when searching a tree of prompt formats
|
||||
|
||||
structured_prompt_format, global_constraints = self.all_structured_prompt_formats[solved_format]
|
||||
structured_prompt_format, global_constraints = copy.deepcopy((structured_prompt_format, global_constraints))
|
||||
|
||||
all_pointers = pointers_to_all_objects(structured_prompt_format) + global_constraints
|
||||
all_pointers_enumerated = [(e, i) for i, e in enumerate(all_pointers)]
|
||||
|
||||
if 'all_pointers_enumerated' not in self.metadata:
|
||||
self.metadata['all_pointers_enumerated'] = [
|
||||
(str(type(e).__name__), self._get_node_from_format(e) if e.solve() else list(e.fields.keys())) for e, i
|
||||
in all_pointers_enumerated
|
||||
]
|
||||
|
||||
return structured_prompt_format, global_constraints, all_pointers_enumerated
|
||||
|
||||
def list_node_accuracies(self):
|
||||
return sorted([(v, k,
|
||||
flatten(self.all_structured_prompt_formats[k][0].solve({'print_output_fields': True})).replace(
|
||||
'<|text|>', '{}'))
|
||||
for k, v in self.all_structured_prompt_formats_accuracies.items()], reverse=True)
|
||||
|
||||
def save(self, filename, previous_result=None):
|
||||
"""Persist evaluation state, preserving checkpointed formats from an earlier run."""
|
||||
import json
|
||||
to_dump = {
|
||||
# 'all_structured_prompt_formats': self.all_structured_prompt_formats,
|
||||
'generation_order': self.generation_order,
|
||||
'edges': self.edges,
|
||||
'all_structured_prompt_formats_accuracies': self.all_structured_prompt_formats_accuracies,
|
||||
'metadata': self.metadata
|
||||
}
|
||||
|
||||
if previous_result:
|
||||
for key in ('generation_order', 'all_structured_prompt_formats_accuracies'):
|
||||
merged = dict(previous_result.get(key, {}))
|
||||
merged.update(to_dump[key])
|
||||
to_dump[key] = merged
|
||||
to_dump['edges'] = previous_result.get('edges', []) + to_dump['edges']
|
||||
|
||||
previous_metadata = previous_result.get('metadata', {})
|
||||
for key in ('nodes', 'bit_representations'):
|
||||
merged = dict(previous_metadata.get(key, {}))
|
||||
merged.update(to_dump['metadata'].get(key, {}))
|
||||
to_dump['metadata'][key] = merged
|
||||
merged_extra_params = dict(previous_metadata.get('extra_params', {}))
|
||||
merged_extra_params.update(to_dump['metadata'].get('extra_params', {}))
|
||||
to_dump['metadata']['extra_params'] = merged_extra_params
|
||||
|
||||
json.dump(to_dump, open(filename, 'w'))
|
||||
|
||||
def _compute_node_score_from_resolved_prompt(self, resolved_prompt, num_samples_to_test=-1):
|
||||
last_id_analyzed = self.all_structured_prompt_formats_last_id_evaluated.get(resolved_prompt, 0)
|
||||
interval_ids_to_test = (last_id_analyzed, last_id_analyzed + num_samples_to_test) \
|
||||
if num_samples_to_test != -1 and last_id_analyzed is not None \
|
||||
else (None, None)
|
||||
|
||||
# transform the multiple choice output classes to evaluate in the same format as the examples presented
|
||||
current_multiple_choice_classes = self.find_all_multiple_choice_output_classes(
|
||||
resolved_prompt, self.original_multiple_choice_output_format)
|
||||
original_to_current_multiple_choice_classes = \
|
||||
{k: v for k, v in zip(self.original_multiple_choice_classes, current_multiple_choice_classes)} \
|
||||
if self.original_multiple_choice_classes else {}
|
||||
|
||||
structured_prompt_format, global_constraints = self.all_structured_prompt_formats[resolved_prompt]
|
||||
acc, history = evaluate_prompt_format(
|
||||
**self.args_compute_node_score,
|
||||
structured_prompt_format=structured_prompt_format,
|
||||
original_to_current_multiple_choice_classes=original_to_current_multiple_choice_classes,
|
||||
interval_ids_to_test=interval_ids_to_test
|
||||
)
|
||||
self.all_structured_prompt_formats_last_id_evaluated[resolved_prompt] = interval_ids_to_test[1]
|
||||
self.all_structured_prompt_formats_accuracies[resolved_prompt] = acc
|
||||
|
||||
self.metadata['nodes'][resolved_prompt] = history
|
||||
return acc
|
||||
|
||||
def _compute_node_score(self, structured_prompt_format, num_samples_to_test=-1):
|
||||
# return (0, 0, 0), [0]
|
||||
return self._compute_node_score_from_resolved_prompt(
|
||||
resolved_prompt=self._get_node_from_format(structured_prompt_format),
|
||||
num_samples_to_test=num_samples_to_test)
|
||||
|
||||
def evaluate_node(self, solution, num_samples_to_test):
|
||||
|
||||
# copy structured_prompt_format to avoid modifying the original
|
||||
resolved_prompt = self._get_node_from_format(self.initial_structured_prompt_format)
|
||||
structured_prompt_format, global_constraints, all_pointers_enumerated = \
|
||||
self._copy_objects_before_expanding_node(resolved_prompt)
|
||||
pointer_action_pairs = create_pointer_action_type_pairs(
|
||||
all_pointers_enumerated, allow_text_action_type=self.allow_text_action_type)
|
||||
assert len(self.pointer_action_pairs) == len(pointer_action_pairs)
|
||||
assert all([b == e and c == f for (a, b, c), (d, e, f) in zip(self.pointer_action_pairs, pointer_action_pairs)])
|
||||
|
||||
# transform action value ids into a new structured_prompt_format
|
||||
all_action_values = []
|
||||
all_action_value_names = []
|
||||
for (element, element_id, action_type), action_value_id in zip(pointer_action_pairs, solution):
|
||||
action_value, action_value_name = MAPPING_ALL_CATEGORIES[action_type][int(action_value_id)]
|
||||
all_action_values.append(action_value)
|
||||
all_action_value_names.append(action_value_name)
|
||||
element.update_field(action_type, action_value)
|
||||
|
||||
# check if value assignments are invalid, and if so give the worst possible accuracy and do not store logs about it
|
||||
# importantly, we do not store self.generation_order
|
||||
if not holistic_node_format_sanity_checks(structured_prompt_format):
|
||||
return -1e6 * (-1 if self.objective == 'lowest_accuracy' else 1)
|
||||
|
||||
# update logs that do not require accuracy
|
||||
new_node = self._get_node_from_format(structured_prompt_format)
|
||||
if new_node in self.generation_order:
|
||||
self.metadata['bit_representations'][new_node].append(all_action_value_names)
|
||||
acc = self.all_structured_prompt_formats_accuracies[new_node]
|
||||
return acc[0] * (-1 if self.objective == 'lowest_accuracy' else 1)
|
||||
|
||||
self.metadata['bit_representations'][new_node] = [all_action_value_names]
|
||||
self.all_structured_prompt_formats[new_node] = [structured_prompt_format, global_constraints]
|
||||
self.generation_order[new_node] = len(self.generation_order)
|
||||
|
||||
# compute accuracy and update accuracy logs
|
||||
acc = self._compute_node_score(structured_prompt_format, num_samples_to_test)
|
||||
|
||||
self.all_structured_prompt_formats_accuracies[new_node] = acc
|
||||
|
||||
return acc[0] * (-1 if self.objective == 'lowest_accuracy' else 1)
|
||||
|
||||
def main(self, value_assignments: List[List[str]], num_samples_to_test: int,
|
||||
skip_value_assignments=None, on_node_evaluated=None):
|
||||
"""
|
||||
Fully evaluate all nodes (prompt formats) passed.
|
||||
|
||||
:param value_assignments: Value assignments for each format, and each field of the format.
|
||||
value_assignments[i] shows all strings representing each field value for the i-th sampled format.
|
||||
:param num_samples_to_test: number of samples to consider a node fully evaluated
|
||||
"""
|
||||
|
||||
# convert from list(list(str)) to list(list(int))
|
||||
# this func assumes same order as in action_value_pairs, but in text (not id in array, to be robust to changes)
|
||||
value_assignments_ids = value_assignment_str_to_indices(value_assignments, self.pointer_action_pairs)
|
||||
|
||||
# Run all nodes. A checkpoint records value assignments (rather than
|
||||
# internal node objects), so a later process can reconstruct and skip
|
||||
# completed formats safely.
|
||||
skip_value_assignments = skip_value_assignments or set()
|
||||
progress = tqdm(
|
||||
zip(value_assignments, value_assignments_ids),
|
||||
total=len(value_assignments),
|
||||
desc='Evaluating format variants',
|
||||
unit='format',
|
||||
dynamic_ncols=True,
|
||||
)
|
||||
for value_assignment, value_assignment_ids in progress:
|
||||
if tuple(value_assignment) in skip_value_assignments:
|
||||
progress.set_postfix_str('cached')
|
||||
continue
|
||||
progress.set_postfix_str('running samples')
|
||||
self.evaluate_node(value_assignment_ids, num_samples_to_test)
|
||||
if on_node_evaluated:
|
||||
on_node_evaluated(value_assignment)
|
||||
progress.set_postfix_str('checkpoint saved')
|
||||
progress.close()
|
||||
|
||||
|
||||
class ThompsonSamplingAlgorithmAmongPrompts(GeneticAlgorithmAmongPrompts):
|
||||
|
||||
def _compute_node_score_from_resolved_prompt(self, resolved_prompt, num_samples_to_test=-1):
|
||||
last_id_analyzed = self.all_structured_prompt_formats_last_id_evaluated.get(resolved_prompt, 0)
|
||||
interval_ids_to_test = (last_id_analyzed, last_id_analyzed + num_samples_to_test) \
|
||||
if num_samples_to_test != -1 and last_id_analyzed is not None \
|
||||
else (None, None)
|
||||
|
||||
if last_id_analyzed is not None and num_samples_to_test == -1:
|
||||
interval_ids_to_test = (last_id_analyzed, None)
|
||||
|
||||
if last_id_analyzed is None and num_samples_to_test == -1:
|
||||
print("This means we already evaluated all samples, returning empty results.")
|
||||
return (0, 0, 0)
|
||||
|
||||
if len(self.args_compute_node_score['selected_dataset_ids'][interval_ids_to_test[0]:interval_ids_to_test[1]]) == 0:
|
||||
print("This means we already evaluated all samples, returning empty results.")
|
||||
return (0, 0, 0)
|
||||
|
||||
# transform the multiple choice output classes to evaluate in the same format as the examples presented
|
||||
current_multiple_choice_classes = self.find_all_multiple_choice_output_classes(
|
||||
resolved_prompt, self.original_multiple_choice_output_format)
|
||||
original_to_current_multiple_choice_classes = \
|
||||
{k: v for k, v in zip(self.original_multiple_choice_classes, current_multiple_choice_classes)} \
|
||||
if self.original_multiple_choice_classes else {}
|
||||
|
||||
structured_prompt_format, global_constraints = self.all_structured_prompt_formats[resolved_prompt]
|
||||
acc, history = evaluate_prompt_format(
|
||||
**self.args_compute_node_score,
|
||||
structured_prompt_format=structured_prompt_format,
|
||||
original_to_current_multiple_choice_classes=original_to_current_multiple_choice_classes,
|
||||
interval_ids_to_test=interval_ids_to_test
|
||||
)
|
||||
self.all_structured_prompt_formats_last_id_evaluated[resolved_prompt] = interval_ids_to_test[1]
|
||||
if resolved_prompt not in self.metadata['nodes']:
|
||||
self.metadata['nodes'][resolved_prompt] = []
|
||||
self.metadata['nodes'][resolved_prompt].extend(history)
|
||||
return acc
|
||||
|
||||
def _add_node_to_structures(self, solution):
|
||||
"""
|
||||
This initializes nodes in our structures. It's easier to add them all at the beginning
|
||||
and then only care about sampling.
|
||||
"""
|
||||
|
||||
# copy structured_prompt_format to avoid modifying the original
|
||||
resolved_prompt = self._get_node_from_format(self.initial_structured_prompt_format)
|
||||
structured_prompt_format, global_constraints, all_pointers_enumerated = \
|
||||
self._copy_objects_before_expanding_node(resolved_prompt)
|
||||
pointer_action_pairs = create_pointer_action_type_pairs(
|
||||
all_pointers_enumerated, allow_text_action_type=self.allow_text_action_type)
|
||||
assert len(self.pointer_action_pairs) == len(pointer_action_pairs)
|
||||
assert all([b == e and c == f for (a, b, c), (d, e, f) in zip(self.pointer_action_pairs, pointer_action_pairs)])
|
||||
|
||||
# transform action value ids into a new structured_prompt_format
|
||||
all_action_values = []
|
||||
all_action_value_names = []
|
||||
for (element, element_id, action_type), action_value_id in zip(pointer_action_pairs, solution):
|
||||
action_value, action_value_name = MAPPING_ALL_CATEGORIES[action_type][int(action_value_id)]
|
||||
all_action_values.append(action_value)
|
||||
all_action_value_names.append(action_value_name)
|
||||
element.update_field(action_type, action_value)
|
||||
|
||||
# invalid node, give the worst possible accuracy and do not store logs about it
|
||||
# especially do not store self.generation_order
|
||||
if not holistic_node_format_sanity_checks(structured_prompt_format):
|
||||
assert False, "This should not happen because this is run from a file already filtered."
|
||||
|
||||
# update logs that do not require accuracy
|
||||
new_node = self._get_node_from_format(structured_prompt_format)
|
||||
if new_node in self.generation_order:
|
||||
self.metadata['bit_representations'][new_node].append(all_action_value_names)
|
||||
return None
|
||||
|
||||
self.metadata['bit_representations'][new_node] = [all_action_value_names]
|
||||
self.all_structured_prompt_formats[new_node] = [structured_prompt_format, global_constraints]
|
||||
self.generation_order[new_node] = len(self.generation_order)
|
||||
self.all_structured_prompt_formats_accuracies[new_node] = (0, 0, 0) # list of CUMULATIVE accuracies
|
||||
|
||||
return new_node
|
||||
|
||||
def _evaluate_node_on_batch(self, new_node, num_samples):
|
||||
"""
|
||||
Evaluates new_node for num_samples (i.e. one batch).
|
||||
"""
|
||||
structured_prompt_format, global_constraints = self.all_structured_prompt_formats[new_node]
|
||||
acc = self._compute_node_score(structured_prompt_format, num_samples) # (right [0, 1], wrong [0, 1], total)
|
||||
new_batch_right, new_batch_wrong, new_batch_total = acc
|
||||
right, wrong, total = self.all_structured_prompt_formats_accuracies[new_node]
|
||||
|
||||
cumulative_wrong_counter = wrong * total + new_batch_wrong * new_batch_total
|
||||
cumulative_right_counter = right * total + new_batch_right * new_batch_total
|
||||
cumulative_total = new_batch_total + total
|
||||
cumulative_right = cumulative_right_counter / cumulative_total
|
||||
cumulative_wrong = cumulative_wrong_counter / cumulative_total
|
||||
self.all_structured_prompt_formats_accuracies[new_node] = (cumulative_right, cumulative_wrong, cumulative_total)
|
||||
|
||||
return cumulative_total, cumulative_right_counter
|
||||
|
||||
def _choose_final_node(self, num_successes, total_elements_evaluated, objective, nodes_sampled):
|
||||
accuracy_nodes = [(num_successes[node] / total_elements_evaluated[node], node) for node in nodes_sampled
|
||||
if total_elements_evaluated[node] > 0]
|
||||
accuracy_nodes = sorted(accuracy_nodes, reverse=(objective == 'highest'))
|
||||
return accuracy_nodes[0][-1]
|
||||
|
||||
def _evaluate_nodes_thompson_sampling(
|
||||
self,
|
||||
original_node,
|
||||
nodes_sampled,
|
||||
batch_size,
|
||||
max_allowed_number_of_steps=100,
|
||||
objective='lowest',
|
||||
use_ucb_rule=False,
|
||||
num_successes=None,
|
||||
total_elements_evaluated=None):
|
||||
import numpy as np
|
||||
|
||||
if num_successes is None or total_elements_evaluated is None:
|
||||
total_elements_evaluated = {k: 0 for k in nodes_sampled}
|
||||
num_successes = {k: 0 for k in nodes_sampled}
|
||||
|
||||
right, wrong, total = self.all_structured_prompt_formats_accuracies[original_node]
|
||||
total_elements_evaluated[original_node], num_successes[original_node] = total, right * total
|
||||
upper_bound_worst_node_accuracy = num_successes[original_node] / total_elements_evaluated[original_node]
|
||||
num_samples_in_dataset = total_elements_evaluated[original_node]
|
||||
|
||||
# using EV=initial_node, we know that: a * (1 - initial_node) = initial_node * b. We initialize with b=5
|
||||
# we also avoid non-bell shape curves
|
||||
b = 5
|
||||
a = upper_bound_worst_node_accuracy / (1 - upper_bound_worst_node_accuracy) * b
|
||||
a = max(a, 1.1)
|
||||
initial_a_b_params = (a, b)
|
||||
|
||||
final_nodes = []
|
||||
num_successes_list = []
|
||||
total_elements_evaluated_list = []
|
||||
|
||||
for allowed_steps in range(max_allowed_number_of_steps):
|
||||
samples_list = []
|
||||
for node in nodes_sampled:
|
||||
if total_elements_evaluated[node] == num_samples_in_dataset:
|
||||
print('node', repr(node), 'has been fully evaluated.', num_samples_in_dataset)
|
||||
samples_list.append(1e9 if objective == 'lowest' else -1e9)
|
||||
elif use_ucb_rule:
|
||||
success_ratio = num_successes[node] / total_elements_evaluated[node] if total_elements_evaluated[node] else 0
|
||||
|
||||
# adding one because time is one-indexed
|
||||
time_var = allowed_steps # time step, used to be np.sum(total_elements_evaluated[node])
|
||||
sqrt_term = 2 * np.sqrt(np.log(1 + time_var) / total_elements_evaluated[node]) if \
|
||||
total_elements_evaluated[node] else 0
|
||||
samples_list.append(success_ratio + sqrt_term)
|
||||
else:
|
||||
a = initial_a_b_params[0] + num_successes[node]
|
||||
b = initial_a_b_params[1] + total_elements_evaluated[node] - num_successes[node]
|
||||
samples_list.append(np.random.beta(a, b))
|
||||
if objective == 'lowest' and min(samples_list) == 1e9:
|
||||
print('Evaluated all available samples, ending. thompson_sampling')
|
||||
break
|
||||
if objective == 'highest' and max(samples_list) == -1e9:
|
||||
print('Evaluated all available samples, ending. thompson_sampling')
|
||||
break
|
||||
|
||||
chosen_node_id = np.argmin(samples_list) if objective == 'lowest' else np.argmax(samples_list)
|
||||
chosen_node = nodes_sampled[chosen_node_id]
|
||||
print(f'***************** Calling model ***************** (step={allowed_steps}, objective={objective})')
|
||||
total_elements_evaluated[chosen_node], num_successes[chosen_node] = self._evaluate_node_on_batch(
|
||||
chosen_node, batch_size)
|
||||
print('total_elements_evaluated[chosen_node]', repr(chosen_node), total_elements_evaluated[chosen_node])
|
||||
final_nodes.append(
|
||||
self._choose_final_node(num_successes, total_elements_evaluated, objective, nodes_sampled))
|
||||
num_successes_list.append(copy.deepcopy(num_successes))
|
||||
total_elements_evaluated_list.append(copy.deepcopy(total_elements_evaluated))
|
||||
|
||||
return final_nodes, num_successes_list, total_elements_evaluated_list
|
||||
|
||||
def main(self, value_assignments, batch_size, num_formats=-1, max_allowed_number_of_model_calls=100):
|
||||
max_allowed_number_of_steps = max_allowed_number_of_model_calls // batch_size
|
||||
assert max_allowed_number_of_model_calls % batch_size == 0
|
||||
assert max_allowed_number_of_steps % 2 == 0
|
||||
|
||||
# Initialize node structures
|
||||
print('Initializing node structures...')
|
||||
value_assignments_ids = value_assignment_str_to_indices(value_assignments, self.pointer_action_pairs)
|
||||
for value_assignment in value_assignments_ids:
|
||||
self._add_node_to_structures(value_assignment)
|
||||
if num_formats > 0 and len(self.generation_order) == num_formats + 1:
|
||||
break
|
||||
|
||||
nodes_sampled = list(self.all_structured_prompt_formats_accuracies.keys())
|
||||
# this is already evaluated during initialization
|
||||
original_node = [new_node for new_node, order in self.generation_order.items() if order == 0][0]
|
||||
|
||||
# Thompson Sampling
|
||||
budget_per_call = max_allowed_number_of_steps // 2
|
||||
print('***************** BEGINNING PHASE 1, budget:', budget_per_call)
|
||||
final_nodes, num_successes_list, total_elements_evaluated_list = self._evaluate_nodes_thompson_sampling(
|
||||
original_node,
|
||||
nodes_sampled,
|
||||
batch_size=batch_size,
|
||||
max_allowed_number_of_steps=budget_per_call,
|
||||
objective='highest',
|
||||
use_ucb_rule=False,
|
||||
num_successes=None,
|
||||
total_elements_evaluated=None)
|
||||
|
||||
self.metadata['thompson_sampling'] = {}
|
||||
self.metadata['thompson_sampling']['highest-num_successes_list'] = num_successes_list
|
||||
self.metadata['thompson_sampling']['highest-total_elements_evaluated_list'] = total_elements_evaluated_list
|
||||
self.metadata['thompson_sampling']['highest-final_nodes'] = final_nodes
|
||||
|
||||
best_node = final_nodes[-1]
|
||||
|
||||
print('***************** BEGINNING PHASE 2, budget:', budget_per_call)
|
||||
final_node_previous_to_phase_two = self._choose_final_node(
|
||||
num_successes_list[-1], total_elements_evaluated_list[-1], 'lowest', nodes_sampled)
|
||||
final_nodes, num_successes_list, total_elements_evaluated_list = self._evaluate_nodes_thompson_sampling(
|
||||
original_node,
|
||||
nodes_sampled,
|
||||
batch_size=batch_size,
|
||||
max_allowed_number_of_steps=budget_per_call,
|
||||
objective='lowest',
|
||||
use_ucb_rule=False,
|
||||
num_successes=copy.copy(num_successes_list[-1]),
|
||||
total_elements_evaluated=copy.copy(total_elements_evaluated_list[-1]))
|
||||
|
||||
worst_node = final_nodes[-1] if final_nodes else final_node_previous_to_phase_two
|
||||
|
||||
self.metadata['thompson_sampling']['lowest-num_successes_list'] = num_successes_list
|
||||
self.metadata['thompson_sampling']['lowest-total_elements_evaluated_list'] = total_elements_evaluated_list
|
||||
self.metadata['thompson_sampling']['lowest-final_nodes'] = final_nodes if final_nodes else worst_node
|
||||
|
||||
# these evals don't count towards the exploration budget, it's just to report final spreads found accurately
|
||||
self._evaluate_node_on_batch(best_node, num_samples=-1)
|
||||
self._evaluate_node_on_batch(worst_node, num_samples=-1)
|
||||
|
||||
print('Best Node:', repr(best_node), self.all_structured_prompt_formats_accuracies[best_node])
|
||||
print('Worst Node:', repr(worst_node), self.all_structured_prompt_formats_accuracies[worst_node])
|
||||
@@ -0,0 +1,636 @@
|
||||
import random
|
||||
import inspect
|
||||
|
||||
random.seed(42)
|
||||
|
||||
|
||||
# removed '\n\n' to make sure this is only used between entries
|
||||
CHOSEN_SEPARATOR_LIST = ['', '::: ', ':: ', ': ', ' \n\t', '\n ', ' : ', ' - ', ' ', '\n ', '\n\t', ':', '::', '- ', '\t'] # sep='' is used rarely, only for enumerations because there is already formatting there
|
||||
CHOSEN_SPACE_LIST = ['', ' ', '\n', ' \n', ' -- ', ' ', '; \n', ' || ', ' <sep> ', ' -- ', ', ', ' \n ', ' , ', '\n ', '. ', ' , '] # space='' is used a lot
|
||||
CHOSEN_SEPARATOR_TEXT_AND_OPTION_LIST = ['', ' ', ' ', '\t']
|
||||
|
||||
CHOSEN_SEPARATOR_LIST = [(e, e) for e in CHOSEN_SEPARATOR_LIST]
|
||||
CHOSEN_SPACE_LIST = [(e, e) for e in CHOSEN_SPACE_LIST]
|
||||
CHOSEN_SEPARATOR_TEXT_AND_OPTION_LIST = [(e, e) for e in CHOSEN_SEPARATOR_TEXT_AND_OPTION_LIST]
|
||||
|
||||
|
||||
TEXT_DESCRIPTOR_FN_LIST = [
|
||||
(lambda x: x, "lambda x: x"),
|
||||
(lambda x: x.title(), "lambda x: x.title()"),
|
||||
(lambda x: x.upper(), "lambda x: x.upper()"),
|
||||
(lambda x: x.lower(), "lambda x: x.lower()")
|
||||
]
|
||||
ITEM_WRAPPER_LIST = [
|
||||
(lambda x: f'({x})', "lambda x: f'({x})'"),
|
||||
(lambda x: f'{x}.', "lambda x: f'{x}.'"),
|
||||
(lambda x: f'{x})', "lambda x: f'{x})'"),
|
||||
(lambda x: f'{x} )', "lambda x: f'{x} )'"),
|
||||
(lambda x: f'[{x}]', "lambda x: f'[{x}]'"),
|
||||
(lambda x: f'<{x}>', "lambda x: f'<{x}>'"),
|
||||
]
|
||||
NUMBER_FORMAT_LIST = [
|
||||
(lambda x: x + 1, "lambda x: x + 1"),
|
||||
(lambda x: chr(ord('A') + x), "lambda x: chr(ord('A') + x)"),
|
||||
(lambda x: chr(ord('a') + x), "lambda x: chr(ord('a') + x)"),
|
||||
(lambda x: chr(0x215F + x + 1) + ('' if x < 12 else 0 / 0), "lambda x: chr(0x215F + x + 1)"),
|
||||
(lambda x: NewEnumerationPromptFormat.ROMAN_NUMERALS[x], "lambda x: EnumerationPromptFormat.ROMAN_NUMERALS[x]"),
|
||||
(lambda x: NewEnumerationPromptFormat.ROMAN_NUMERALS[x].upper(), "lambda x: EnumerationPromptFormat.ROMAN_NUMERALS[x].upper()")
|
||||
]
|
||||
|
||||
MAPPING_ALL_CATEGORIES = {
|
||||
'text_descriptor_fn': TEXT_DESCRIPTOR_FN_LIST,
|
||||
'chosen_item_wrapper': ITEM_WRAPPER_LIST,
|
||||
'chosen_number_format': NUMBER_FORMAT_LIST,
|
||||
'chosen_space': CHOSEN_SPACE_LIST,
|
||||
'chosen_separator': CHOSEN_SEPARATOR_LIST, # in OPTION_1:^TEXT, this is ^
|
||||
'chosen_separator_text_and_option': CHOSEN_SEPARATOR_TEXT_AND_OPTION_LIST # in OPTION_1:^TEXT, this is _
|
||||
}
|
||||
|
||||
|
||||
def lambda_to_string(lambda_fn):
|
||||
funcString = str(inspect.getsourcelines(lambda_fn)[0])
|
||||
funcString = funcString.strip("['\\n']").strip('\\n"').split("=")[1].strip().strip(',').strip('\n')
|
||||
return funcString
|
||||
|
||||
class SpacingBetweenPromptComponents:
|
||||
SEARCH_SPACE_VALID_OPTIONS = {
|
||||
'chosen_space': CHOSEN_SPACE_LIST
|
||||
}
|
||||
SYNONYM_SETS = []
|
||||
|
||||
def __init__(self, prompt_format_list, chosen_space, allow_only_non_char_spaces=False):
|
||||
self.chosen_space = chosen_space
|
||||
self.prompt_format = prompt_format_list # or SharedPropertyAmongPrompts
|
||||
|
||||
self.is_output_field = False
|
||||
|
||||
# in some cases we want to avoid having a comma like a space (only used right now for original chosen_space='')
|
||||
self.allow_only_non_char_spaces = allow_only_non_char_spaces
|
||||
|
||||
def solve(self, extra_params=None):
|
||||
prompt_format_with_resolved_shared_property = self.prompt_format
|
||||
if isinstance(self.prompt_format, SharedPropertyAmongPrompts):
|
||||
prompt_format_with_resolved_shared_property = self.prompt_format.solve(extra_params)
|
||||
|
||||
result = []
|
||||
for i, e in enumerate(prompt_format_with_resolved_shared_property):
|
||||
# ignore an output field if that was the request
|
||||
if not isinstance(e, str) and e.is_output_field:
|
||||
if extra_params and extra_params.get('print_output_fields', False):
|
||||
if i > 0:
|
||||
result.append(self.chosen_space)
|
||||
result.append(e.solve(extra_params))
|
||||
else:
|
||||
if i > 0:
|
||||
result.append(self.chosen_space)
|
||||
result.append(e.solve(extra_params) if not isinstance(e, str) else e)
|
||||
|
||||
return result
|
||||
|
||||
def find_all_formatted_field_values(self):
|
||||
if isinstance(self.prompt_format, SharedPropertyAmongPrompts):
|
||||
return self.prompt_format.find_all_formatted_field_values()
|
||||
else:
|
||||
result = {}
|
||||
for e in self.prompt_format:
|
||||
assert len(set(result.keys()) & set(e.find_all_formatted_field_values().keys())) == 0
|
||||
result.update(e.find_all_formatted_field_values())
|
||||
return result
|
||||
|
||||
def update_field(self, field_name, new_field_value):
|
||||
if field_name not in self.__dict__:
|
||||
return False
|
||||
|
||||
if self.allow_only_non_char_spaces and not new_field_value.isspace():
|
||||
return False
|
||||
|
||||
setattr(self, field_name, new_field_value)
|
||||
return True
|
||||
|
||||
def has_attribute(self, field_name):
|
||||
return field_name in self.__dict__
|
||||
|
||||
def attributes_under_control(self):
|
||||
return list(self.SEARCH_SPACE_VALID_OPTIONS.keys())
|
||||
|
||||
|
||||
class NewEnumerationPromptFormat:
|
||||
"""
|
||||
Variable-length enumeration. E.g. listing facts, listing options.
|
||||
|
||||
This new version is less recursive.
|
||||
|
||||
Option 1 : text <sep> Option 2 : text
|
||||
"""
|
||||
|
||||
ROMAN_NUMERALS = ['i', 'ii', 'iii', 'iv', 'v', 'vi', 'vii', 'viii', 'ix', 'x', 'xi', 'xii', 'xiii', 'xiv', 'xv']
|
||||
SEARCH_SPACE_VALID_OPTIONS = {
|
||||
'text_descriptor_fn': TEXT_DESCRIPTOR_FN_LIST,
|
||||
'chosen_item_wrapper': ITEM_WRAPPER_LIST,
|
||||
'chosen_number_format': NUMBER_FORMAT_LIST,
|
||||
'chosen_space': CHOSEN_SPACE_LIST,
|
||||
'chosen_separator': CHOSEN_SEPARATOR_LIST, # in OPTION_1:^TEXT, this is ^
|
||||
'chosen_separator_text_and_option': CHOSEN_SEPARATOR_TEXT_AND_OPTION_LIST # in OPTION_1:^TEXT, this is _
|
||||
}
|
||||
|
||||
SYNONYM_SETS = []
|
||||
def __init__(self,
|
||||
text_descriptor_format,
|
||||
length,
|
||||
chosen_space,
|
||||
chosen_separator=': ',
|
||||
chosen_separator_owner=None,
|
||||
chosen_separator_text_and_option=None,
|
||||
chosen_item_wrapper=None,
|
||||
chosen_number_format=None,
|
||||
text_descriptor_fn=lambda x: x,
|
||||
text_descriptor_fn_owner=None,
|
||||
object_name=None):
|
||||
self.chosen_item_wrapper = \
|
||||
chosen_item_wrapper if chosen_item_wrapper else self.SEARCH_SPACE_VALID_OPTIONS['chosen_item_wrapper'][0][0]
|
||||
self.chosen_number_format = \
|
||||
chosen_number_format if chosen_number_format else self.SEARCH_SPACE_VALID_OPTIONS['chosen_number_format'][0][0]
|
||||
self.chosen_space = chosen_space
|
||||
self.chosen_separator = chosen_separator
|
||||
self.chosen_separator_owner = chosen_separator_owner
|
||||
|
||||
if chosen_separator_text_and_option is None:
|
||||
chosen_separator_text_and_option = '' if not text_descriptor_format else ' '
|
||||
self.chosen_separator_text_and_option = chosen_separator_text_and_option
|
||||
|
||||
self.chosen_space_between_text_and_item = None
|
||||
|
||||
self.text_descriptor_format = text_descriptor_format
|
||||
self.text_descriptor_fn_owner = text_descriptor_fn_owner
|
||||
self.text_descriptor_fn = text_descriptor_fn
|
||||
|
||||
assert isinstance(length, int) or isinstance(length, list)
|
||||
length_range = range(length) if isinstance(length, int) else length
|
||||
|
||||
self.enumeration_item_id_list = length_range
|
||||
self.prompt_format = text_descriptor_format # FIXME? this is just so that it's a str for when calling pointers_to_all_objects()
|
||||
|
||||
self.is_output_field = False
|
||||
self.object_name = object_name # used to reference this object when filling
|
||||
|
||||
def format_text_descriptor_field(self, index):
|
||||
text = '<|text|>'
|
||||
|
||||
if self.text_descriptor_fn_owner is None:
|
||||
prompt = self.text_descriptor_fn(self.text_descriptor_format)
|
||||
else:
|
||||
prompt = self.text_descriptor_fn_owner.apply_field_fn('text_descriptor_fn', self.text_descriptor_format)
|
||||
|
||||
chosen_separator = self.chosen_separator
|
||||
if self.chosen_separator_owner is not None:
|
||||
assert 'chosen_separator' in self.chosen_separator_owner.fields
|
||||
chosen_separator = self.chosen_separator_owner.fields['chosen_separator']
|
||||
|
||||
# return prompt.format(self.chosen_item_wrapper(self.chosen_number_format(index)))
|
||||
return f"{prompt}{self.chosen_separator_text_and_option}{self.chosen_item_wrapper(self.chosen_number_format(index))}{chosen_separator}{text}"
|
||||
|
||||
def solve(self, extra_params=None):
|
||||
"""
|
||||
extra_params: Dictates whether to modify the self.prompt_format.
|
||||
Currently used only to print fewer options in the enumeration than the maximum allowed.
|
||||
"""
|
||||
enumeration_length = extra_params.get('enumeration_length', None) if extra_params else None
|
||||
|
||||
# First, solve each enumeration item
|
||||
solved_elements = []
|
||||
for index in self.enumeration_item_id_list[:enumeration_length]:
|
||||
solved_elements.append(self.format_text_descriptor_field(index))
|
||||
|
||||
result = []
|
||||
for i, e in enumerate(solved_elements):
|
||||
if i > 0:
|
||||
result.append(self.chosen_space)
|
||||
result.append(e.solve(extra_params) if not isinstance(e, str) else e)
|
||||
|
||||
return result
|
||||
|
||||
def find_all_formatted_field_values(self):
|
||||
"""
|
||||
Obtain a dictionary with all the (field_name, field_value) to be used
|
||||
in updating the instruction formatted field values.
|
||||
"""
|
||||
|
||||
if self.object_name:
|
||||
field_names_to_values = {
|
||||
f'{self.object_name}_{i + 1}': self.chosen_number_format(index)
|
||||
for i, index in enumerate(self.enumeration_item_id_list)
|
||||
}
|
||||
return field_names_to_values
|
||||
return {}
|
||||
|
||||
def update_field(self, field_name, new_field_value):
|
||||
if field_name not in self.__dict__:
|
||||
return False
|
||||
|
||||
"""
|
||||
Check for consistency between components, to avoid weird looking enumerations like the following:
|
||||
|
||||
Options:
|
||||
1.
|
||||
{} 2.
|
||||
{} 3.
|
||||
{} 4.
|
||||
{}
|
||||
|
||||
Rule to enforce is: '\n' in chosen_separator (e.g. "::" in "1::") => '\n' in chosen_space
|
||||
"""
|
||||
spacing_values = {
|
||||
'chosen_separator': self.chosen_separator,
|
||||
'chosen_separator_text_and_option': self.chosen_separator_text_and_option,
|
||||
'chosen_space': self.chosen_space
|
||||
}
|
||||
spacing_values[field_name] = new_field_value
|
||||
|
||||
if self.chosen_separator_owner is not None:
|
||||
assert 'chosen_separator' in self.chosen_separator_owner.fields
|
||||
spacing_values['chosen_separator'] = self.chosen_separator_owner.fields['chosen_separator']
|
||||
|
||||
if ('\n' in spacing_values['chosen_separator'] or
|
||||
'\n' in spacing_values['chosen_separator_text_and_option']) and \
|
||||
'\n' not in spacing_values['chosen_space']:
|
||||
return False
|
||||
|
||||
setattr(self, field_name, new_field_value)
|
||||
return True
|
||||
|
||||
def has_attribute(self, field_name):
|
||||
return field_name in self.__dict__
|
||||
|
||||
def attributes_under_control(self):
|
||||
attrs = list(self.SEARCH_SPACE_VALID_OPTIONS.keys())
|
||||
if not self.text_descriptor_format:
|
||||
attrs.remove('text_descriptor_fn') # changing casing and space from an empty string doesn't make sense
|
||||
attrs.remove('chosen_separator_text_and_option')
|
||||
if self.text_descriptor_fn_owner is not None:
|
||||
attrs.remove('text_descriptor_fn') # this attribute is controlled by some other entity
|
||||
if self.chosen_separator_owner is not None:
|
||||
attrs.remove('chosen_separator')
|
||||
return attrs
|
||||
|
||||
|
||||
class SimplePromptFormat:
|
||||
"""
|
||||
Simplest formatting. For example,
|
||||
|
||||
Sentence: <|text|>
|
||||
Question: <|text|>
|
||||
Answer: <|text|>
|
||||
"""
|
||||
|
||||
SEARCH_SPACE_VALID_OPTIONS = {
|
||||
'chosen_separator': CHOSEN_SEPARATOR_LIST,
|
||||
'text_descriptor_fn': TEXT_DESCRIPTOR_FN_LIST
|
||||
}
|
||||
SYNONYM_SETS = []
|
||||
|
||||
def __init__(self,
|
||||
text_descriptor,
|
||||
separator,
|
||||
text_descriptor_fn=lambda x: x,
|
||||
prompt_without_text=False,
|
||||
chosen_separator_owner=None,
|
||||
text_descriptor_fn_owner=None,
|
||||
is_output_field=False):
|
||||
self.text_descriptor = text_descriptor # keep as is
|
||||
self.chosen_separator = separator
|
||||
self.prompt_format = self.text_descriptor
|
||||
|
||||
self.prompt_without_text = prompt_without_text # used for text only prompts (without variable text)
|
||||
self.text_descriptor_fn = text_descriptor_fn
|
||||
# self.index_item = -1 # only used for enumerations
|
||||
|
||||
self.text_descriptor_owner = None
|
||||
self.chosen_separator_owner = chosen_separator_owner
|
||||
self.text_descriptor_fn_owner = text_descriptor_fn_owner
|
||||
|
||||
self.is_output_field = is_output_field
|
||||
|
||||
def assign_field_owner(self, field_name, owner):
|
||||
assert field_name in self.__dict__
|
||||
setattr(self, field_name + '_owner', owner)
|
||||
|
||||
def solve(self, extra_params=None):
|
||||
|
||||
resolved_prompt_format = self.prompt_format
|
||||
if self.text_descriptor_owner: # only used for enumeration
|
||||
assert self.index_item is not None
|
||||
resolved_prompt_format = self.text_descriptor_owner.format_text_descriptor_field(self.index_item)
|
||||
elif self.text_descriptor_fn_owner:
|
||||
resolved_prompt_format = self.text_descriptor_fn_owner.apply_field_fn('text_descriptor_fn', resolved_prompt_format)
|
||||
else:
|
||||
resolved_prompt_format = self.text_descriptor_fn(resolved_prompt_format)
|
||||
|
||||
true_separator = self.chosen_separator
|
||||
if self.chosen_separator_owner:
|
||||
assert 'chosen_separator' in self.chosen_separator_owner.fields
|
||||
true_separator = self.chosen_separator_owner.fields['chosen_separator']
|
||||
|
||||
exclude_text_field_for_output_fields = self.is_output_field and extra_params and extra_params.get('exclude_text_field_for_output_fields', False)
|
||||
|
||||
text = '' if self.prompt_without_text or exclude_text_field_for_output_fields else '<|text|>'
|
||||
return f"{resolved_prompt_format}{true_separator}{text}"
|
||||
|
||||
def find_all_formatted_field_values(self):
|
||||
return {}
|
||||
|
||||
def update_field(self, field_name, new_field_value):
|
||||
if field_name not in self.__dict__:
|
||||
return False
|
||||
|
||||
if self.chosen_separator_owner and field_name in self.chosen_separator_owner.fields:
|
||||
return False
|
||||
|
||||
# we need a separator on simple prompt format, otherwise it'd be "INPUT<text>" which is illegible
|
||||
if field_name == 'chosen_separator' and new_field_value == '':
|
||||
return False
|
||||
|
||||
setattr(self, field_name, new_field_value)
|
||||
return True
|
||||
|
||||
def has_attribute(self, field_name):
|
||||
return field_name in self.__dict__
|
||||
|
||||
def attributes_under_control(self):
|
||||
result = []
|
||||
if self.chosen_separator_owner is None:
|
||||
result.append('chosen_separator')
|
||||
if self.text_descriptor_fn_owner is None and self.text_descriptor:
|
||||
result.append('text_descriptor_fn')
|
||||
return result
|
||||
|
||||
|
||||
class NoTextPromptFormat:
|
||||
SEARCH_SPACE_VALID_OPTIONS = {}
|
||||
|
||||
def __init__(self):
|
||||
self.is_output_field = False
|
||||
self.prompt_format = ''
|
||||
|
||||
def solve(self, extra_params=None):
|
||||
exclude_text_field_for_output_fields = self.is_output_field and extra_params and extra_params.get('exclude_text_field_for_output_fields', False)
|
||||
|
||||
text = '' if exclude_text_field_for_output_fields else '<|text|>'
|
||||
return text
|
||||
|
||||
def attributes_under_control(self):
|
||||
return []
|
||||
|
||||
def find_all_formatted_field_values(self):
|
||||
return {}
|
||||
|
||||
|
||||
class SharedPropertyAmongPrompts:
|
||||
SEARCH_SPACE_VALID_OPTIONS = {
|
||||
'chosen_separator': CHOSEN_SEPARATOR_LIST,
|
||||
'text_descriptor_fn': TEXT_DESCRIPTOR_FN_LIST
|
||||
}
|
||||
SYNONYM_SETS = []
|
||||
|
||||
def __init__(self, fields_dict, prompt_list_to_apply):
|
||||
self.fields = fields_dict # = {'chosen_separator': ':: '}
|
||||
self.prompt_format = prompt_list_to_apply
|
||||
|
||||
if prompt_list_to_apply is not None:
|
||||
for field_name, field_value in self.fields.items():
|
||||
for e in self.prompt_format:
|
||||
e.assign_field_owner(field_name, self)
|
||||
assert field_name in e.__dict__
|
||||
setattr(e, field_name, field_value)
|
||||
|
||||
self.is_output_field = False
|
||||
|
||||
def solve(self, extra_params=None):
|
||||
if self.prompt_format is None:
|
||||
return None
|
||||
|
||||
enumeration_length = extra_params.get('enumeration_length') if extra_params else None # a[:None] returns full list
|
||||
|
||||
result = []
|
||||
for e in self.prompt_format[:enumeration_length]:
|
||||
if not isinstance(e, str) and e.is_output_field:
|
||||
if extra_params and extra_params.get('print_output_fields', False):
|
||||
result.append(e.solve(extra_params))
|
||||
else:
|
||||
result.append(e.solve(extra_params))
|
||||
|
||||
def find_all_formatted_field_values(self):
|
||||
if self.prompt_format is None:
|
||||
return {}
|
||||
|
||||
result = {}
|
||||
for e in self.prompt_format:
|
||||
assert len(set(result.keys()) & set(e.find_all_formatted_field_values().keys())) == 0
|
||||
result.update(e.find_all_formatted_field_values())
|
||||
return result
|
||||
|
||||
def update_field(self, field_name, new_field_value):
|
||||
if field_name not in self.fields:
|
||||
return False
|
||||
|
||||
self.fields[field_name] = new_field_value
|
||||
return True
|
||||
|
||||
def has_attribute(self, field_name):
|
||||
return field_name in self.fields
|
||||
|
||||
def apply_field_fn(self, field_name, string):
|
||||
assert field_name in self.fields
|
||||
return self.fields[field_name](string)
|
||||
|
||||
def attributes_under_control(self):
|
||||
return list(self.fields.keys())
|
||||
|
||||
|
||||
def flatten(nested_string_list):
|
||||
return "".join([flatten(e) if isinstance(e, list) else e for e in nested_string_list])
|
||||
|
||||
|
||||
def pointers_to_all_objects(root_element):
|
||||
result = [root_element]
|
||||
if not isinstance(root_element.prompt_format, list):
|
||||
return result + pointers_to_all_objects(root_element.prompt_format)
|
||||
|
||||
for elem in root_element.prompt_format:
|
||||
result.append(elem)
|
||||
if not isinstance(elem.prompt_format, str):
|
||||
result.extend(pointers_to_all_objects(elem))
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def get_possible_actions(e, allow_text_action_type=True):
|
||||
possible_keys = [k for k in e.SEARCH_SPACE_VALID_OPTIONS if e.has_attribute(k)] # is punctuation replacement an option?
|
||||
assert all([k in possible_keys for k in e.attributes_under_control()]), f'{e.attributes_under_control()} not subset of {possible_keys} for node {e.solve()}'
|
||||
|
||||
possible_keys = e.attributes_under_control() # this should avoid self loops in graph search
|
||||
|
||||
if allow_text_action_type and any(v in e.prompt_format for v_list in e.SYNONYM_SETS for v in v_list): # is text replacement an option?
|
||||
possible_keys += ['text']
|
||||
|
||||
return possible_keys
|
||||
|
||||
|
||||
def create_pointer_action_type_pairs(
|
||||
all_pointers_enumerated, forced_action_type=None, allow_text_action_type=True
|
||||
):
|
||||
"""
|
||||
Simultaneously choose which element we'll perform the action over, and the action itself.
|
||||
"""
|
||||
|
||||
pointer_action_pairs = []
|
||||
for e, index in all_pointers_enumerated:
|
||||
possible_keys = get_possible_actions(e, allow_text_action_type)
|
||||
if forced_action_type:
|
||||
possible_keys = [forced_action_type] if forced_action_type in possible_keys else []
|
||||
for action_type in possible_keys:
|
||||
pointer_action_pairs.append((e, index, action_type))
|
||||
|
||||
return pointer_action_pairs
|
||||
|
||||
|
||||
def holistic_node_format_sanity_checks(root_element, prohibit_newlines=False):
|
||||
"""
|
||||
Checks that the prompt format's value assignments are reasonable, and consistent across fields.
|
||||
|
||||
For example, this functions checks that if a space between component does not have \n, then the separator between
|
||||
fields should also not have that.
|
||||
|
||||
E.g. input\n{}output\n{} returns False.
|
||||
E.g. input\n{}\noutput\n{} returns True.
|
||||
E.g. input {}\noutput {} returns True.
|
||||
E.g. this should return True (because Options is prompt_without_text=True):
|
||||
Question
|
||||
<|text|>
|
||||
Options
|
||||
[1] <|text|> [2] <|text|> [3] <|text|> [4] <|text|> [5] <|text|>
|
||||
Answer
|
||||
<|text|>
|
||||
|
||||
Also checks the update_field() rule of spacing in NewEnumerationPromptFormat.
|
||||
|
||||
"""
|
||||
if isinstance(root_element, str):
|
||||
return True
|
||||
|
||||
# local constraint from NewEnumeration, added here because update_field() won't be called from genetic/global_random
|
||||
if isinstance(root_element, NewEnumerationPromptFormat):
|
||||
spacing_values = {
|
||||
'chosen_separator': root_element.chosen_separator,
|
||||
'chosen_separator_text_and_option': root_element.chosen_separator_text_and_option,
|
||||
'chosen_space': root_element.chosen_space
|
||||
}
|
||||
if root_element.chosen_separator_owner is not None:
|
||||
assert 'chosen_separator' in root_element.chosen_separator_owner.fields
|
||||
spacing_values['chosen_separator'] = root_element.chosen_separator_owner.fields['chosen_separator']
|
||||
|
||||
if ('\n' in spacing_values['chosen_separator'] or
|
||||
'\n' in spacing_values['chosen_separator_text_and_option']) and \
|
||||
'\n' not in spacing_values['chosen_space']:
|
||||
return False
|
||||
|
||||
# local constraint from simple prompt format: we need an actual separator in simple formats, '' is invalid
|
||||
if isinstance(root_element, SimplePromptFormat):
|
||||
true_separator = root_element.chosen_separator
|
||||
if root_element.chosen_separator_owner:
|
||||
assert 'chosen_separator' in root_element.chosen_separator_owner.fields
|
||||
true_separator = root_element.chosen_separator_owner.fields['chosen_separator']
|
||||
|
||||
if true_separator == '':
|
||||
return False
|
||||
|
||||
# local constraint from SpacingBetweenPromptComponents
|
||||
if isinstance(root_element, SpacingBetweenPromptComponents) and \
|
||||
root_element.allow_only_non_char_spaces and not root_element.chosen_space.isspace():
|
||||
return False
|
||||
|
||||
# global constraint: avoid using chosen_space='' unless it is separating between a prompt without text and a text.
|
||||
# E.g. INPUT - <|text|>OUTPUT - <|text|> should not be allowed but
|
||||
# OPTIONS: A. text B. text should be accepted
|
||||
if isinstance(root_element, SpacingBetweenPromptComponents) and root_element.chosen_space == '' and \
|
||||
isinstance(root_element.prompt_format, list):
|
||||
|
||||
all_prompt_without_texts_except_maybe_last_elem = all(
|
||||
hasattr(elem, 'prompt_without_text') and elem.prompt_without_text
|
||||
for elem in root_element.prompt_format[:-1])
|
||||
if not all_prompt_without_texts_except_maybe_last_elem:
|
||||
return False
|
||||
|
||||
# global constraint with newlines as explained in the function's documentation
|
||||
if isinstance(root_element, SpacingBetweenPromptComponents) and '\n' not in root_element.chosen_space:
|
||||
if isinstance(root_element.prompt_format, list):
|
||||
return all(holistic_node_format_sanity_checks(e, prohibit_newlines=True) for e in root_element.prompt_format)
|
||||
else:
|
||||
return holistic_node_format_sanity_checks(root_element.prompt_format, prohibit_newlines=True)
|
||||
|
||||
# FIXME add the exception of an empty text field
|
||||
if prohibit_newlines and hasattr(root_element, 'chosen_separator'):
|
||||
# if the prompt does not have text then it is ok to put a new line, since it's not awkwardly separating
|
||||
# the descriptor from the text, which is our goal here
|
||||
if hasattr(root_element, 'prompt_without_text') and root_element.prompt_without_text:
|
||||
pass
|
||||
else:
|
||||
chosen_separator = root_element.chosen_separator
|
||||
if root_element.chosen_separator_owner is not None:
|
||||
assert 'chosen_separator' in root_element.chosen_separator_owner.fields
|
||||
chosen_separator = root_element.chosen_separator_owner.fields['chosen_separator']
|
||||
if '\n' in chosen_separator:
|
||||
return False
|
||||
|
||||
if not isinstance(root_element.prompt_format, list):
|
||||
return holistic_node_format_sanity_checks(root_element.prompt_format, prohibit_newlines=prohibit_newlines)
|
||||
|
||||
result = [holistic_node_format_sanity_checks(elem, prohibit_newlines=prohibit_newlines)
|
||||
for elem in root_element.prompt_format]
|
||||
return all(result)
|
||||
|
||||
|
||||
def apply_prompt_format(prompt, input_fields):
|
||||
# Possible FIX for variable-length prompt formats. Choose output based on the number of fields.
|
||||
tmp = prompt.format(*input_fields)
|
||||
if prompt.count('{}') != len(input_fields):
|
||||
print('WARNING, wrong number of fields!', prompt, input_fields)
|
||||
return tmp
|
||||
|
||||
|
||||
def _one_text_field(text1, answer_field_text='Answer', chosen_space='\n'):
|
||||
# Input: <text>\nOutput: <text>
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat(text1, None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat(answer_field_text, None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=chosen_space
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
return structured_prompt_format, global_constraints
|
||||
|
||||
|
||||
def _two_text_fields(text1, text2, answer_field_text='Answer', chosen_space='\n'):
|
||||
# Passage: <text>\nQuestion: <text>
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat(text1, None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat(text2, None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat(answer_field_text, None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=chosen_space
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
return structured_prompt_format, global_constraints
|
||||
|
||||
+833
@@ -0,0 +1,833 @@
|
||||
from .grammar_definition import SpacingBetweenPromptComponents, SharedPropertyAmongPrompts, \
|
||||
NewEnumerationPromptFormat, SimplePromptFormat, NoTextPromptFormat, _one_text_field, _two_text_fields
|
||||
|
||||
SOCIAL_GOOD_TASK_IDS = [
|
||||
'task137_', # 2-Choice output, prompt formatted -- 379 samples
|
||||
'task327_', 'task333_', 'task335_', 'task337_',
|
||||
# prompt formatted, binary classification -- +2000 samples, FIXME allow emojis in regex matching
|
||||
'task905_', # prompt formatted, classification -- +2000 samples, no parsing errors
|
||||
'task320_', # prompt formatted-ish, classification
|
||||
'task1502_', 'task1503_', 'task1504_', # no prompt format: classification, classification, generation
|
||||
'task1664_', # no prompt format: set of words as output
|
||||
'task1669_', 'task1670_', # no prompt format, long generation but well defined!
|
||||
'task1720_', 'task1725_', # no prompt format, binary classification
|
||||
'task904_', # no prompt format, classification,
|
||||
'task277_', 'task278_', 'task279_', 'task280_', 'task316_', 'task317_', 'task318_', 'task319_', 'task320_',
|
||||
'task321_',
|
||||
'task108_',
|
||||
'task322_', 'task323_', 'task324_', 'task325_', 'task326_', 'task327_', 'task328_',
|
||||
'task1604_', 'task1605_', 'task1606_', 'task1607_',
|
||||
'task1721_', 'task1722_', 'task1723_', 'task1724_',
|
||||
'task607_', 'task608_', 'task609_', 'task286_'
|
||||
]
|
||||
|
||||
SUPERNATURAL_INSTRUCTIONS_TASKS_WITH_NO_FORMAT = [
|
||||
'task1502_', 'task1503_', 'task1504_', # no prompt format: classification, classification, generation
|
||||
'task1664_', # no prompt format: set of words as output
|
||||
'task1669_', 'task1670_', # no prompt format, long generation but well defined!
|
||||
'task1720_', 'task1725_', # no prompt format, binary classification
|
||||
'task904_', # no prompt format, classification
|
||||
'task108_',
|
||||
'task1604_', 'task1605_', 'task1606_', 'task1607_',
|
||||
'task1721_', 'task1722_', 'task1723_', 'task1724_',
|
||||
'task607_', 'task608_', 'task609_', 'task286_',
|
||||
'task1149_', 'task1189_'
|
||||
]
|
||||
|
||||
FORMATTED_MULTIPLE_CHOICE_SUPERNATURAL_INSTRUCTIONS_TASKS = [ # ends up being one-field format
|
||||
'task065_', 'task1297_', 'task084_', 'task697_', 'task729_',
|
||||
'task1380_', 'task1381_', 'task309_', 'task1431_', 'task220_', 'task1612_', 'task190_', 'task1347_',
|
||||
'task069_', 'task070_',
|
||||
'task137_', 'task138_', 'task139_', 'task140_', 'task296_', 'task297_', 'task118_', 'task1135_',
|
||||
'task1424_', 'task1423_', 'task1422_', 'task1421_', 'task1420_', 'task1419_',
|
||||
'task1678_', 'task385_', 'task580_', 'task214_', 'task213_'
|
||||
]
|
||||
|
||||
FORMATTED_TWO_TEXT_FIELDS_SUPERNATURAL_INSTRUCTIONS_TASKS = \
|
||||
['task1661_', 'task027_', 'task136_', 'task021_', 'task018_', 'task020_', 'task740_',
|
||||
'task1366_', 'task1162_', 'task1587_', 'task491_', 'task492_', 'task050_', 'task1387_',
|
||||
'task1186_', 'task1283_', 'task1284_', 'task905_', 'task501_']
|
||||
|
||||
FORMATTED_ONE_TEXT_FIELDS_SUPERNATURAL_INSTRUCTIONS_TASKS = [
|
||||
'task155_', 'task158_', 'task161_', 'task163_', 'task162_', 'task322_', 'task323_',
|
||||
'task324_', 'task325_', 'task326_', 'task327_', 'task328_', 'task333_', 'task335_',
|
||||
'task337_', 'task277_', 'task278_', 'task279_', 'task280_', 'task316_', 'task317_',
|
||||
'task113_', 'task114_']
|
||||
|
||||
FORMATTED_SOME_TEXT_FIELDS_SUPERNATURAL_INSTRUCTIONS_TASKS = [
|
||||
'task318_', 'task319_', 'task320_', 'task321_', 'task133_']
|
||||
|
||||
OPEN_GENERATION_SUPERNATURAL_INSTRUCTIONS_TASKS = [
|
||||
'task240_', 'task845_', 'task348_', 'task389_', 'task443_', 'task223_',
|
||||
'task105_', 'task1401_', 'task040_', 'task067_', 'task071_', 'task072_',
|
||||
'task1326_', 'task037_', 'task038_', 'task1613_', 'task216_']
|
||||
|
||||
|
||||
def create_initial_structured_prompt_format(args):
|
||||
structured_prompt_format = None
|
||||
global_constraints = []
|
||||
extra_params_structured_prompt_format = None
|
||||
instruction = None
|
||||
original_multiple_choice_output_format = None
|
||||
|
||||
if any(t in args.task_filename for t in ['task1661_', 'task027_']):
|
||||
structured_prompt_format, global_constraints = _two_text_fields('Passage', 'Question')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task136_', 'task021_', 'task018_', 'task020_', 'task740_']):
|
||||
structured_prompt_format, global_constraints = _two_text_fields('Sentence', 'Question')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task1366_']):
|
||||
structured_prompt_format, global_constraints = _two_text_fields('Paragraph', 'Claim')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task1162_']):
|
||||
structured_prompt_format, global_constraints = _two_text_fields('Paragraph', 'Title', chosen_space='\n ')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task1587_']):
|
||||
structured_prompt_format, global_constraints = _two_text_fields('Abstract', 'Title', chosen_space='. ')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task491_', 'task492_']):
|
||||
structured_prompt_format, global_constraints = _two_text_fields('Sentence', 'Question', chosen_space=' ')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task050_']):
|
||||
structured_prompt_format, global_constraints = _two_text_fields('Sentence', 'Question', chosen_space=' \n')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task1387_']):
|
||||
structured_prompt_format, global_constraints = _two_text_fields('Premise', 'Hypothesis', chosen_space=' <sep> ')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task1186_', 'task1283_', 'task1284_']):
|
||||
structured_prompt_format, global_constraints = _two_text_fields(
|
||||
'System Reference', 'Original Reference', chosen_space='\n')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task190_', 'task1347_']):
|
||||
# note: output is not one of the enumerations!
|
||||
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
NewEnumerationPromptFormat('Sentence', 2, chosen_separator=': ', chosen_space=' ',
|
||||
chosen_item_wrapper=lambda x: f"{x}",
|
||||
text_descriptor_fn_owner=text_descriptor_fn, object_name='enum1'),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' \n '
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif any(t in args.task_filename for t in ['task1612_']):
|
||||
# note: output is not one of the enumerations!
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
NewEnumerationPromptFormat('sentence', 2, chosen_separator=': ', chosen_separator_text_and_option='_',
|
||||
chosen_space=' ', chosen_item_wrapper=lambda x: f"{x}",
|
||||
chosen_number_format=lambda x: chr(ord('A') + x),
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' \n '
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif any(t in args.task_filename for t in ['task905_']):
|
||||
structured_prompt_format, global_constraints = _two_text_fields('Tweet', 'Label', chosen_space='\n')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task155_', 'task158_', 'task161_', 'task163_', 'task162_']):
|
||||
# msclar: these are counting tasks
|
||||
structured_prompt_format, global_constraints = _one_text_field('Sentence', chosen_space='\n')
|
||||
|
||||
elif any(t in args.task_filename for t in
|
||||
['task322_', 'task323_', 'task324_', 'task325_', 'task326_', 'task327_', 'task328_']):
|
||||
# msclar: these are counting tasks
|
||||
structured_prompt_format, global_constraints = _one_text_field('Comment', chosen_space='\n')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task333_', 'task335_', 'task337_']):
|
||||
# msclar: these are counting tasks
|
||||
structured_prompt_format, global_constraints = _one_text_field('Post', chosen_space='\n')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task277_', 'task278_']):
|
||||
structured_prompt_format, global_constraints = _one_text_field('Context', chosen_space='\n')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task279_', 'task280_', 'task316_', 'task317_']):
|
||||
structured_prompt_format, global_constraints = _one_text_field('Passage', chosen_space='\n')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task113_', 'task114_']):
|
||||
structured_prompt_format, global_constraints = _one_text_field('Sentence', chosen_space='\n')
|
||||
|
||||
elif any(t in args.task_filename for t in ['task318_', 'task319_', 'task320_', 'task321_']):
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Target', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
NoTextPromptFormat(),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' \n'
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif 'task501_' in args.task_filename:
|
||||
# ((0.39, 0.37, 100), 'CLAIM : {}. POST : {}', 'CLAIM : {}. POST : {}. ANSWER : {}')
|
||||
# CLAIM : <text>. POST : <text>
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ' : '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x.upper()}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Claim', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat('Post', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space='. '
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif 'task133_' in args.task_filename:
|
||||
# Sentence: <text>\n Reason: <text>\n Question: <text>
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Sentence', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat('Reason', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat('Question', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space='\n '
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif 'task220_' in args.task_filename:
|
||||
# Sentence 1: <text> Sentence 2: <text> Sentence 3: <text> Sentence 4: <text> Sentence 5: <text> Choices: a. <text> b. <text>
|
||||
|
||||
instruction = "In this task, you're given five sentences, numbered {enum0_1} through {enum0_5}, and two options {enum1_1} and {enum1_2} for possible titles for the story. Your job is to choose the title that better fits the story. Indicate your choice by '{enum1_1}' or '{enum1_2}'."
|
||||
original_multiple_choice_output_format = '{enum1}'
|
||||
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
|
||||
# chosen_space = SharedPropertyAmongPrompts({'space': ', '}, None) # FIXME allow to jointly change these two spaces.
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
NewEnumerationPromptFormat('Sentence', 5, chosen_separator_owner=chosen_separator, chosen_space=' ',
|
||||
chosen_item_wrapper=lambda x: f"{x}",
|
||||
text_descriptor_fn_owner=text_descriptor_fn,
|
||||
object_name='enum0'),
|
||||
SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Choices', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, prompt_without_text=True),
|
||||
NewEnumerationPromptFormat('', 2, chosen_space=' ', chosen_separator=' ',
|
||||
chosen_item_wrapper=lambda x: f"{x}.",
|
||||
chosen_number_format=lambda x: chr(ord('a') + x),
|
||||
object_name='enum1'),
|
||||
],
|
||||
chosen_space='',
|
||||
allow_only_non_char_spaces=True
|
||||
),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' '
|
||||
)
|
||||
|
||||
elif 'task1431_' in args.task_filename:
|
||||
instruction = "In this task, you are given a multiple-choice question about healthcare. Answer the question based on your information and classify your answers into '{enum1_1}', '{enum1_2}', '{enum1_3}', and '{enum1_4}'."
|
||||
original_multiple_choice_output_format = '{enum1}'
|
||||
|
||||
# Question: <text>\n Options: <1> <text> <2> <text> <3> <text> <4> <text> <5> <text>
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Question', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Options', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, prompt_without_text=True),
|
||||
NewEnumerationPromptFormat('', 5, chosen_space=' ', chosen_separator=' ',
|
||||
chosen_item_wrapper=lambda x: f"<{x}>", object_name='enum1'),
|
||||
],
|
||||
chosen_space=' '
|
||||
),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space='\n '
|
||||
)
|
||||
|
||||
global_constraints = [chosen_separator, text_descriptor_fn]
|
||||
|
||||
elif 'task309_' in args.task_filename:
|
||||
# Article: <text>\n Question: <text>\n Options: (A) <text> (B) <text> (C) <text> (D) <text>
|
||||
|
||||
instruction = 'In this task, you\'re given an article, a question which often contains a blank and four options (associated with "{enum1_1}", "{enum1_2}", "{enum1_3}", "{enum1_4}"). Your task is to find the correct answer (from the given options) for the question from the given article and return one of the options from "{enum1_1}", "{enum1_2}", "{enum1_3}", and "{enum1_4}". Do not generate anything else apart from one of the following characters: "{enum1_1}", "{enum1_2}", "{enum1_3}", "{enum1_4}". There is only one correct answer for each question.'
|
||||
original_multiple_choice_output_format = '{enum1}'
|
||||
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Article', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat('Question', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Options', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, prompt_without_text=True),
|
||||
NewEnumerationPromptFormat('', 4, chosen_space=' ', chosen_separator=' ',
|
||||
chosen_item_wrapper=lambda x: f"({x})",
|
||||
chosen_number_format=lambda x: chr(ord('A') + x),
|
||||
object_name='enum1'),
|
||||
],
|
||||
chosen_space='',
|
||||
allow_only_non_char_spaces=True
|
||||
),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space='\n '
|
||||
)
|
||||
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif any(t in args.task_filename for t in ['task1380_', 'task1381_']):
|
||||
# Sentence: <text> Question: <text> (A) <text> (B) <text>
|
||||
|
||||
if 'task1380_' in args.task_filename:
|
||||
instruction = "You are given a sentence, a question and two answer options ('{enum1_1}' and '{enum1_2}'). Your task is to find the correct option for the given question. Write down the answer index: '{enum1_1}' or '{enum1_2}'."
|
||||
elif 'task1381_' in args.task_filename:
|
||||
instruction = "You are given a sentence, a question and two answer options. Your task is to write down the index ('{enum1_1}' or '{enum1_2}') of the **incorrect** option for the given question."
|
||||
original_multiple_choice_output_format = '{enum1}'
|
||||
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Sentence', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat('Question', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
NewEnumerationPromptFormat('', 2, chosen_space=' ', chosen_separator=' ',
|
||||
chosen_item_wrapper=lambda x: f"({x})",
|
||||
chosen_number_format=lambda x: chr(ord('A') + x), object_name='enum1'),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' '
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif any(t in args.task_filename for t in ['task697_', 'task729_']):
|
||||
# task697 = ((0.19230769230769232, 0.38461538461538464, 26), '{}\n(A){} (B){} (C){} (D){}', '{}\n(A){} (B){} (C){} (D){}\nAnswer: {}')
|
||||
# <text>\n(A)<text> (B)<text> (C)<text> (D)<text>
|
||||
|
||||
# both tasks share instruction text
|
||||
instruction = 'You are given a question on formal logic. You are also given 4 answer options (associated with "{enum1_1}", "{enum1_2}", "{enum1_3}", "{enum1_4}"), out of which only one is correct. You need to answer the question by selecting the correct option. You should only answer with the choice letter, not the whole answer.' # FIXME letter -> number when needed
|
||||
original_multiple_choice_output_format = '{enum1}'
|
||||
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('', ''),
|
||||
NewEnumerationPromptFormat('', 4, chosen_space=' ', chosen_separator='',
|
||||
chosen_item_wrapper=lambda x: f"({x})",
|
||||
chosen_number_format=lambda x: chr(ord('A') + x), object_name='enum1'),
|
||||
SimplePromptFormat('Answer', ': ', is_output_field=True)
|
||||
],
|
||||
chosen_space='\n'
|
||||
)
|
||||
|
||||
elif 'task903_' in args.task_filename:
|
||||
# Review: <text>\nPolarity: <text>
|
||||
instruction = "Given a hotel review and the corresponding polarity of review (i.e., Negative or Positive) identify if the polarity is correct. Write 'true' if it's correct, 'false' otherwise."
|
||||
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Review', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat('Polarity', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space='\n'
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif 'task084_' in args.task_filename:
|
||||
# Passage: Fact 1- <text>. Fact 2- <text>. Question: <text> Answer: <text>
|
||||
|
||||
instruction = "You will be given a passage with an enumerated set of facts, a question of form 'Where is <person_name>?', and its answer. The task is to identify a supporting fact that is necessary to answer the question. The output would be the corresponding fact number." # FIXME "number" -> "letter" when it should change
|
||||
original_multiple_choice_output_format = "{enum1}"
|
||||
|
||||
min_elements, max_elements = 2, 15
|
||||
extra_params_structured_prompt_format = {'enumeration_length_range': (min_elements, max_elements + 1)}
|
||||
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Passage', None, chosen_separator_owner=chosen_separator,
|
||||
prompt_without_text=True, text_descriptor_fn_owner=text_descriptor_fn),
|
||||
NewEnumerationPromptFormat('Fact', max_elements, chosen_separator='- ', chosen_space=' ',
|
||||
chosen_item_wrapper=lambda x: f"{x}",
|
||||
text_descriptor_fn_owner=text_descriptor_fn, object_name='enum1'),
|
||||
],
|
||||
chosen_space='',
|
||||
allow_only_non_char_spaces=True
|
||||
),
|
||||
SimplePromptFormat('Question', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat('Final Output', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' '
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif 'task1297_' in args.task_filename:
|
||||
# Fact1: <text>, Fact2: <text>, Question: <text> (A) <text> (B) <text> (C) <text> (D) <text> (E) <text> (F) <text> (G) <text> (H) <text>
|
||||
instruction = 'In this task, you are given two facts, and a multiple-choice question. Based on the given facts, answer the question with index of the correct option (e.g, "{enum1_1}").'
|
||||
original_multiple_choice_output_format = "{enum1}"
|
||||
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
NewEnumerationPromptFormat('Fact', 2, chosen_separator=': ', chosen_separator_text_and_option='',
|
||||
chosen_space=', ', chosen_item_wrapper=lambda x: f"{x}",
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Question', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
NewEnumerationPromptFormat('', 8, chosen_separator=' ', chosen_space=' ',
|
||||
chosen_item_wrapper=lambda x: f"({x})",
|
||||
chosen_number_format=lambda x: chr(ord('A') + x),
|
||||
object_name='enum1'),
|
||||
],
|
||||
chosen_space=' '
|
||||
),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=', '
|
||||
)
|
||||
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif 'task065_' in args.task_filename:
|
||||
# Sentence 1: <text>\n Sentence 3: <text>\n Sentence 4: <text>\n Sentence 5: <text>\n Option 1: <text>\n Option 2: <text>
|
||||
|
||||
instruction = "In this task, you are given a short story consisting of exactly 5 sentences where the second sentence is missing. You are given two options and you need to select the one that best connects the first sentence with the rest of the story. Indicate your answer by 'Option {enum1_1}' if the first option is correct, otherwise 'Option {enum1_2}'. The incorrect option will change the subsequent storyline, so that at least one of the three subsequent sentences is no longer consistent with the story."
|
||||
original_multiple_choice_output_format = "Option {enum1}" # Idea: save chosen_number_format from the initial text, and compute chosen_number_format^-1. Then it's just a lookup table from "Option a"->1, and then we apply the current function in chosen_number_format
|
||||
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
# [0, 2, 3, 4] -> [1, 3, 4, 5] because of indexing
|
||||
NewEnumerationPromptFormat('Sentence', [0, 2, 3, 4], chosen_separator=': ', chosen_space=' \n ',
|
||||
chosen_item_wrapper=lambda x: f"{x}",
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
NewEnumerationPromptFormat('Option', 2, chosen_separator=': ', chosen_space=' \n ',
|
||||
chosen_item_wrapper=lambda x: f"{x}",
|
||||
text_descriptor_fn_owner=text_descriptor_fn, object_name='enum1'),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' \n '
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif any(t in args.task_filename for t in ['task069_', 'task070_']):
|
||||
# Beginning: <text> Middle 1: <text> Middle 2: <text> Ending: <text>
|
||||
if 'task069_' in args.task_filename:
|
||||
instruction = "In this task, you will be shown a short story with a beginning, two potential middles, and an ending. Your job is to choose the middle statement that makes the story coherent / plausible by writing \"{enum1_1}\" or \"{enum1_2}\" in the output. If both sentences are plausible, pick the one that makes most sense."
|
||||
else:
|
||||
instruction = "In this task, you will be shown a short story with a beginning, two potential middles, and an ending. Your job is to choose the middle statement that makes the story incoherent / implausible by indicating \"{enum1_1}\" or \"{enum1_2}\" in the output. If both sentences are plausible, pick the one that makes less sense."
|
||||
|
||||
original_multiple_choice_output_format = "{enum1}"
|
||||
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Beginning', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
NewEnumerationPromptFormat('Middle', 2, chosen_separator=': ', chosen_space=' ',
|
||||
chosen_number_format=lambda x: x + 1, chosen_item_wrapper=lambda x: x,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, object_name='enum1'),
|
||||
SimplePromptFormat('Ending', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' '
|
||||
)
|
||||
# FIXME allow the two chosen_spaces to change jointly (enumeration + spacing)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif any(t in args.task_filename for t in ['task137_', 'task138_', 'task139_', 'task140_']):
|
||||
# Prompt: <text> \nCompletion A: <text> \nCompletion B: <text>
|
||||
|
||||
if 'task137_' in args.task_filename:
|
||||
instruction = "Given a prompt and two completions from a language model, select the completion that is the least toxic with respect to continuing the prompt. Toxicity is defined as a rude, disrespectful, or unreasonable statement that is likely to make you leave a discussion. Generate either Completion {enum1_1} or Completion {enum1_2}."
|
||||
elif 'task138_' in args.task_filename:
|
||||
instruction = "Given a prompt and two completions from a language model, select the completion that is the most fluent with respect to continuing the prompt. Fluency is defined by coherency and ease of understanding, not necessarily grammatical correctness. Generate either Completion {enum1_1} or Completion {enum1_2}."
|
||||
elif 'task139_' in args.task_filename:
|
||||
instruction = "Given a prompt and two completions from a language model, select the completion that is more topical with respect to continuing the prompt. A prompt-completion pair is defined to be topical if the completion maintains relevance and logical succession (i.e. stays on topic) with the prompt. The flow from the prompt to the completion should be as reasonable as possible. Generate either Completion {enum1_1} or Completion {enum1_2}."
|
||||
elif 'task140_' in args.task_filename:
|
||||
instruction = "Given a prompt and two completions from a language model, select the completion that has the most similar style to the prompt. Style is defined as the tone, word choice, grammar, and sentence structure throughout the prompt-completion pair. If a prompt is colloquial, then the completion should also be colloquial, as opposed to a completion that is encyclopedic or overly formal. Generate either Completion {enum1_1} or Completion {enum1_2}."
|
||||
original_multiple_choice_output_format = "Completion {enum1}"
|
||||
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
# [0, 2, 3, 4] -> [1, 3, 4, 5] because of indexing
|
||||
SimplePromptFormat('Prompt', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
NewEnumerationPromptFormat('Completion', 2, chosen_separator=': ', chosen_space=' \n',
|
||||
chosen_number_format=lambda x: chr(ord('A') + x),
|
||||
chosen_item_wrapper=lambda x: x, text_descriptor_fn_owner=text_descriptor_fn,
|
||||
object_name='enum1'),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' \n'
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif 'task638_' in args.task_filename:
|
||||
0 / 0
|
||||
instruction = 'You are shown a conversation between a user and system. Identify who has spoken the indicated sentence based on the conversation.'
|
||||
# original_multiple_choice_output_format is complex here, but the task has been discarded anyways because of low perf
|
||||
|
||||
# Sentence1:<text> Sentence2: <text> Sentence3: <text> Question: <text> (A) <text> (B) <text>
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
|
||||
min_elements = 1
|
||||
max_elements = 45
|
||||
extra_params_structured_prompt_format = {'enumeration_length_range': (min_elements, max_elements + 1)}
|
||||
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
NewEnumerationPromptFormat('Sentence', max_elements, chosen_separator=': ', chosen_space=', ',
|
||||
chosen_separator_text_and_option='',
|
||||
chosen_item_wrapper=lambda x: f"{x}",
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Question', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
NewEnumerationPromptFormat('', 2, chosen_separator=' ', chosen_space=' ',
|
||||
chosen_item_wrapper=lambda x: f"({x})",
|
||||
chosen_number_format=lambda x: chr(ord('A') + x),
|
||||
object_name='enum1'),
|
||||
],
|
||||
chosen_space=' '
|
||||
),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' '
|
||||
)
|
||||
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif any(t in args.task_filename for t in ['task296_', 'task297_']):
|
||||
instruction = "In this task, you're given four sentences of a story written in natural language. The given story is not complete and your job is to complete the story by selecting one of the sentence choices from ({enum1_1}) and ({enum1_2}), such that the story sounds fully coherent." # FIXME also include formatting options in enum1
|
||||
original_multiple_choice_output_format = "{enum1}"
|
||||
|
||||
# Sentence1: <text> Sentence2: <text> Sentence3: <text> Sentence4: <text> \n (A) <text> (B) <text>
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
|
||||
min_elements = 1
|
||||
max_elements = 10
|
||||
extra_params_structured_prompt_format = {'enumeration_length_range': (min_elements, max_elements + 1)}
|
||||
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
NewEnumerationPromptFormat('Sentence', max_elements, chosen_separator=': ', chosen_space=' ',
|
||||
chosen_separator_text_and_option='',
|
||||
chosen_item_wrapper=lambda x: f"{x}",
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SpacingBetweenPromptComponents(
|
||||
[
|
||||
NewEnumerationPromptFormat('', 2, chosen_separator=' ', chosen_space=' ',
|
||||
chosen_item_wrapper=lambda x: f"({x})",
|
||||
chosen_number_format=lambda x: chr(ord('A') + x),
|
||||
object_name='enum1'),
|
||||
],
|
||||
chosen_space=' '
|
||||
),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' \n '
|
||||
)
|
||||
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif 'task1565_' in args.task_filename:
|
||||
# Question:<text> , Options: [A.jack miller B.bobby brown]
|
||||
# FIXME: we'd need to implement the wrapping with [...]
|
||||
0 / 0
|
||||
original_multiple_choice_output_format = "{enum1}"
|
||||
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Question', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Options', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, prompt_without_text=True),
|
||||
NewEnumerationPromptFormat('', 2, chosen_separator='', chosen_space=' ',
|
||||
chosen_item_wrapper=lambda x: f'{x}.',
|
||||
chosen_number_format=lambda x: chr(ord('A') + x),
|
||||
object_name='enum1'),
|
||||
],
|
||||
chosen_space='',
|
||||
allow_only_non_char_spaces=True
|
||||
),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' , '
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif 'task118_' in args.task_filename:
|
||||
# """<text>\n(A)68 (B)64 (C)60 (D)16 (E)15"""
|
||||
|
||||
instruction = "You are given a mathematical question described with an open-ended vocabulary. Questions in this task involve real-world situations, describing a mathematical problem. You are also given 4 or 5 answer options (associated with \"{enum1_1}\", \"{enum1_2}\", \"{enum1_3}\", \"{enum1_4}\", \"{enum1_5}\"). Do not generate anything else apart from one of the following characters: 'A', 'B, 'C', 'D', 'E'. LaTeX mathematical format (the standard way to express mathematical expressions in the typesetting software known as LaTeX) is used to express equations. Each question is solvable with high school math knowledge. Give only one answer for each question."
|
||||
original_multiple_choice_output_format = '{enum1}'
|
||||
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
NoTextPromptFormat(),
|
||||
NewEnumerationPromptFormat('', 5, chosen_separator='', chosen_separator_text_and_option='',
|
||||
chosen_space=' ', chosen_item_wrapper=lambda x: f"({x})",
|
||||
chosen_number_format=lambda x: chr(ord('A') + x), object_name='enum1'),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space='\n'
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif 'task1135_' in args.task_filename:
|
||||
instruction = "In this task, you will be presented with a question that has multiple possible answers. You should choose the most suitable option out of \"{enum1_1}\", \"{enum1_2}\", \"{enum1_3}\", \"{enum1_4}\", and \"{enum1_5}\", based on your commonsense knowledge."
|
||||
original_multiple_choice_output_format = '{enum1}'
|
||||
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Question', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Options', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, prompt_without_text=True),
|
||||
NewEnumerationPromptFormat('', 5, chosen_separator=' ', chosen_separator_text_and_option='',
|
||||
chosen_space=' ', chosen_item_wrapper=lambda x: x,
|
||||
chosen_number_format=lambda x: chr(ord('A') + x),
|
||||
object_name='enum1'),
|
||||
],
|
||||
chosen_space='',
|
||||
allow_only_non_char_spaces=True
|
||||
),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' \n '
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif any(t in args.task_filename for t in
|
||||
['task1424_', 'task1423_', 'task1422_', 'task1421_', 'task1420_', 'task1419_']):
|
||||
# Problem: <text> \nOptions: a ) <text> , b ) <text> , c ) <text> , d ) <text> , e ) <text>
|
||||
if 'task1419_' in args.task_filename:
|
||||
instruction = "In this task, you need to answer the given multiple-choice question on the gain. Gain is the value by which to multiply the input. Classify your answers into '{enum1_1}', '{enum1_2}', '{enum1_3}', '{enum1_4}', and '{enum1_5}'."
|
||||
elif 'task1420_' in args.task_filename:
|
||||
instruction = "In this task, you need to answer the given multiple-choice question on the general math. Classify your answers into '{enum1_1}', '{enum1_2}', '{enum1_3}', '{enum1_4}', and '{enum1_5}'."
|
||||
elif 'task1421_' in args.task_filename:
|
||||
instruction = "In this task, you need to provide the correct option for a given problem from the provided options."
|
||||
elif 'task1422_' in args.task_filename:
|
||||
instruction = "In this task, you need to answer the given multiple-choice question on the physics. Classify your answers into '{enum1_1}', '{enum1_2}', '{enum1_3}', '{enum1_4}', and '{enum1_5}'."
|
||||
elif 'task1423_' in args.task_filename:
|
||||
instruction = "In this task, you need to answer the given multiple-choice question on geometry. Classify your answers into '{enum1_1}', '{enum1_2}', '{enum1_3}', '{enum1_4}', and '{enum1_5}'."
|
||||
elif 'task1424_' in args.task_filename:
|
||||
instruction = "In this task, you need to provide the correct option for a given problem on probability from the provided options."
|
||||
original_multiple_choice_output_format = "{enum1}"
|
||||
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Problem', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Options', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, prompt_without_text=True),
|
||||
NewEnumerationPromptFormat('', 5, chosen_separator=' ', chosen_separator_text_and_option='',
|
||||
chosen_space=' , ', chosen_item_wrapper=lambda x: f'{x} )',
|
||||
chosen_number_format=lambda x: chr(ord('a') + x),
|
||||
object_name='enum1'),
|
||||
],
|
||||
chosen_space='',
|
||||
allow_only_non_char_spaces=True
|
||||
),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' \n'
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif 'task1678_' in args.task_filename:
|
||||
# Problem: <|text|>\nOptions: a. <|text|>, b. <|text|>, c. <|text|>, d. <|text|>, e. <|text|>
|
||||
instruction = "Given a math problem with context and a question and 5 answer choices, the task is to provide the correct answer choice based on the problem. You must choose one of the given answer choices by letter: {enum1_1}, {enum1_2}, {enum1_3}, {enum1_4}, and {enum1_5}; anything else is invalid."
|
||||
original_multiple_choice_output_format = "{enum1}"
|
||||
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Problem', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Options', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, prompt_without_text=True),
|
||||
NewEnumerationPromptFormat('', 5, chosen_separator=' ', chosen_separator_text_and_option='',
|
||||
chosen_space=', ', chosen_item_wrapper=lambda x: f'{x}.',
|
||||
chosen_number_format=lambda x: chr(ord('a') + x),
|
||||
object_name='enum1'),
|
||||
],
|
||||
chosen_space='',
|
||||
allow_only_non_char_spaces=True
|
||||
),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space='\n'
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif 'task385_' in args.task_filename or 'task580_' in args.task_filename:
|
||||
# Context: Even though she had homework to do that night, Jesse helped Skylar study.
|
||||
# Question: What will Jesse want to do next?
|
||||
# Options: (A) read homework to Skylar (B) help Skylar finish (C) skip her studying
|
||||
|
||||
if 'task385_' in args.task_filename:
|
||||
instruction = "In this task, you're given a context passage, a question, and three answer options. Your task is to return an incorrect answer option to the question from the choices given. For all questions, only one of the three answer options is correct. Pick one of the two incorrect answer options as the output."
|
||||
elif 'task580_' in args.task_filename:
|
||||
instruction = "In this task, you're given a context, a question, and three options. Your task is to find the correct answer to the question using the given context and options. Also, you may need to use commonsense reasoning about social situations to answer the questions. Classify your answers into '{enum1_1}', '{enum1_2}', and '{enum1_3}'."
|
||||
else:
|
||||
assert False
|
||||
original_multiple_choice_output_format = '{enum1}'
|
||||
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Context', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SimplePromptFormat('Question', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Options', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, prompt_without_text=True),
|
||||
NewEnumerationPromptFormat('', 3, chosen_separator=' ', chosen_separator_text_and_option='',
|
||||
chosen_space=' ', chosen_item_wrapper=lambda x: f"({x})",
|
||||
chosen_number_format=lambda x: chr(ord('A') + x),
|
||||
object_name='enum1'),
|
||||
],
|
||||
chosen_space='',
|
||||
allow_only_non_char_spaces=True
|
||||
),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' \n '
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
|
||||
elif 'task214_' in args.task_filename or 'task213_' in args.task_filename:
|
||||
# Title: The Lawsuit. Sentence 1: Denise got hit by a car. Sentence 2: She sued the driver. Sentence 3: She got a huge settlement. Sentence 4: Denise retired and moved to the beach. Choices: a. He signed up for another class to learn more. b. Her fortune was worth the pain!
|
||||
|
||||
if 'task213_' in args.task_filename:
|
||||
instruction = "In this task, you're given the title of a five-sentence story, the first four sentences, and two options for the fifth sentence as {enum1_1} and {enum1_2}. Your job is to pick the sentence option that seamlessly connects with the rest of the story, indicating your choice as '{enum1_1}' or '{enum1_2}'. If both sentences are plausible, pick the one that makes more sense."
|
||||
elif 'task214_' in args.task_filename:
|
||||
instruction = "In this task, you're given the title of a five-sentence story, the first four sentences, and two options for the fifth sentence as {enum1_1} and {enum1_2}. Your job is to pick the sentence option that does not connect with the rest of the story, indicating your choice as '{enum1_1}' or '{enum1_2}'. If both sentences are plausible, pick the one that makes less sense."
|
||||
else:
|
||||
assert False
|
||||
original_multiple_choice_output_format = '{enum1}'
|
||||
|
||||
chosen_separator = SharedPropertyAmongPrompts({'chosen_separator': ': '}, None)
|
||||
text_descriptor_fn = SharedPropertyAmongPrompts({'text_descriptor_fn': lambda x: x}, None)
|
||||
structured_prompt_format = SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Title', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn),
|
||||
NewEnumerationPromptFormat('Sentence', 4, chosen_separator_owner=chosen_separator,
|
||||
chosen_separator_text_and_option=' ',
|
||||
chosen_space=' ', chosen_item_wrapper=lambda x: f"{x}",
|
||||
chosen_number_format=lambda x: x + 1, object_name='enum0'),
|
||||
SpacingBetweenPromptComponents(
|
||||
[
|
||||
SimplePromptFormat('Choices', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, prompt_without_text=True),
|
||||
NewEnumerationPromptFormat('', 2, chosen_separator='. ', chosen_separator_text_and_option='',
|
||||
chosen_space=' ', chosen_item_wrapper=lambda x: x,
|
||||
chosen_number_format=lambda x: chr(ord('a') + x),
|
||||
object_name='enum1'),
|
||||
],
|
||||
chosen_space='',
|
||||
allow_only_non_char_spaces=True
|
||||
),
|
||||
SimplePromptFormat('Answer', None, chosen_separator_owner=chosen_separator,
|
||||
text_descriptor_fn_owner=text_descriptor_fn, is_output_field=True)
|
||||
],
|
||||
chosen_space=' '
|
||||
)
|
||||
global_constraints = [text_descriptor_fn, chosen_separator]
|
||||
else:
|
||||
# task058 = cannot be done because it has two moving length variables
|
||||
print("Unrecognized task", args.task_filename)
|
||||
return None, None, None, None, None
|
||||
|
||||
return structured_prompt_format, global_constraints, extra_params_structured_prompt_format, \
|
||||
instruction, original_multiple_choice_output_format
|
||||
@@ -0,0 +1,721 @@
|
||||
import argparse
|
||||
import copy
|
||||
import itertools
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
|
||||
from .data_loading import load_supernatural_instructions_task, load_instruction_induction_task
|
||||
from .format_evaluation import GeneticAlgorithmAmongPrompts, value_assignment_str_to_indices, \
|
||||
ThompsonSamplingAlgorithmAmongPrompts
|
||||
from .grammar_definition import pointers_to_all_objects, create_pointer_action_type_pairs, MAPPING_ALL_CATEGORIES, \
|
||||
holistic_node_format_sanity_checks
|
||||
from ...paths import PROFILE_RESULTS_ROOT, PROJECT_ROOT, model_directory, model_profile_path
|
||||
from scripts.provider_router import provider_environment
|
||||
|
||||
random.seed(0)
|
||||
|
||||
MODULE_DIRECTORY = Path(__file__).resolve().parent
|
||||
DEFAULT_NATURAL_INSTRUCTIONS_DIRECTORY = PROJECT_ROOT / 'data' / 'format-preference' / 'natural-instructions' / 'tasks'
|
||||
DEFAULT_INSTRUCTION_INDUCTION_DIRECTORY = PROJECT_ROOT / 'data' / 'format-preference' / 'instruction-induction'
|
||||
OUTPUT_ROOT = PROFILE_RESULTS_ROOT / 'format-preference'
|
||||
REMOTE_PROVIDER_ENVIRONMENT = provider_environment()
|
||||
|
||||
|
||||
def _load_model(args):
|
||||
model, tokenizer, model_will_repeat_input = None, None, False
|
||||
|
||||
if args.model_name and not args.use_gpt3:
|
||||
import torch
|
||||
cache_dir = args.cache_dir
|
||||
|
||||
if 'Llama-2-70b-hf' in args.model_name or args.use_4bit:
|
||||
# assert args.batch_size_llm == 1
|
||||
from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig
|
||||
|
||||
# torch_dtype=torch.float16 is incompatible with batching
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
args.model_name, use_fast=True, cache_dir=cache_dir, return_token_type_ids=False)
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
args.model_name, cache_dir=cache_dir, trust_remote_code=True,
|
||||
torch_dtype=torch.bfloat16,
|
||||
quantization_config=BitsAndBytesConfig(
|
||||
load_in_4bit=True,
|
||||
bnb_4bit_compute_dtype=torch.bfloat16,
|
||||
)
|
||||
)
|
||||
model_will_repeat_input = True
|
||||
|
||||
# Add special padding token
|
||||
special_tokens_dict = {'pad_token': '<pad>'}
|
||||
num_added_toks = tokenizer.add_special_tokens(special_tokens_dict)
|
||||
tokenizer.padding_side = "left"
|
||||
print('We have added', num_added_toks, 'tokens')
|
||||
|
||||
# Resize the token embeddings
|
||||
model.resize_token_embeddings(len(tokenizer))
|
||||
|
||||
# Set `pad_token_id` in model's configuration
|
||||
model.config.pad_token_id = tokenizer.pad_token_id
|
||||
|
||||
elif any(t in args.model_name.lower() for t in ['llama', 'falcon', 'mistral', 'mixtral']) \
|
||||
and args.batch_size_llm is not None:
|
||||
from transformers import AutoTokenizer, AutoModelForCausalLM
|
||||
|
||||
# torch_dtype=torch.float16 is incompatible with batching
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.model_name, use_fast=True, cache_dir=cache_dir,
|
||||
return_token_type_ids=False)
|
||||
model = AutoModelForCausalLM.from_pretrained(args.model_name, cache_dir=cache_dir, trust_remote_code=True)
|
||||
model = model.to('cuda')
|
||||
model_will_repeat_input = True
|
||||
|
||||
# Add special padding token
|
||||
special_tokens_dict = {'pad_token': '<pad>'}
|
||||
num_added_toks = tokenizer.add_special_tokens(special_tokens_dict)
|
||||
tokenizer.padding_side = "left"
|
||||
print('We have added', num_added_toks, 'tokens')
|
||||
|
||||
# Resize the token embeddings
|
||||
model.resize_token_embeddings(len(tokenizer))
|
||||
|
||||
# Set `pad_token_id` in model's configuration
|
||||
model.config.pad_token_id = tokenizer.pad_token_id
|
||||
|
||||
elif not args.use_gpt3:
|
||||
from transformers import AutoTokenizer, AutoModelForCausalLM
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
args.model_name, use_fast=True, cache_dir=cache_dir, return_token_type_ids=False)
|
||||
model = AutoModelForCausalLM.from_pretrained(args.model_name, cache_dir=cache_dir, trust_remote_code=True)
|
||||
model = model.to('cuda')
|
||||
model_will_repeat_input = True
|
||||
|
||||
model.tie_weights()
|
||||
model.eval()
|
||||
model.tie_weights()
|
||||
|
||||
return model, tokenizer, model_will_repeat_input
|
||||
|
||||
|
||||
def _load_task(args):
|
||||
if args.dataset_name == 'natural-instructions':
|
||||
from parsing_supernatural_instructions_tasks import OPEN_GENERATION_SUPERNATURAL_INSTRUCTIONS_TASKS
|
||||
args.max_new_tokens = 50 if args.task_filename in OPEN_GENERATION_SUPERNATURAL_INSTRUCTIONS_TASKS else 10
|
||||
structured_prompt_format, global_constraints, extra_params_structured_prompt_format, \
|
||||
original_multiple_choice_output_format, args_compute_node_score, raw_dataset_size = load_supernatural_instructions_task(
|
||||
args)
|
||||
elif args.dataset_name == 'instruction-induction':
|
||||
structured_prompt_format, global_constraints, extra_params_structured_prompt_format, \
|
||||
original_multiple_choice_output_format, args_compute_node_score, raw_dataset_size = load_instruction_induction_task(
|
||||
args)
|
||||
args.max_new_tokens = 15
|
||||
else:
|
||||
assert False, "No custom loading function found for this dataset."
|
||||
|
||||
return structured_prompt_format, global_constraints, extra_params_structured_prompt_format, \
|
||||
original_multiple_choice_output_format, args_compute_node_score, raw_dataset_size
|
||||
|
||||
|
||||
def _value_assignment_is_valid(structured_prompt_format, global_constraints, value_assignment, allow_text_action_type):
|
||||
# A. copy structured_prompt_format to avoid modifying the original
|
||||
new_structured_prompt_format, new_global_constraints = \
|
||||
copy.deepcopy((structured_prompt_format, global_constraints))
|
||||
all_pointers = pointers_to_all_objects(new_structured_prompt_format) + new_global_constraints
|
||||
all_pointers_enumerated = [(e, i) for i, e in enumerate(all_pointers)]
|
||||
pointer_action_pairs = create_pointer_action_type_pairs(
|
||||
all_pointers_enumerated, allow_text_action_type=allow_text_action_type)
|
||||
|
||||
# B. apply the value assignment
|
||||
value_assignments_ids = value_assignment_str_to_indices([value_assignment], pointer_action_pairs)[0]
|
||||
for (element, element_id, action_type), action_value_id in zip(pointer_action_pairs, value_assignments_ids):
|
||||
action_value, action_value_name = MAPPING_ALL_CATEGORIES[action_type][int(action_value_id)]
|
||||
element.update_field(action_type, action_value)
|
||||
|
||||
# C. evaluate new node holistically
|
||||
return holistic_node_format_sanity_checks(new_structured_prompt_format)
|
||||
|
||||
|
||||
def _sample_value_assignments(args):
|
||||
# load task [we might do it twice, but this first time is to load the structured_prompt_format]
|
||||
structured_prompt_format, global_constraints, extra_params_structured_prompt_format, \
|
||||
original_multiple_choice_output_format, args_compute_node_score, raw_dataset_size = _load_task(args)
|
||||
|
||||
# sample nodes to evaluate if file has not been passed
|
||||
all_pointers = pointers_to_all_objects(structured_prompt_format) + global_constraints
|
||||
all_pointers_enumerated = [(e, i) for i, e in enumerate(all_pointers)]
|
||||
pointer_action_pairs = create_pointer_action_type_pairs(
|
||||
all_pointers_enumerated, allow_text_action_type=args.allow_text_action_type)
|
||||
|
||||
action_value_options = []
|
||||
for a, b, action_type in pointer_action_pairs:
|
||||
action_value_options.append([f_name for f_value, f_name in MAPPING_ALL_CATEGORIES[action_type]])
|
||||
|
||||
num_combinations = 1
|
||||
for e in action_value_options:
|
||||
num_combinations *= len(e)
|
||||
|
||||
if num_combinations <= args.num_formats_to_analyze:
|
||||
valid_value_assignments = []
|
||||
for value_assignment in itertools.product(*action_value_options):
|
||||
if _value_assignment_is_valid(
|
||||
structured_prompt_format, global_constraints, value_assignment, args.allow_text_action_type):
|
||||
valid_value_assignments.append(value_assignment)
|
||||
else:
|
||||
valid_value_assignments = set()
|
||||
while len(valid_value_assignments) < args.num_formats_to_analyze:
|
||||
value_assignment = [random.choice(sublist) for sublist in action_value_options]
|
||||
if _value_assignment_is_valid(
|
||||
structured_prompt_format, global_constraints, value_assignment, args.allow_text_action_type):
|
||||
valid_value_assignments.add(tuple(value_assignment))
|
||||
valid_value_assignments = [list(e) for e in valid_value_assignments]
|
||||
|
||||
# set an order in which to shuffle the whole dataset (including demonstrations)
|
||||
dataset_ordered_ids = list(range(raw_dataset_size))
|
||||
random.shuffle(dataset_ordered_ids)
|
||||
|
||||
return valid_value_assignments, dataset_ordered_ids
|
||||
|
||||
|
||||
def _generate_neighbor_value_assignment(value_assignment, idx_to_change, action_types):
|
||||
action_type_to_change = action_types[idx_to_change]
|
||||
neighbor_value_assignment = copy.copy(value_assignment)
|
||||
|
||||
cur_value = value_assignment[idx_to_change]
|
||||
new_value = cur_value
|
||||
while new_value == cur_value:
|
||||
new_value = random.choice(MAPPING_ALL_CATEGORIES[action_type_to_change])[1]
|
||||
neighbor_value_assignment[idx_to_change] = new_value
|
||||
return neighbor_value_assignment
|
||||
|
||||
|
||||
def _sample_value_assignments_edges(args):
|
||||
# load task [we might do it twice, but this first time is to load the structured_prompt_format]
|
||||
structured_prompt_format, global_constraints, extra_params_structured_prompt_format, \
|
||||
original_multiple_choice_output_format, args_compute_node_score, raw_dataset_size = _load_task(args)
|
||||
|
||||
# sample nodes to evaluate if file has not been passed
|
||||
all_pointers = pointers_to_all_objects(structured_prompt_format) + global_constraints
|
||||
all_pointers_enumerated = [(e, i) for i, e in enumerate(all_pointers)]
|
||||
pointer_action_pairs = create_pointer_action_type_pairs(
|
||||
all_pointers_enumerated, allow_text_action_type=args.allow_text_action_type)
|
||||
|
||||
action_value_options = []
|
||||
action_types = []
|
||||
for a, b, action_type in pointer_action_pairs:
|
||||
action_value_options.append([f_name for f_value, f_name in MAPPING_ALL_CATEGORIES[action_type]])
|
||||
action_types.append(action_type)
|
||||
|
||||
valid_value_assignments = []
|
||||
while len(valid_value_assignments) < args.num_edges_to_analyze * 2:
|
||||
value_assignment = [random.choice(sublist) for sublist in action_value_options]
|
||||
|
||||
# generate value assignment with only one difference w.r.t. the current one (an "edge")
|
||||
# we decide which one to change using round robin
|
||||
idx_to_change = (len(valid_value_assignments) // 2) % len(action_types)
|
||||
neighbor_value_assignment = _generate_neighbor_value_assignment(value_assignment, idx_to_change, action_types)
|
||||
if tuple(value_assignment) in valid_value_assignments or \
|
||||
tuple(neighbor_value_assignment) in valid_value_assignments:
|
||||
continue
|
||||
|
||||
if _value_assignment_is_valid(structured_prompt_format, global_constraints, value_assignment,
|
||||
args.allow_text_action_type) and \
|
||||
_value_assignment_is_valid(structured_prompt_format, global_constraints, neighbor_value_assignment,
|
||||
args.allow_text_action_type):
|
||||
valid_value_assignments.append(tuple(value_assignment))
|
||||
valid_value_assignments.append(tuple(neighbor_value_assignment))
|
||||
|
||||
valid_value_assignments = [list(e) for e in valid_value_assignments]
|
||||
|
||||
# set an order in which to shuffle the whole dataset (including demonstrations)
|
||||
dataset_ordered_ids = list(range(raw_dataset_size))
|
||||
random.shuffle(dataset_ordered_ids)
|
||||
|
||||
return valid_value_assignments, dataset_ordered_ids
|
||||
|
||||
|
||||
def _sample_value_assignment_paths(args, existing_value_assignments):
|
||||
# load task [we might do it twice, but this first time is to load the structured_prompt_format]
|
||||
structured_prompt_format, global_constraints, extra_params_structured_prompt_format, \
|
||||
original_multiple_choice_output_format, args_compute_node_score, raw_dataset_size = _load_task(args)
|
||||
|
||||
# sample nodes to evaluate if file has not been passed
|
||||
all_pointers = pointers_to_all_objects(structured_prompt_format) + global_constraints
|
||||
all_pointers_enumerated = [(e, i) for i, e in enumerate(all_pointers)]
|
||||
pointer_action_pairs = create_pointer_action_type_pairs(
|
||||
all_pointers_enumerated, allow_text_action_type=args.allow_text_action_type)
|
||||
|
||||
action_value_options = []
|
||||
action_types = []
|
||||
for a, b, action_type in pointer_action_pairs:
|
||||
action_value_options.append([f_name for f_value, f_name in MAPPING_ALL_CATEGORIES[action_type]])
|
||||
action_types.append(action_type)
|
||||
|
||||
valid_value_assignments = []
|
||||
for value_assignment_0 in existing_value_assignments:
|
||||
found_valid_path = False
|
||||
while not found_valid_path:
|
||||
idx_to_change_1 = random.randrange(len(action_types))
|
||||
value_assignment_1 = _generate_neighbor_value_assignment(value_assignment_0, idx_to_change_1, action_types)
|
||||
|
||||
idx_to_change_2 = random.randrange(len(action_types))
|
||||
value_assignment_2 = _generate_neighbor_value_assignment(value_assignment_1, idx_to_change_2, action_types)
|
||||
|
||||
if len({tuple(value_assignment_0), tuple(value_assignment_1), tuple(value_assignment_2)}) != 3:
|
||||
continue
|
||||
|
||||
if _value_assignment_is_valid(structured_prompt_format, global_constraints, value_assignment_1,
|
||||
args.allow_text_action_type) and \
|
||||
_value_assignment_is_valid(structured_prompt_format, global_constraints, value_assignment_2,
|
||||
args.allow_text_action_type):
|
||||
valid_value_assignments.append(tuple(value_assignment_1))
|
||||
valid_value_assignments.append(tuple(value_assignment_2))
|
||||
found_valid_path = True
|
||||
|
||||
return valid_value_assignments
|
||||
|
||||
|
||||
def _get_task_filename_to_print(args):
|
||||
if args.dataset_name == 'natural-instructions':
|
||||
task_filename = args.task_filename
|
||||
to_print = task_filename.split("_")[0]
|
||||
to_print = to_print[:-5] if to_print.endswith('.json') else to_print
|
||||
elif args.dataset_name == 'instruction-induction':
|
||||
task_filename = args.task_filename.replace('_', '-')
|
||||
to_print = task_filename[:-5] if task_filename.endswith('.json') else task_filename
|
||||
else:
|
||||
assert False, "Dataset not supported."
|
||||
return to_print
|
||||
|
||||
|
||||
def _get_output_filename(args):
|
||||
scoring_type = 'rankscore' if args.evaluation_metric == 'probability_ranking' else 'genscore'
|
||||
use_4bit_str = '_4bit' if args.use_4bit else ''
|
||||
if args.evaluation_type == 'format_spread':
|
||||
filename = f'metadataholistic_{disable_text_action_type}_{scoring_type}_{task_filename_to_print}_search_model_{args.model_name.split("/")[-1]}_nshot_{args.n_shot}_numnodes_{args.num_formats_to_analyze}_numsamples_{args.num_samples}_thompson_numformats_{args.num_formats_format_spread}_batch_{args.batch_size_format_spread}_maxcalls_{args.budget_format_spread}{use_4bit_str}'
|
||||
elif args.num_formats_to_analyze:
|
||||
filename = f'metadataholistic_{disable_text_action_type}_{scoring_type}_{task_filename_to_print}_search_model_{args.model_name.split("/")[-1]}_nshot_{args.n_shot}_numnodes_{args.num_formats_to_analyze}_numsamples_{args.num_samples}{use_4bit_str}'
|
||||
elif args.num_edges_to_analyze:
|
||||
filename = f'metadataholistic_{disable_text_action_type}_{scoring_type}_{task_filename_to_print}_search_model_{args.model_name.split("/")[-1]}_nshot_{args.n_shot}_numedges_{args.num_edges_to_analyze}_numsamples_{args.num_samples}{use_4bit_str}'
|
||||
elif args.extend_graph_paths_from_file:
|
||||
# it is exactly like args.num_formats_to_analyze, but from a specific file
|
||||
filename = f'metadataholistic_{disable_text_action_type}_{scoring_type}_{task_filename_to_print}_search_model_{args.model_name.split("/")[-1]}_nshot_{args.n_shot}_numnodes-extension_{num_new_paths}_numsamples_{args.num_samples}{use_4bit_str}'
|
||||
else:
|
||||
assert False, "No output file format defined."
|
||||
|
||||
return filename
|
||||
|
||||
|
||||
def _checkpoint_config_matches(existing_config, expected_config):
|
||||
"""Accept legacy checkpoints that predate ``model_identifier``.
|
||||
|
||||
Older checkpoints stored only the provider-local model name. Their
|
||||
remaining settings still identify the exact same run, so rejecting them
|
||||
forces unnecessary API calls after a provider-qualified model migration.
|
||||
"""
|
||||
if not isinstance(existing_config, dict):
|
||||
return False
|
||||
for key, value in expected_config.items():
|
||||
if key == 'model_identifier' and key not in existing_config:
|
||||
continue
|
||||
if existing_config.get(key) != value:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _result_has_only_nonempty_generations(result_path):
|
||||
"""Reject completed caches whose API calls produced empty final answers."""
|
||||
try:
|
||||
with open(result_path, 'r') as result_file:
|
||||
result = json.load(result_file)
|
||||
except (OSError, json.JSONDecodeError):
|
||||
return False
|
||||
|
||||
generations = []
|
||||
|
||||
def collect(value):
|
||||
if isinstance(value, dict):
|
||||
if 'generation' in value:
|
||||
generations.append(value['generation'])
|
||||
for child in value.values():
|
||||
collect(child)
|
||||
elif isinstance(value, list):
|
||||
for child in value:
|
||||
collect(child)
|
||||
|
||||
collect(result)
|
||||
return bool(generations) and all(
|
||||
isinstance(generation, str) and generation.strip()
|
||||
for generation in generations
|
||||
)
|
||||
|
||||
|
||||
def _best_worst_accuracy(node_accuracies):
|
||||
"""Extract scalar right-answer rates from list_node_accuracies entries."""
|
||||
if not node_accuracies:
|
||||
raise ValueError('format evaluation produced no node accuracies')
|
||||
right_rates = [entry[0][0] for entry in node_accuracies]
|
||||
return max(right_rates), min(right_rates)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# python main.py --task_filename singular_to_plural.json --num_formats_to_analyze 5 --batch_size_llm 10 --model_name "meta-llama/Llama-2-7b-hf" --n_shot 5
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
# params to load a task
|
||||
parser.add_argument('--task_filename', type=str, default='task158_',
|
||||
help='Benchmark task. Defaults to the format-preference baseline task158_.')
|
||||
parser.add_argument('--dataset_name', type=str, choices=['natural-instructions', 'instruction-induction'],
|
||||
default='natural-instructions', help='Dataset containing --task_filename.')
|
||||
parser.add_argument('--natural_instructions_dir', type=str,
|
||||
default=os.getenv('NATURAL_INSTRUCTIONS_DIR', str(DEFAULT_NATURAL_INSTRUCTIONS_DIRECTORY)),
|
||||
help='Path to the natural-instructions tasks directory.')
|
||||
parser.add_argument('--instruction_induction_dir', type=str,
|
||||
default=os.getenv('INSTRUCTION_INDUCTION_DIR', str(DEFAULT_INSTRUCTION_INDUCTION_DIRECTORY)),
|
||||
help='Path to the instruction-induction repository directory.')
|
||||
|
||||
# params to create or load a set of formats to evaluate
|
||||
parser.add_argument('--num_formats_to_analyze', type=int, default=9,
|
||||
help='Number of sampled format variants; the original format is evaluated as well.')
|
||||
parser.add_argument('--num_edges_to_analyze', type=int, default=None, help='Use for atomic changes experiment.')
|
||||
parser.add_argument('--extend_graph_paths_from_file', type=str, default=None,
|
||||
help='Use solely for non-monotonic paths experiment. Only include filename of old 499 samples file.')
|
||||
parser.add_argument('--nodes_to_evaluate_filepath', type=str, default=None,
|
||||
help='Filepath containing the formats to evaluate. If no file is passed, '
|
||||
'the script loads the default file if available, or creates it if it does not exist.')
|
||||
|
||||
# params to set up evaluation settings
|
||||
parser.add_argument('--num_samples', type=int, default=100, help='Maximum number of samples to evaluate for each format.')
|
||||
parser.add_argument('--evaluation_metric', choices=['exact_prefix_matching', 'probability_ranking'],
|
||||
default='exact_prefix_matching')
|
||||
parser.add_argument('--evaluation_type', type=str, choices=['full', 'format_spread'],
|
||||
default='full',
|
||||
help='Determines how to evaluate the array of formats defined. '
|
||||
'Options are full evaluation of each node, or use Thompson Sampling to quickly find the format spread.')
|
||||
parser.add_argument('--n_shot', type=int, default=1)
|
||||
|
||||
# params to load models and how to use them
|
||||
parser.add_argument('--model', '--model_name', dest='model_name', type=str, required=True,
|
||||
help='Canonical provider/model-id, e.g. siliconflow/Qwen/Qwen2.5-72B-Instruct.')
|
||||
parser.add_argument('--api_provider', choices=['auto', 'local', *REMOTE_PROVIDER_ENVIRONMENT], default='auto',
|
||||
help='Optional legacy provider override. By default it is parsed from --model.')
|
||||
parser.add_argument('--api_url_env', type=str, default=None,
|
||||
help='Environment-variable name containing the Chat Completions URL. Defaults depend on --api_provider.')
|
||||
parser.add_argument('--api_key_env', type=str, default=None,
|
||||
help='Environment-variable name containing the API key. Defaults depend on --api_provider.')
|
||||
parser.add_argument('--api_concurrency', type=int, default=3,
|
||||
help='Maximum number of simultaneous remote API requests. Only used with a remote --api_provider.')
|
||||
parser.add_argument('--batch_size_llm', type=int, default=2, help='Batch size to call the LLM.')
|
||||
parser.add_argument('--use_4bit', action='store_true')
|
||||
parser.add_argument('--cache_dir', type=str, default='/gscratch/xlab/msclar/.cache')
|
||||
|
||||
# FormatSpread-specific parameters, corresponding to Thompson Sampling
|
||||
parser.add_argument('--num_formats_format_spread', type=int, default=320, help='Number of formats to sample.')
|
||||
parser.add_argument('--batch_size_format_spread', type=int, default=20, help='Batch size used by FormatSpread when running Thompson Sampling. Only used with `--evaluation_type format_spread`')
|
||||
parser.add_argument('--budget_format_spread', type=int, default=40000, help='Maximum number of model calls allowed when exploring best and worst formats, i.e. budget for thompson sampling. Only used with `--evaluation_type format_spread`')
|
||||
|
||||
# saving parameters
|
||||
parser.add_argument('--output_dir', type=str, default=None,
|
||||
help='Directory for checkpoints and result metadata. Defaults to this module\'s results directory.')
|
||||
parser.add_argument('--checkpoint_path', type=str, default=None,
|
||||
help='JSON checkpoint for full evaluation. Defaults to a task-specific file in --output_dir.')
|
||||
parser.add_argument('--profile_path', type=str, default=None,
|
||||
help='Final profile JSON to write after evaluation. Defaults below results/static-opimization/profiles/models/.')
|
||||
parser.add_argument('--base_profile_path', '--base-profile-path', dest='base_profile_path', type=str,
|
||||
default=None,
|
||||
help='Read-only upstream profile used as a template for the final profile.')
|
||||
parser.add_argument('--format_sensitivity_threshold', type=float, default=0.05,
|
||||
help='Strict accuracy-spread threshold used for profile classification.')
|
||||
parser.add_argument('--profile_top_k', type=int, default=3,
|
||||
help='Number of best and worst formats retained in the profile field.')
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Preferred input is provider/model-id, preserving model namespace slashes:
|
||||
# ``siliconflow/Qwen/Qwen2.5-72B-Instruct`` becomes provider
|
||||
# ``siliconflow`` and API model ID ``Qwen/Qwen2.5-72B-Instruct``.
|
||||
input_model_identifier = args.model_name.strip('/')
|
||||
input_provider, separator, provider_model_name = input_model_identifier.partition('/')
|
||||
if args.api_provider == 'auto':
|
||||
if not separator or input_provider not in {*REMOTE_PROVIDER_ENVIRONMENT, 'local'}:
|
||||
parser.error('--model must use provider/model-id, for example siliconflow/Qwen/Qwen2.5-72B-Instruct.')
|
||||
args.api_provider = input_provider
|
||||
args.model_name = provider_model_name
|
||||
elif separator and input_provider == args.api_provider:
|
||||
args.model_name = provider_model_name
|
||||
else:
|
||||
args.model_name = input_model_identifier
|
||||
args.model_identifier = f'{args.api_provider}/{args.model_name}'
|
||||
|
||||
try:
|
||||
if args.output_dir is None:
|
||||
args.output_dir = str(model_directory(OUTPUT_ROOT, args.model_identifier))
|
||||
if args.profile_path is None:
|
||||
args.profile_path = str(
|
||||
model_profile_path(OUTPUT_ROOT.parent / 'models', args.model_identifier)
|
||||
)
|
||||
except ValueError as error:
|
||||
parser.error(str(error))
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# note: earlier version of the code allowed to vary the text for synonyms, but that has been deprecated
|
||||
args.disable_text_action_type = True
|
||||
args.allow_text_action_type = not args.disable_text_action_type
|
||||
disable_text_action_type = 'textdisabled'
|
||||
|
||||
# ``use_gpt3`` is retained as an internal flag for backward compatibility with
|
||||
# the evaluation code. It now means any remote Chat Completions provider.
|
||||
args.use_gpt3 = args.api_provider in REMOTE_PROVIDER_ENVIRONMENT
|
||||
args.gpt3_engine = args.model_name if args.use_gpt3 else None
|
||||
if args.use_gpt3:
|
||||
default_url_env, default_key_env = REMOTE_PROVIDER_ENVIRONMENT[args.api_provider]
|
||||
args.api_url_env = args.api_url_env or default_url_env
|
||||
args.api_key_env = args.api_key_env or default_key_env
|
||||
if args.use_gpt3 and not args.model_name:
|
||||
parser.error('--model_name is required for remote API evaluation.')
|
||||
if args.use_gpt3 and args.evaluation_metric == 'probability_ranking':
|
||||
parser.error('probability_ranking requires local model logits; use exact_prefix_matching with OpenCode.')
|
||||
if args.api_concurrency < 1:
|
||||
parser.error('--api_concurrency must be at least 1.')
|
||||
if not 0 <= args.format_sensitivity_threshold <= 1:
|
||||
parser.error('--format_sensitivity_threshold must be between 0 and 1.')
|
||||
if args.profile_top_k < 1:
|
||||
parser.error('--profile_top_k must be at least 1.')
|
||||
|
||||
assert args.num_samples % args.batch_size_llm == 0 # for simplicity
|
||||
assert args.batch_size_format_spread % args.batch_size_llm == 0 if args.evaluation_type == 'format_spread' else True # for simplicity
|
||||
assert len(
|
||||
[e for e in [args.num_formats_to_analyze, args.num_edges_to_analyze, args.extend_graph_paths_from_file] if
|
||||
e is not None]) == 1
|
||||
if args.extend_graph_paths_from_file is not None:
|
||||
assert args.task_filename in args.extend_graph_paths_from_file
|
||||
|
||||
demonstrations_filename_suffix = ''
|
||||
|
||||
# 0. load sampled formats (or sample formats if they are not available)
|
||||
task_filename_to_print = _get_task_filename_to_print(args)
|
||||
if args.num_formats_to_analyze:
|
||||
shared_sample_path = PROJECT_ROOT / 'data' / 'format-preference' / 'format-samples' / (
|
||||
f'holistic_random_sample_{task_filename_to_print}_nodes_{args.num_formats_to_analyze}_'
|
||||
f'{disable_text_action_type}.json'
|
||||
)
|
||||
sample_path = Path(args.nodes_to_evaluate_filepath) if args.nodes_to_evaluate_filepath else shared_sample_path
|
||||
if sample_path.exists():
|
||||
tmp = json.load(open(sample_path, 'r'))
|
||||
valid_value_assignments = tmp['valid_value_assignments']
|
||||
dataset_ordered_ids = tmp['dataset_ordered_ids']
|
||||
else:
|
||||
valid_value_assignments, dataset_ordered_ids = _sample_value_assignments(args)
|
||||
sample_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
json.dump({'valid_value_assignments': valid_value_assignments,
|
||||
'dataset_ordered_ids': dataset_ordered_ids}, open(sample_path, 'w'))
|
||||
print('Created shared sample and stored it in', sample_path)
|
||||
|
||||
args.dataset_ordered_ids = dataset_ordered_ids # used in data loading
|
||||
elif args.num_edges_to_analyze:
|
||||
filepath = os.path.join(args.output_dir,
|
||||
f'holistic_random_sample_{task_filename_to_print}_edges_{args.num_edges_to_analyze}_{disable_text_action_type}.json')
|
||||
if args.nodes_to_evaluate_filepath:
|
||||
tmp = json.load(open(args.nodes_to_evaluate_filepath, 'r'))
|
||||
valid_value_assignments = tmp['valid_value_assignments']
|
||||
dataset_ordered_ids = tmp['dataset_ordered_ids']
|
||||
elif os.path.exists(filepath):
|
||||
tmp = json.load(open(filepath, 'r'))
|
||||
valid_value_assignments = tmp['valid_value_assignments']
|
||||
dataset_ordered_ids = tmp['dataset_ordered_ids']
|
||||
else:
|
||||
valid_value_assignments, dataset_ordered_ids = _sample_value_assignments_edges(args)
|
||||
json.dump({'valid_value_assignments': valid_value_assignments,
|
||||
'dataset_ordered_ids': dataset_ordered_ids}, open(filepath, 'w'))
|
||||
print('Created sample and stored it in', filepath)
|
||||
|
||||
args.dataset_ordered_ids = dataset_ordered_ids # used in data loading
|
||||
elif args.extend_graph_paths_from_file:
|
||||
"""
|
||||
We have a file with already analyzed nodes (~499) and we want to sample a bunch of paths v_1->v_2->v_3.
|
||||
We cap it to 300*2 new nodes to analyze.
|
||||
"""
|
||||
num_new_paths = 300
|
||||
filepath = os.path.join(
|
||||
args.output_dir,
|
||||
f'extension_{num_new_paths}_paths_from_{args.extend_graph_paths_from_file}'
|
||||
)
|
||||
|
||||
if os.path.exists(filepath):
|
||||
tmp = json.load(open(filepath, 'r'))
|
||||
valid_value_assignments = tmp['valid_value_assignments']
|
||||
dataset_ordered_ids = tmp['dataset_ordered_ids']
|
||||
else:
|
||||
assert os.path.exists(os.path.join(args.output_dir, args.extend_graph_paths_from_file))
|
||||
tmp = json.load(open(os.path.join(args.output_dir, args.extend_graph_paths_from_file), 'r'))
|
||||
existing_value_assignments = tmp['valid_value_assignments']
|
||||
dataset_ordered_ids = tmp['dataset_ordered_ids']
|
||||
|
||||
assert len(existing_value_assignments) >= num_new_paths
|
||||
valid_value_assignments = _sample_value_assignment_paths(args, existing_value_assignments[:num_new_paths])
|
||||
json.dump({'valid_value_assignments': valid_value_assignments,
|
||||
'dataset_ordered_ids': dataset_ordered_ids}, open(filepath, 'w'))
|
||||
|
||||
# A fully checkpointed result needs no model loading or API calls. Check
|
||||
# this before constructing the evaluation tree, whose baseline node would
|
||||
# otherwise be evaluated again.
|
||||
result_path = Path(args.output_dir) / f'{_get_output_filename(args)}.json'
|
||||
checkpoint_path = args.checkpoint_path or os.path.join(
|
||||
args.output_dir,
|
||||
f'checkpoint_{task_filename_to_print}_{args.model_name.replace("/", "_")}_nshot_{args.n_shot}_'
|
||||
f'numnodes_{args.num_formats_to_analyze}_numsamples_{args.num_samples}.json')
|
||||
checkpoint_config = {
|
||||
'task_filename': args.task_filename,
|
||||
'dataset_name': args.dataset_name,
|
||||
'model_identifier': args.model_identifier,
|
||||
'model_name': args.model_name,
|
||||
'n_shot': args.n_shot,
|
||||
'num_formats_to_analyze': args.num_formats_to_analyze,
|
||||
'num_samples': args.num_samples,
|
||||
'evaluation_metric': args.evaluation_metric,
|
||||
}
|
||||
checkpoint = {'config': checkpoint_config, 'completed_value_assignments': []}
|
||||
if os.path.exists(checkpoint_path):
|
||||
checkpoint = json.load(open(checkpoint_path, 'r'))
|
||||
if not _checkpoint_config_matches(checkpoint.get('config'), checkpoint_config):
|
||||
parser.error(f'Checkpoint settings do not match this run: {checkpoint_path}')
|
||||
if args.evaluation_type == 'full' and \
|
||||
len(checkpoint['completed_value_assignments']) >= len(valid_value_assignments) and result_path.exists():
|
||||
if _result_has_only_nonempty_generations(result_path):
|
||||
print('Format evaluation is already complete; reusing cached results.')
|
||||
if args.profile_path:
|
||||
from .update_profile import update_profile
|
||||
profile = update_profile(
|
||||
Path(args.profile_path), result_path,
|
||||
threshold=args.format_sensitivity_threshold,
|
||||
top_k=args.profile_top_k,
|
||||
model_id=args.model_identifier,
|
||||
display_name=args.model_identifier,
|
||||
base_profile_path=(Path(args.base_profile_path) if args.base_profile_path else None),
|
||||
)
|
||||
print(
|
||||
f"Updated profile {args.profile_path}: "
|
||||
f"{profile['format_preference']['classification']} "
|
||||
f"(spread={profile['format_preference']['strict_accuracy_spread']:.1%})."
|
||||
)
|
||||
raise SystemExit(0)
|
||||
print(
|
||||
'Completed format cache contains empty generations; '
|
||||
'discarding its completion markers and rebuilding it.'
|
||||
)
|
||||
backup_suffix = '.invalid-empty-generations.bak'
|
||||
for invalid_path in (Path(result_path), Path(checkpoint_path)):
|
||||
backup_path = invalid_path.with_name(invalid_path.name + backup_suffix)
|
||||
if invalid_path.exists() and not backup_path.exists():
|
||||
shutil.copy2(invalid_path, backup_path)
|
||||
print(f'Backed up invalid cache to {backup_path}.')
|
||||
checkpoint = {'config': checkpoint_config, 'completed_value_assignments': []}
|
||||
|
||||
# 1. load task
|
||||
structured_prompt_format, global_constraints, extra_params_structured_prompt_format, \
|
||||
original_multiple_choice_output_format, args_compute_node_score, _ = _load_task(args)
|
||||
print('Task loaded.')
|
||||
|
||||
# 1.b. check that the evaluation metric is reasonable
|
||||
# Specifically, we can compute probability ranking metric only if the task is a classification task
|
||||
output_options_size = len(set([e for d in args_compute_node_score['dataset'] for e in d['output']]))
|
||||
assert output_options_size < 10 if args.evaluation_metric == 'probability_ranking' else True
|
||||
|
||||
# 2. load model
|
||||
model, tokenizer, model_will_repeat_input = _load_model(args)
|
||||
print('Model loaded.')
|
||||
|
||||
args_compute_node_score['model'] = model
|
||||
args_compute_node_score['tokenizer'] = tokenizer
|
||||
args_compute_node_score['model_will_repeat_input'] = model_will_repeat_input
|
||||
args_compute_node_score['args'].use_gpt3 = args.use_gpt3
|
||||
args_compute_node_score['args'].gpt3_engine = args.gpt3_engine
|
||||
|
||||
# 3. evaluate formats
|
||||
print('Start evaluation of formats.')
|
||||
if args.evaluation_type == 'format_spread':
|
||||
search_tree = ThompsonSamplingAlgorithmAmongPrompts(
|
||||
structured_prompt_format,
|
||||
global_constraints,
|
||||
extra_params_structured_prompt_format,
|
||||
args_compute_node_score=args_compute_node_score,
|
||||
objective='lowest_accuracy', # dummy in this mode
|
||||
allow_text_action_type=args.allow_text_action_type,
|
||||
original_multiple_choice_output_format=original_multiple_choice_output_format
|
||||
)
|
||||
|
||||
search_tree.main(
|
||||
value_assignments=valid_value_assignments[:args.num_formats_format_spread + 1],
|
||||
batch_size=args.batch_size_format_spread,
|
||||
num_formats=args.num_formats_format_spread,
|
||||
max_allowed_number_of_model_calls=args.budget_format_spread
|
||||
)
|
||||
|
||||
elif args.evaluation_type == 'full':
|
||||
# exhaustive node evaluation
|
||||
search_tree = GeneticAlgorithmAmongPrompts(
|
||||
structured_prompt_format,
|
||||
global_constraints,
|
||||
extra_params_structured_prompt_format,
|
||||
args_compute_node_score=args_compute_node_score,
|
||||
objective='lowest_accuracy', # dummy in this mode
|
||||
allow_text_action_type=args.allow_text_action_type,
|
||||
original_multiple_choice_output_format=original_multiple_choice_output_format
|
||||
)
|
||||
|
||||
completed_value_assignments = checkpoint['completed_value_assignments']
|
||||
completed_value_assignment_keys = {tuple(assignment) for assignment in completed_value_assignments}
|
||||
previous_result = None
|
||||
if completed_value_assignments and result_path.exists():
|
||||
previous_result = json.load(open(result_path, 'r'))
|
||||
|
||||
def save_checkpoint(value_assignment):
|
||||
completed_value_assignments.append(value_assignment)
|
||||
temporary_path = checkpoint_path + '.tmp'
|
||||
with open(temporary_path, 'w') as checkpoint_file:
|
||||
json.dump(checkpoint, checkpoint_file)
|
||||
os.replace(temporary_path, checkpoint_path)
|
||||
# Save detailed metadata at the same boundary as the checkpoint.
|
||||
# If the process is interrupted later, completed assignments and
|
||||
# their scores/logs remain consistent for a genuine resume.
|
||||
search_tree.save(result_path, previous_result=previous_result)
|
||||
|
||||
print(f'Prepared {len(valid_value_assignments)} format variant(s) for evaluation.')
|
||||
if completed_value_assignments:
|
||||
print(f'Resuming from checkpoint: {len(completed_value_assignments)} completed format(s).')
|
||||
search_tree.main(
|
||||
value_assignments=valid_value_assignments,
|
||||
num_samples_to_test=args.num_samples,
|
||||
skip_value_assignments=completed_value_assignment_keys,
|
||||
on_node_evaluated=save_checkpoint,
|
||||
)
|
||||
|
||||
acc = search_tree.list_node_accuracies()
|
||||
best_accuracy, worst_accuracy = _best_worst_accuracy(acc)
|
||||
print(
|
||||
f'Format evaluation finished: best accuracy={best_accuracy:.1%}, '
|
||||
f'worst accuracy={worst_accuracy:.1%}.'
|
||||
)
|
||||
|
||||
if args.evaluation_type == 'full':
|
||||
search_tree.save(result_path, previous_result=previous_result)
|
||||
else:
|
||||
result_path = Path(args.output_dir) / f'{_get_output_filename(args)}.json'
|
||||
search_tree.save(result_path)
|
||||
if args.profile_path:
|
||||
from .update_profile import update_profile
|
||||
profile = update_profile(
|
||||
Path(args.profile_path),
|
||||
result_path,
|
||||
threshold=args.format_sensitivity_threshold,
|
||||
top_k=args.profile_top_k,
|
||||
model_id=args.model_identifier,
|
||||
display_name=args.model_identifier,
|
||||
base_profile_path=(Path(args.base_profile_path) if args.base_profile_path else None),
|
||||
)
|
||||
print(
|
||||
f"Updated profile {args.profile_path}: "
|
||||
f"{profile['format_preference']['classification']} "
|
||||
f"(spread={profile['format_preference']['strict_accuracy_spread']:.1%})."
|
||||
)
|
||||
@@ -0,0 +1,145 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Add a FormatSpread-derived format-preference section to a model profile."""
|
||||
|
||||
import json
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
from ...paths import PROJECT_ROOT
|
||||
|
||||
|
||||
|
||||
def relative_to_project(path):
|
||||
"""Return a project-relative artifact path when possible."""
|
||||
try:
|
||||
return str(path.resolve().relative_to(PROJECT_ROOT))
|
||||
except ValueError:
|
||||
return str(path)
|
||||
|
||||
|
||||
def first_numeric_token_is_correct(log):
|
||||
"""A task-specific content proxy for numeric-answer tasks such as task158."""
|
||||
match = re.search(r'(?<!\d)\d+(?!\d)', str(log.get('generation', '')))
|
||||
return match is not None and match.group() == str(log['entry']['output'][0])
|
||||
|
||||
|
||||
def load_nodes(result):
|
||||
accuracies = result['all_structured_prompt_formats_accuracies']
|
||||
generation_order = result['generation_order']
|
||||
histories = result.get('metadata', {}).get('nodes', {})
|
||||
nodes = []
|
||||
|
||||
for prompt, (strict_accuracy, wrong_rate, total) in accuracies.items():
|
||||
score, logs = histories.get(prompt, ({}, []))
|
||||
numeric_accuracy = (
|
||||
sum(first_numeric_token_is_correct(log) for log in logs) / len(logs)
|
||||
if logs else None
|
||||
)
|
||||
nodes.append({
|
||||
'format_order': generation_order[prompt],
|
||||
'is_original_format': generation_order[prompt] == 0,
|
||||
'prompt_format': prompt,
|
||||
'strict_accuracy': strict_accuracy,
|
||||
'first_numeric_token_accuracy': numeric_accuracy,
|
||||
'right_count': sum(score.get('right', [])),
|
||||
'wrong_answer_count': sum(score.get('wrong', [])),
|
||||
'format_or_other_count': sum(score.get('other', [])),
|
||||
'sample_count': total,
|
||||
'wrong_rate': wrong_rate,
|
||||
})
|
||||
return sorted(nodes, key=lambda node: node['format_order'])
|
||||
|
||||
|
||||
def build_format_preference(result_path, result, threshold, top_k):
|
||||
nodes = load_nodes(result)
|
||||
if not nodes:
|
||||
raise ValueError('The FormatSpread result contains no evaluated formats.')
|
||||
|
||||
ranked_best = sorted(nodes, key=lambda node: (-node['strict_accuracy'], node['format_order']))
|
||||
ranked_worst = sorted(nodes, key=lambda node: (node['strict_accuracy'], node['format_order']))
|
||||
best_accuracy = ranked_best[0]['strict_accuracy']
|
||||
worst_accuracy = ranked_worst[0]['strict_accuracy']
|
||||
numeric_accuracies = [node['first_numeric_token_accuracy'] for node in nodes]
|
||||
total_observations = sum(node['sample_count'] for node in nodes)
|
||||
strict_spread = round(best_accuracy - worst_accuracy, 4)
|
||||
numeric_spread = round(max(numeric_accuracies) - min(numeric_accuracies), 4)
|
||||
|
||||
result_prefix = result_path.stem.split('_search_model_', 1)[0]
|
||||
task_label = re.split(r'_(?:gen|rank)score_', result_prefix, maxsplit=1)[-1]
|
||||
|
||||
def compact_format(node):
|
||||
return {
|
||||
'prompt_format': node['prompt_format'],
|
||||
'strict_accuracy': node['strict_accuracy'],
|
||||
}
|
||||
|
||||
return {
|
||||
'classification': (
|
||||
'format_sensitive'
|
||||
if strict_spread >= threshold
|
||||
else 'format_insensitive'
|
||||
),
|
||||
'strict_accuracy_spread': strict_spread,
|
||||
'best_formats': [compact_format(node) for node in ranked_best[:top_k]],
|
||||
'worst_formats': [compact_format(node) for node in ranked_worst[:top_k]],
|
||||
}
|
||||
|
||||
|
||||
def default_profile(model_id, display_name):
|
||||
return {
|
||||
'schema_version': '1.0',
|
||||
'model': {
|
||||
'id': model_id,
|
||||
'display_name': display_name or model_id,
|
||||
'profile_status': 'partial',
|
||||
},
|
||||
'provenance': {},
|
||||
'behavioral_profile': {},
|
||||
'artifacts': {},
|
||||
'validation': {},
|
||||
'interpretation_cautions': [],
|
||||
}
|
||||
|
||||
|
||||
def update_profile(profile_path, result_path, threshold=0.05, top_k=3, model_id=None, display_name=None,
|
||||
base_profile_path=None):
|
||||
"""Write a format-preference-enriched profile to *profile_path*.
|
||||
|
||||
When *base_profile_path* is supplied, it is the authoritative read-only
|
||||
behavioral profile for this merge. This prevents a stale combined output
|
||||
from overriding freshly rebuilt behavioral results.
|
||||
"""
|
||||
if not 0 <= threshold <= 1:
|
||||
raise ValueError('threshold must be between 0 and 1')
|
||||
if top_k < 1:
|
||||
raise ValueError('top_k must be at least 1')
|
||||
|
||||
with result_path.open(encoding='utf-8') as result_file:
|
||||
result = json.load(result_file)
|
||||
if base_profile_path is not None:
|
||||
if not base_profile_path.is_file():
|
||||
raise ValueError(f'Base profile not found: {base_profile_path}')
|
||||
with base_profile_path.open(encoding='utf-8') as profile_file:
|
||||
profile = json.load(profile_file)
|
||||
elif profile_path.exists():
|
||||
with profile_path.open(encoding='utf-8') as profile_file:
|
||||
profile = json.load(profile_file)
|
||||
else:
|
||||
if not model_id:
|
||||
raise ValueError('--model-id is required when creating a new profile')
|
||||
profile = default_profile(model_id, display_name)
|
||||
|
||||
profile['format_preference'] = build_format_preference(result_path, result, threshold, top_k)
|
||||
profile.setdefault('artifacts', {})['format_preference_result'] = relative_to_project(result_path)
|
||||
cautions = profile.setdefault('interpretation_cautions', [])
|
||||
caution = (
|
||||
'Format-preference scores are task-, metric-, sample-, and provider-specific; strict exact-match '
|
||||
'sensitivity can reflect output rendering rather than task-content errors.'
|
||||
)
|
||||
if caution not in cautions:
|
||||
cautions.append(caution)
|
||||
|
||||
profile_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
profile_path.write_text(json.dumps(profile, ensure_ascii=False, indent=2) + '\n', encoding='utf-8')
|
||||
return profile
|
||||
|
||||
@@ -0,0 +1,510 @@
|
||||
import copy
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
|
||||
import requests
|
||||
from dotenv import load_dotenv
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from .grammar_definition import apply_prompt_format, flatten
|
||||
|
||||
# Load this project's .env when the script is launched from the project root.
|
||||
# Existing shell environment variables still take precedence.
|
||||
load_dotenv()
|
||||
PRINT_HIDDEN_STATE = False
|
||||
|
||||
|
||||
def call_openai_api_with_retry(args, prompt, max_tokens=10):
|
||||
"""Call an OpenAI-compatible Chat Completions endpoint without local ML dependencies."""
|
||||
url = os.getenv(args.api_url_env)
|
||||
api_key = os.getenv(args.api_key_env)
|
||||
if not url or not api_key:
|
||||
raise RuntimeError(
|
||||
f'Missing {args.api_url_env} or {args.api_key_env}. Put both values in .env or export them.')
|
||||
|
||||
payload = {
|
||||
'model': args.gpt3_engine,
|
||||
'messages': [
|
||||
{'role': 'system', 'content': 'You are a helpful assistant.'},
|
||||
{'role': 'user', 'content': prompt},
|
||||
],
|
||||
'max_tokens': max_tokens,
|
||||
'temperature': 0,
|
||||
'top_p': 1.0,
|
||||
}
|
||||
if (
|
||||
args.api_provider == 'siliconflow'
|
||||
and args.gpt3_engine.startswith('Qwen/Qwen3.5-')
|
||||
):
|
||||
payload['enable_thinking'] = False
|
||||
for attempt in range(4):
|
||||
try:
|
||||
response = requests.post(
|
||||
url,
|
||||
headers={'Authorization': f'Bearer {api_key}', 'Content-Type': 'application/json'},
|
||||
json=payload,
|
||||
timeout=120,
|
||||
)
|
||||
# Only transient errors are retried. Configuration and authentication
|
||||
# errors should be reported immediately.
|
||||
if response.status_code == 429 or response.status_code >= 500:
|
||||
response.raise_for_status()
|
||||
response.raise_for_status()
|
||||
result = response.json()
|
||||
generation = result['choices'][0]['message']['content']
|
||||
if not isinstance(generation, str):
|
||||
raise RuntimeError(f'Unexpected completion content: {generation!r}')
|
||||
tokens_used = result.get('usage', {}).get('total_tokens', 0)
|
||||
return generation.strip(), tokens_used
|
||||
except (requests.Timeout, requests.ConnectionError, requests.HTTPError) as error:
|
||||
retryable = isinstance(error, (requests.Timeout, requests.ConnectionError)) or \
|
||||
getattr(error.response, 'status_code', 0) == 429 or \
|
||||
getattr(error.response, 'status_code', 0) >= 500
|
||||
if not retryable or attempt == 3:
|
||||
raise RuntimeError(f'OpenCode request failed: {error}') from error
|
||||
wait_seconds = 2 ** attempt
|
||||
print(f'OpenCode request failed ({error}); retrying in {wait_seconds}s.')
|
||||
time.sleep(wait_seconds)
|
||||
|
||||
|
||||
def query_model_parallelized(model, tokenizer, prompt_list, max_tokens, top_p, temperature):
|
||||
import torch
|
||||
inputs = tokenizer(prompt_list, padding=True, return_tensors='pt', return_token_type_ids=False).to('cuda')
|
||||
|
||||
with torch.no_grad():
|
||||
outputs = model.generate(
|
||||
**inputs, top_p=top_p, temperature=temperature, max_new_tokens=max_tokens,
|
||||
return_dict_in_generate=True, output_hidden_states=True, output_attentions=False, output_scores=True
|
||||
)
|
||||
|
||||
logits_list = [[] for _ in range(len(prompt_list))]
|
||||
|
||||
# we do not print hidden state and scores because it is too much memory spenditure
|
||||
if PRINT_HIDDEN_STATE:
|
||||
# take the first (0th) inference. Its last layer (-1) will have shape [1, prompt_size, 4096]. Take last one.
|
||||
final_prompt_hidden_state_list = [
|
||||
outputs['hidden_states'][0][-1][i, -1, :].tolist() for i in range(len(prompt_list))]
|
||||
else:
|
||||
for new_token_idx in range(len(outputs['scores'])):
|
||||
for i in range(len(prompt_list)):
|
||||
logits = torch.topk(outputs['scores'][new_token_idx][i, :], k=100)
|
||||
logits = [(value, index) for value, index in zip(logits.values.tolist(), logits.indices.tolist())]
|
||||
logits_list[i].append(logits)
|
||||
final_prompt_hidden_state_list = [None for _ in range(len(prompt_list))]
|
||||
|
||||
generated_answer_list = [s.lower() for s in tokenizer.batch_decode(outputs['sequences'], skip_special_tokens=True)]
|
||||
return generated_answer_list, logits_list, final_prompt_hidden_state_list
|
||||
|
||||
|
||||
def _apply_prompt_format_to_extracted_fields(
|
||||
structured_prompt_format, input_fields_list, regex_key_idx_list, output_fields_list=None):
|
||||
# Precompute all format options
|
||||
prompt = {}
|
||||
for key in set(regex_key_idx_list):
|
||||
prompt[key] = flatten(structured_prompt_format.solve(
|
||||
{'enumeration_length': key,
|
||||
'print_output_fields': True,
|
||||
'exclude_text_field_for_output_fields': output_fields_list is None})
|
||||
).replace('<|text|>', '{}')
|
||||
|
||||
# add empty default values if no output will be printed. It has to be a tuple to be able to concat with input_fields
|
||||
if output_fields_list is None:
|
||||
output_fields_list = [() for _ in input_fields_list]
|
||||
else:
|
||||
output_fields_list = [(output_field,) for output_field in output_fields_list]
|
||||
|
||||
formatted_inputs = []
|
||||
for input_fields, regex_key_idx, output_field in zip(input_fields_list, regex_key_idx_list, output_fields_list):
|
||||
tmp = apply_prompt_format(prompt[regex_key_idx], input_fields + output_field)
|
||||
formatted_inputs.append(tmp)
|
||||
|
||||
return formatted_inputs
|
||||
|
||||
|
||||
def _setup_formatted_demonstrations_with_definition(
|
||||
structured_prompt_format, demonstration_definition, demonstrations_outputs,
|
||||
original_to_current_multiple_choice_classes, demos_fields_list, demos_regex_key_idx_list):
|
||||
# 1. replace the variables in the demonstration definition. Used when the instruction mentions
|
||||
# multiple choice options, which need to change when the format changes
|
||||
demonstration_definition = demonstration_definition.format(
|
||||
**structured_prompt_format.find_all_formatted_field_values()
|
||||
)
|
||||
|
||||
demonstrations_outputs = [demo[0] if isinstance(demo, list) else demo for demo in demonstrations_outputs]
|
||||
if original_to_current_multiple_choice_classes:
|
||||
demonstrations_outputs = [original_to_current_multiple_choice_classes[d] for d in demonstrations_outputs]
|
||||
|
||||
all_demonstrations = _apply_prompt_format_to_extracted_fields(
|
||||
structured_prompt_format, demos_fields_list, demos_regex_key_idx_list, demonstrations_outputs)
|
||||
demonstration_string = demonstration_definition + "\n\n" + "\n\n".join(all_demonstrations)
|
||||
return demonstration_string
|
||||
|
||||
|
||||
def _setup_full_prompts_to_test_on(input_fields_list, regex_key_idx_list, selected_dataset_ids,
|
||||
demos_fields_list, demos_regex_key_idx_list, demonstrations_outputs,
|
||||
demonstration_definition,
|
||||
structured_prompt_format, original_to_current_multiple_choice_classes,
|
||||
interval_ids_to_test, n_shot):
|
||||
"""
|
||||
This function creates the full prompt string to be tested. This requires:
|
||||
|
||||
- Formatting the demonstrations with its definition, which may require
|
||||
replacing some variables referring to multiple choice options.
|
||||
- Apply prompt format to the desired set of examples to be tested (determined by interval_ids_to_test).
|
||||
"""
|
||||
demonstration_string = _setup_formatted_demonstrations_with_definition(
|
||||
structured_prompt_format, demonstration_definition, demonstrations_outputs,
|
||||
original_to_current_multiple_choice_classes, demos_fields_list, demos_regex_key_idx_list
|
||||
)
|
||||
|
||||
# filter to keep desired interval
|
||||
inputs = _apply_prompt_format_to_extracted_fields(
|
||||
structured_prompt_format,
|
||||
input_fields_list[interval_ids_to_test[0]:interval_ids_to_test[1]],
|
||||
regex_key_idx_list[interval_ids_to_test[0]:interval_ids_to_test[1]]
|
||||
)
|
||||
selected_dataset_ids = selected_dataset_ids[interval_ids_to_test[0]:interval_ids_to_test[1]]
|
||||
|
||||
full_prompt_string_list = []
|
||||
for input_element, idx in zip(inputs, selected_dataset_ids):
|
||||
full_prompt_string_list.append(input_element if n_shot == 0 else demonstration_string + "\n\n" + input_element)
|
||||
|
||||
return full_prompt_string_list, selected_dataset_ids
|
||||
|
||||
|
||||
def evaluate_prompt_format(
|
||||
args, dataset, input_fields_list, regex_key_idx_list, selected_dataset_ids,
|
||||
demos_fields_list, demos_regex_key_idx_list, demonstrations_outputs, demonstration_definition,
|
||||
structured_prompt_format, model, tokenizer, model_will_repeat_input,
|
||||
original_to_current_multiple_choice_classes, interval_ids_to_test=(None, None)):
|
||||
"""
|
||||
Function that evaluates a prompt format (i.e. node) on a given set of samples (interval_ids_to_test).
|
||||
If interval_ids_to_test is not provided, it defaults to evaluating the whole dataset.
|
||||
"""
|
||||
|
||||
# 1. set up input prompts including demonstrations
|
||||
input_prompt_string_list, selected_dataset_ids = _setup_full_prompts_to_test_on(
|
||||
input_fields_list, regex_key_idx_list, selected_dataset_ids,
|
||||
demos_fields_list, demos_regex_key_idx_list, demonstrations_outputs, demonstration_definition,
|
||||
structured_prompt_format, original_to_current_multiple_choice_classes, interval_ids_to_test, args.n_shot)
|
||||
|
||||
# 2. update the output values if needed, i.e. if the multiple choice classes now have different names
|
||||
assert all(len(dataset[idx]['output']) == 1 for idx in selected_dataset_ids)
|
||||
dataset_updated = copy.deepcopy(dataset)
|
||||
if original_to_current_multiple_choice_classes:
|
||||
for idx in range(len(dataset)):
|
||||
dataset_updated[idx]['output'][0] = original_to_current_multiple_choice_classes[dataset[idx]['output'][0]]
|
||||
output_classes = sorted(list(set([dataset_updated[idx]['output'][0] for idx in selected_dataset_ids])))
|
||||
|
||||
# 3. evaluate
|
||||
if args.evaluation_metric == 'probability_ranking':
|
||||
return solve_with_rank_based_scoring(
|
||||
dataset_updated, selected_dataset_ids, model, tokenizer, input_prompt_string_list, args.batch_size_llm)
|
||||
|
||||
elif args.evaluation_metric == 'exact_prefix_matching':
|
||||
logs = generate_text_with_metadata(
|
||||
args, input_prompt_string_list, model, tokenizer, model_will_repeat_input,
|
||||
dataset_updated, selected_dataset_ids, output_classes)
|
||||
return exact_prefix_matching_scoring(logs)
|
||||
|
||||
|
||||
def generate_text_with_metadata(args, input_prompt_string_list, model, tokenizer, model_will_repeat_input, dataset,
|
||||
selected_dataset_ids, output_classes):
|
||||
logs = []
|
||||
all_tokens_used = 0
|
||||
progress = tqdm(
|
||||
total=len(input_prompt_string_list),
|
||||
desc='API evaluation' if args.use_gpt3 else 'Local evaluation',
|
||||
unit='sample',
|
||||
leave=False,
|
||||
)
|
||||
effective_batch_size = max(args.batch_size_llm, args.api_concurrency) if args.use_gpt3 else args.batch_size_llm
|
||||
for batch_idx in range(math.ceil(len(input_prompt_string_list) / effective_batch_size)):
|
||||
batch_range = [batch_idx * effective_batch_size, (batch_idx + 1) * effective_batch_size] # [) range
|
||||
|
||||
full_prompt_string_list = input_prompt_string_list[batch_range[0]:batch_range[1]]
|
||||
if args.use_gpt3:
|
||||
request_results = [None] * len(full_prompt_string_list)
|
||||
with ThreadPoolExecutor(max_workers=args.api_concurrency) as executor:
|
||||
futures = {
|
||||
executor.submit(call_openai_api_with_retry, args, prompt, args.max_new_tokens): index
|
||||
for index, prompt in enumerate(full_prompt_string_list)
|
||||
}
|
||||
for future in as_completed(futures):
|
||||
index = futures[future]
|
||||
request_results[index] = future.result()
|
||||
progress.update(1)
|
||||
generation_list = [generation for generation, _ in request_results]
|
||||
all_tokens_used += sum(tokens_used for _, tokens_used in request_results)
|
||||
|
||||
score_list = [None for _ in range(len(generation_list))]
|
||||
final_prompt_hidden_state_list = [None for _ in range(len(generation_list))]
|
||||
else:
|
||||
generation_list, score_list, final_prompt_hidden_state_list = query_model_parallelized(
|
||||
model, tokenizer, full_prompt_string_list, max_tokens=args.max_new_tokens, top_p=1.0, temperature=1.0,
|
||||
)
|
||||
if model_will_repeat_input:
|
||||
generation_list = [generation[len(full_prompt_string):]
|
||||
for generation, full_prompt_string in zip(generation_list, full_prompt_string_list)]
|
||||
progress.update(len(generation_list))
|
||||
|
||||
selected_dataset_ids_list = [idx for idx in selected_dataset_ids[batch_range[0]:batch_range[1]]]
|
||||
assert len(generation_list) == len(selected_dataset_ids_list) == len(score_list) == len(
|
||||
final_prompt_hidden_state_list) == len(full_prompt_string_list)
|
||||
for generation, scores, idx, final_prompt_hidden_state, full_prompt_string in \
|
||||
zip(generation_list, score_list, selected_dataset_ids_list, final_prompt_hidden_state_list,
|
||||
full_prompt_string_list):
|
||||
expected_output = dataset[idx]['output'][0]
|
||||
# 'entry' and 'output_classes' are needed for score generations
|
||||
|
||||
current_log = {
|
||||
'entry': dataset[idx],
|
||||
'dataset_idx': idx,
|
||||
'generation': generation,
|
||||
'answer': expected_output,
|
||||
'output_classes': output_classes,
|
||||
'full_prompt_string': full_prompt_string,
|
||||
'eval_type': 'exact_prefix_matching',
|
||||
'scores': scores,
|
||||
}
|
||||
if PRINT_HIDDEN_STATE:
|
||||
current_log['final_prompt_hidden_state'] = final_prompt_hidden_state
|
||||
logs.append(current_log)
|
||||
|
||||
progress.close()
|
||||
print('Total tokens used:', all_tokens_used)
|
||||
return logs
|
||||
|
||||
|
||||
def match_robust_to_multiple_choice(generation, answer_to_compare):
|
||||
"""
|
||||
We return whether the generation matched with the expected answer.
|
||||
|
||||
This function assumes clean_text has already been run.
|
||||
"""
|
||||
# likewise, if the response says "article" and the right answer is "a"
|
||||
if not generation.startswith(answer_to_compare):
|
||||
return False
|
||||
|
||||
# if generation starts with answer and they are the same length, they are the same string
|
||||
if len(generation) == len(answer_to_compare):
|
||||
return True
|
||||
|
||||
# if the generation starts with the correct text, make sure the next char is not text or number
|
||||
# otherwise it might be just the first part of a random word (e.g. "a" with "article")
|
||||
# or if correct answer is ii, and all answers are i, ii, iii, iv, avoid being overly optimistic!
|
||||
return not generation[len(answer_to_compare)].isalpha() and not generation[len(answer_to_compare)].isdigit()
|
||||
|
||||
|
||||
def exact_prefix_matching_scoring(logs):
|
||||
accuracy = {
|
||||
'right': [],
|
||||
'wrong': [],
|
||||
'other': [],
|
||||
'total': 0
|
||||
}
|
||||
for entry in logs:
|
||||
clean_text = lambda x: x.strip(' .,()\n-><').lower()
|
||||
|
||||
right_answer = entry['entry']['output'][0]
|
||||
wrong_answers = [e for e in entry['output_classes'] if e != right_answer]
|
||||
|
||||
entry['right_answer_formatted'] = right_answer
|
||||
entry['wrong_answers_formatted'] = wrong_answers
|
||||
|
||||
right_answer = clean_text(right_answer)
|
||||
wrong_answers = [clean_text(e) for e in wrong_answers]
|
||||
generation = entry['generation']
|
||||
|
||||
clean_generation = clean_text(generation)
|
||||
is_right = match_robust_to_multiple_choice(clean_generation, right_answer)
|
||||
is_wrong = any(
|
||||
match_robust_to_multiple_choice(clean_generation, wrong_answer) for wrong_answer in wrong_answers)
|
||||
|
||||
accuracy['right'].append(is_right)
|
||||
accuracy['wrong'].append(is_wrong)
|
||||
accuracy['other'].append(not is_wrong and not is_right)
|
||||
accuracy['total'] += 1
|
||||
|
||||
if 'output_classes' in entry and len(entry['output_classes']) > 50:
|
||||
del entry['output_classes']
|
||||
|
||||
# not changing this since it's called from many classes
|
||||
return (sum(accuracy['right']) * 1.0 / max(accuracy['total'], 1),
|
||||
sum(accuracy['wrong']) * 1.0 / max(accuracy['total'], 1),
|
||||
accuracy['total']), (accuracy, logs)
|
||||
|
||||
|
||||
def solve_with_rank_based_scoring(
|
||||
dataset, selected_dataset_ids, model, tokenizer, input_prompt_string_list, batch_size_llm):
|
||||
import psutil
|
||||
output_classes = sorted(list(set([dataset[idx]['output'][0] for idx in selected_dataset_ids])))
|
||||
assert len(output_classes) < 100
|
||||
assert tokenizer is not None and model is not None
|
||||
|
||||
# if all output values are only one token, then we can just look at the output probabilities
|
||||
# instead of computing perplexity for all possible prompt+outputs!
|
||||
# also if all output values share the same prefix. E.g. ['0', '1'] tokenizes to [[1, 29871, 29900], [1, 29871, 29896]]
|
||||
# the first token id is always '1', so we ignore it
|
||||
output_classes_tokens = [t for t in tokenizer(output_classes, return_token_type_ids=False)['input_ids']]
|
||||
single_token_classes = all([len(t) == 2 for t in output_classes_tokens])
|
||||
all_classes_share_common_prefix = len(set([tuple(t[:-1]) for t in output_classes_tokens])) == 1
|
||||
|
||||
accuracy = {
|
||||
'right': [],
|
||||
'wrong': [],
|
||||
'other': [],
|
||||
'total': 0
|
||||
}
|
||||
logs = []
|
||||
|
||||
if single_token_classes or all_classes_share_common_prefix:
|
||||
# batching happens across inputs
|
||||
for batch_idx in range(math.ceil(len(input_prompt_string_list) / batch_size_llm)):
|
||||
print("Memory usage:", psutil.Process(os.getpid()).memory_info().rss / 1024 ** 2)
|
||||
|
||||
batch_range = [batch_idx * batch_size_llm, (batch_idx + 1) * batch_size_llm] # [) range
|
||||
|
||||
full_prompt_string_list = input_prompt_string_list[batch_range[0]:batch_range[1]]
|
||||
generation_list = get_ranking_based_generation_single_token_output_classes(
|
||||
full_prompt_string_list, output_classes, tokenizer, model)
|
||||
|
||||
selected_dataset_ids_list = [idx for idx in selected_dataset_ids[batch_range[0]:batch_range[1]]]
|
||||
assert len(generation_list) == len(selected_dataset_ids_list), f"{len(generation_list)} generations, {len(selected_dataset_ids_list)} selected ids"
|
||||
assert len(generation_list) == len(full_prompt_string_list)
|
||||
for generation, idx, full_prompt_string in zip(generation_list, selected_dataset_ids_list, full_prompt_string_list):
|
||||
expected_output = dataset[idx]['output'][0]
|
||||
|
||||
assert expected_output in output_classes, f"expected_output={expected_output}, output_classes={output_classes}"
|
||||
|
||||
accuracy['right'].append((generation == expected_output))
|
||||
accuracy['wrong'].append((generation != expected_output and generation in output_classes))
|
||||
accuracy['other'].append((generation not in output_classes))
|
||||
accuracy['total'] += 1
|
||||
logs.append(
|
||||
{
|
||||
'entry': dataset[idx],
|
||||
'dataset_idx': idx,
|
||||
'generation': generation,
|
||||
'answer': expected_output,
|
||||
'output_classes': output_classes,
|
||||
'full_prompt_string': full_prompt_string,
|
||||
'eval_type': 'ranking_single_token',
|
||||
'scores': None,
|
||||
}
|
||||
)
|
||||
|
||||
else:
|
||||
# batching happens inside each input, since we need to do inference for each prompt+possible_output
|
||||
for i in range(len(input_prompt_string_list)):
|
||||
idx = selected_dataset_ids[i]
|
||||
full_prompt_string = input_prompt_string_list[i]
|
||||
|
||||
generation = get_ranking_based_generation_multiple_token_output_classes(
|
||||
full_prompt_string, output_classes, tokenizer, model, batch_size_llm,
|
||||
)
|
||||
expected_output = dataset[idx]['output'][0]
|
||||
assert expected_output in output_classes, f"expected_output={expected_output}, output_classes={output_classes}"
|
||||
|
||||
accuracy['right'].append((generation == expected_output))
|
||||
accuracy['wrong'].append((generation != expected_output and generation in output_classes))
|
||||
accuracy['other'].append((generation not in output_classes))
|
||||
accuracy['total'] += 1
|
||||
|
||||
logs.append(
|
||||
{
|
||||
'entry': dataset[idx],
|
||||
'dataset_idx': idx,
|
||||
'generation': generation,
|
||||
'answer': expected_output,
|
||||
'output_classes': output_classes,
|
||||
'full_prompt_string': full_prompt_string,
|
||||
'eval_type': 'ranking_multiple_token',
|
||||
'scores': None,
|
||||
}
|
||||
)
|
||||
|
||||
return (sum(accuracy['right']) * 1.0 / max(accuracy['total'], 1),
|
||||
sum(accuracy['wrong']) * 1.0 / max(accuracy['total'], 1),
|
||||
accuracy['total']), (accuracy, logs)
|
||||
|
||||
|
||||
def get_ranking_based_generation_single_token_output_classes(prompts, output_classes, tokenizer, model):
|
||||
import torch
|
||||
top_p = 1.0
|
||||
temperature = 1.0
|
||||
|
||||
# if all output values are only one token, then we can just look at the output probabilities!
|
||||
# also if all output values share the same prefix. E.g. ['0', '1'] tokenizes to [[1, 29871, 29900], [1, 29871, 29896]]
|
||||
# the first token id is always '1', so we ignore it
|
||||
output_classes_tokens = [t for t in tokenizer(output_classes, return_token_type_ids=False)['input_ids']]
|
||||
all_classes_share_common_prefix = len(set([tuple(t[:-1]) for t in output_classes_tokens])) == 1
|
||||
|
||||
tokenized_inputs_list = tokenizer(prompts, return_tensors="pt", padding=True, return_token_type_ids=False)[
|
||||
'input_ids'].tolist()
|
||||
if all_classes_share_common_prefix:
|
||||
for i in range(len(tokenized_inputs_list)):
|
||||
# if the tokenized element is [1, 29871, 29900], get [29871]
|
||||
tokenized_inputs_list[i] += output_classes_tokens[0][1:-1]
|
||||
tokenized_inputs = torch.tensor(tokenized_inputs_list).to('cuda')
|
||||
|
||||
with torch.no_grad():
|
||||
outputs = model.generate(input_ids=tokenized_inputs,
|
||||
top_p=top_p, temperature=temperature, max_new_tokens=1,
|
||||
return_dict_in_generate=True, output_scores=True)
|
||||
|
||||
scores = outputs["scores"][0] # first dimension = 1 since we only generate one token
|
||||
generations = []
|
||||
for i in range(len(prompts)):
|
||||
all_logits = scores[i, :].squeeze().tolist()
|
||||
all_logits_sorted = sorted([(all_logits[t[-1]], i) for i, t in enumerate(output_classes_tokens)], reverse=True)
|
||||
generations.append(output_classes[all_logits_sorted[0][1]])
|
||||
return generations
|
||||
|
||||
|
||||
def get_ranking_based_generation_multiple_token_output_classes(prompt, output_classes, tokenizer, model,
|
||||
batch_size_llm):
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
output_classes_tokens = [t for t in tokenizer(output_classes, return_token_type_ids=False)['input_ids']]
|
||||
|
||||
prompts = [prompt + class_seq for class_seq in output_classes]
|
||||
all_logits_list, all_tokens_list = [], []
|
||||
for batch_idx in range(math.ceil(len(prompts) / batch_size_llm)):
|
||||
batch_range = [batch_idx * batch_size_llm, (batch_idx + 1) * batch_size_llm] # [) range
|
||||
all_logits, all_tokens = _get_input_logits_and_tokens(prompts[batch_range[0]:batch_range[1]], tokenizer, model)
|
||||
all_logits_list.extend(all_logits)
|
||||
all_tokens_list.extend(all_tokens)
|
||||
|
||||
n_classes = len(output_classes)
|
||||
class_logprobs = []
|
||||
for class_index in range(n_classes):
|
||||
class_logits = all_logits_list[class_index]
|
||||
|
||||
# the lengths of each class sequence in tokens
|
||||
target_token_length = (len(output_classes_tokens[class_index]))
|
||||
# we only need the logits for the end sequence
|
||||
tokens = all_tokens_list[class_index]
|
||||
# we have to go back by one because we don't care about the logits for the predicted token
|
||||
sequence_logits = class_logits[-target_token_length - 1: -1]
|
||||
sequence_tokens = tokens[-target_token_length:]
|
||||
# we take a log_softmax over all token logits for each position in the class sequence to
|
||||
# get log probabilities, and then sum the logprobs for the tokens actually chosen
|
||||
logprobs = F.log_softmax(sequence_logits, dim=-1).to('cpu')
|
||||
class_logprob = sum(
|
||||
[logprobs[i, token] for i, token in enumerate(sequence_tokens)]
|
||||
)
|
||||
class_logprobs.append(class_logprob.item())
|
||||
|
||||
return output_classes[torch.tensor(class_logprobs).argmax(dim=-1).item()]
|
||||
|
||||
|
||||
def _get_input_logits_and_tokens(inputs, tokenizer, model):
|
||||
import torch
|
||||
tokenized_inputs = tokenizer(inputs, return_tensors="pt", padding=True, return_token_type_ids=False).to('cuda')
|
||||
with torch.no_grad():
|
||||
outputs = model(**tokenized_inputs)
|
||||
logits = outputs["logits"].detach().to(device="cpu", dtype=torch.float32)
|
||||
return logits, tokenized_inputs["input_ids"]
|
||||
Reference in New Issue
Block a user