diff --git a/perturb_EB.py b/perturb_EB.py index 7654a8b..80ad039 100644 --- a/perturb_EB.py +++ b/perturb_EB.py @@ -11,6 +11,8 @@ import jsonlines from tqdm import tqdm from transformers import PegasusForConditionalGeneration, PegasusTokenizer +import torch +import numpy as np gold_path = 'entailment_bank/data/public_dataset/entailment_trees_emnlp2021_data_v2/dataset/task_1/test.jsonl' more_path = 'entailment_bank/data/public_dataset/entailment_trees_emnlp2021_data_v2/dataset/task_2/test.jsonl' @@ -73,7 +75,7 @@ def reconstruct_proof(steps, sentences): proof += '; ' return proof -def repeat_steps(in_steps, in_sentences): +def repeat_steps(in_steps, in_sentences, *args): steps = deepcopy(in_steps) sentences = deepcopy(in_sentences) int_idxs = [s for s in range(len(steps)) if 'int' in steps[s]['child']] @@ -87,7 +89,7 @@ def repeat_steps(in_steps, in_sentences): steps.insert(idx + 1, {'parents': [key], 'child': repeated_node}) return steps, sentences -def delete_steps(in_steps, in_sentences): +def delete_steps(in_steps, in_sentences, *args): steps = deepcopy(in_steps) sentences = deepcopy(in_sentences) int_idxs = [s for s in range(len(steps)) if 'int' in steps[s]['child']] @@ -102,7 +104,7 @@ def delete_steps(in_steps, in_sentences): step['parents'] = [p for p in step['parents'] if p != del_node] return steps, sentences -def swapped_steps(in_steps, in_sentences): +def swapped_steps(in_steps, in_sentences, *args): steps = deepcopy(in_steps) sentences = deepcopy(in_sentences) int_idxs = [s for s in range(len(steps)) if 'int' in steps[s]['child']] @@ -121,7 +123,7 @@ def swapped_steps(in_steps, in_sentences): step['parents'] = [p for p in step['parents'] if p != swap_node] return steps, sentences -def negate_step(in_steps, in_sentences): +def negate_step(in_steps, in_sentences, *args): steps = deepcopy(in_steps) sentences = deepcopy(in_sentences) int_idxs = [s for s in range(len(steps)) if 'int' in steps[s]['child']] @@ -145,7 +147,7 @@ def hallucinate_step(in_steps, in_sentences, extra_sentences): print(sentences[hallucinate_node]) return steps, sentences -def paraphrase_steps(in_steps, in_sentences): +def paraphrase_steps(in_steps, in_sentences, *args): steps = deepcopy(in_steps) sentences = deepcopy(in_sentences) int_idxs = [s for s in range(len(steps)) if 'int' in steps[s]['child']] @@ -209,7 +211,8 @@ def redundant_steps(in_steps, in_sentences, extra_sentences): perturbed = False custom_unperturbed_ids[perturb_type + "_test.jsonl"].append(id) else: perturbed = True - tree_entry.append({'perturbed': perturbed, 'perturbations': perturb_type, 'steps':{'original': steps, 'perturbed': perturbed_steps}, 'sentences':{'original': sentences, 'perturbed': perturbed_sentences}, 'written':{'original': original_written_steps, 'perturbed': written_steps}, 'question': question, 'answer': answer}) + # tree_entry.append({'perturbed': perturbed, 'perturbations': perturb_type, 'steps':{'original': steps, 'perturbed': perturbed_steps}, 'sentences':{'original': sentences, 'perturbed': perturbed_sentences}, 'written':{'original': original_written_steps, 'perturbed': written_steps}, 'question': question, 'answer': answer}) + tree_entry.append({'perturbed': perturbed, 'perturbations': perturb_type, 'steps':{'original': steps, 'perturbed': perturbed_steps}, 'sentences':{'original': sentences, 'perturbed': perturbed_sentences}, 'question': question, 'answer': answer}) with jsonlines.open(os.path.join(tree_dest_path, fname), 'w') as writer: writer.write_all(tree_entry)