Repository navigation
Expand file tree
/
Copy pathgenerate_frames.py
More file actions
102 lines (81 loc) · 6.33 KB
/
Copy pathgenerate_frames.py
File metadata and controls
102 lines (81 loc) · 6.33 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
import argparse
from utils import define_text_input, select_texts_4_eval
from RAG import generate_response
from evaluation import evaluate
import json
import os
def parse_args():
parser = argparse.ArgumentParser()
# Arguments for reading in and writing out (parsed) frames
parser.add_argument("--frame_dir", type=str,
help="Directory containing the frames from FrameNet. Default is 'data/framenet/frames'"
"If frame embeddings do not exist yet in the current directory, please provide the path to the directory with FrameNet frames (in xml, or json if already parsed) so that the embeddings can be created.",
default="data/framenet/frames")
parser.add_argument("--frame_output_filepath", type=str,
help="Path to save parsed frames. If none is given, the system tries to save parsed frames in the same directory as the raw xml frames.",
default=None)
# Arguments for reading in and writing out (parsed) annotated texts
parser.add_argument("--annotated_texts_dir", type=str,
help="Directory containing the annotated texts from FrameNet. Default is 'data/framenet/annotated_texts'. "
"Needed for few-shot mode, and evaluation against annotations.",
default="data/framenet/fulltexts")
# Arguments for reading input data
parser.add_argument("--input_text", type=str, help="Text to generate frames from.", default=None)
parser.add_argument("--input_filepath", type=str, help="Path to input texts. If csv, provide name of column with texts, and optionally provide separator (--sep, default ',')", default=None)
parser.add_argument("--text_column_name", type=str, help="Column name of column with input texts in provided csv.", default=None)
parser.add_argument("--sep", type=str, help="Separator used in input csv file. Default is ','.", default=',')
# Arguments for writing out output data
parser.add_argument("--output_filepath", type=str, help="Path to save output of frame extraction. If none is given, the system tries to save response in current working directory.", default=None)
# Arguments for retrieval set-up
parser.add_argument("--top_k", type=int, help="Number of most relevant frames to retrieve from the vector store. Default: 29", default=29)
parser.add_argument("--search_type", type=str, help="FAISS retrieval strategy to select the most relevant frames. Default: similarity", default="similarity")
parser.add_argument("--embedding_model_name", type=str, help="Name of the HuggingFace embedding model to use for frame retrieval. Default: BAAI/bge-m3", default="BAAI/bge-m3")
parser.add_argument("--retrieval_only", dest='retrieval_only', action='store_true', help='Disable generation (LLM prompting). Retrieve candidate frames based on similarity with input text, but do not select and align them to text using an LLM.')
# Arguments for LLM set-up
parser.add_argument("--model_name", type=str, help="Name of OpenAI model to be used for frame extraction.", default="gpt-5.4")
parser.add_argument("--api_key_name", type=str, help="Name of user's API key to call model.", default="OPENAI_API_KEY")
parser.add_argument("--nebula", dest='nebula', action="store_true", help="Whether to use Nebula models for frame extraction. If not set, nebula models are not used.")
parser.add_argument("--model_provider", type=str, help="Name of model provider in case of using nebula (LangChain cannot infer the model provider if using nebula)", default=None)
parser.add_argument("--zero_shot", dest="zero_shot", action="store_true", help="Change mode from few-shot (with examples in LLM prompt) to zero-shot (without examples in LLM prompt).")
# Arguments for executing evaluation with FrameNet annotated texts
parser.add_argument("--evaluate", action="store_true", help="Evaluate model using FrameNet annotated texts.")
parser.add_argument("--generated_frames_filepath", type=str, help='If frames have already been generated, provide path to the saved generated frames.', default=None)
parser.add_argument("--proportion_eval_texts", type=float, help="Proportion of FrameNet texts to use for evaluation. Default is 0.1 (=10%).", default=0.1)
parser.add_argument("--min_frames_per_eval_text", type=int, help="Minimum number of frames annotated in a text for it to be included in evaluation. Default is 2.", default=2)
parser.add_argument("--match_type", type=str, default="exact", help="Type of matching between model-generated frames and human-annotated frames. Choose between exact, semantic, and graph. Default: exact.")
parser.add_argument("--output_eval_dir", type=str, help="Directory to save evaluation output. If none is given, the system tries to save output files in current working directory.", default=None)
return parser.parse_args()
def main():
args = parse_args()
# Handle where to save output
if not args.output_filepath:
output_filepath = os.getcwd() + f'/output_{args.model_name}.json'
args.output_filepath = output_filepath
# Only generation mode
if not args.evaluate:
# Define input text
# if no text string is given in the command line, check if a filepath is given, and read text(s) from there
if args.input_text:
texts = [args.input_text]
else:
texts = define_text_input(args)
# Set up RAG and get RAG responses
_ = generate_response(texts=texts, args=args)
# Evaluation mode
else:
# If evaluating, input texts are a portion of annotated FrameNet texts
annotated_texts = select_texts_4_eval(args)
# Generate frames from selected texts
generated_frames = None
if not args.generated_frames_filepath:
texts = [text["text"] for text in annotated_texts["texts"]]
generated_frames = generate_response(texts=texts, args=args)
else:
# Try to load frames if already generated
with open(args.generated_frames_filepath, "r") as f:
generated_frames = json.load(f)
print(f'Loaded generated frames from {args.generated_frames_filepath}.')
# Evaluate generated frames against annotations
evaluate(args, annotated_texts, generated_frames)
if __name__ == "__main__":
main()