flashqda 0.0.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- flashqda/__init__.py +0 -0
- flashqda/analyze.py +555 -0
- flashqda/classify.py +379 -0
- flashqda/general.py +41 -0
- flashqda/preprocess.py +108 -0
- flashqda/screen.py +126 -0
- flashqda-0.0.1.dist-info/LICENSE +674 -0
- flashqda-0.0.1.dist-info/METADATA +23 -0
- flashqda-0.0.1.dist-info/RECORD +11 -0
- flashqda-0.0.1.dist-info/WHEEL +5 -0
- flashqda-0.0.1.dist-info/top_level.txt +1 -0
flashqda/__init__.py
ADDED
|
File without changes
|
flashqda/analyze.py
ADDED
|
@@ -0,0 +1,555 @@
|
|
|
1
|
+
from .general import save_to_csv, read_csv_file
|
|
2
|
+
import re, csv, os, threading, string, requests, random, time
|
|
3
|
+
from openai import OpenAI
|
|
4
|
+
from tqdm import tqdm
|
|
5
|
+
from pandas.errors import EmptyDataError
|
|
6
|
+
import pandas as pd
|
|
7
|
+
import json
|
|
8
|
+
import numpy as np
|
|
9
|
+
from datetime import datetime
|
|
10
|
+
|
|
11
|
+
classification_analyses = ["tenses", "relationships_classify"]
|
|
12
|
+
extraction_analyses = ["relationships_extract"]
|
|
13
|
+
#keyword_analyses = ["uncertainty", "opinion", "association"]
|
|
14
|
+
all_analyses = classification_analyses + extraction_analyses #+ keyword_analyses
|
|
15
|
+
|
|
16
|
+
def api_call_wrapper(system_prompt, prompt, result_container, event):
|
|
17
|
+
|
|
18
|
+
"""Try prompting OpenAI."""
|
|
19
|
+
|
|
20
|
+
client = OpenAI(
|
|
21
|
+
api_key=os.environ['OPENAI_API_KEY'], # this is also the default, it can be omitted
|
|
22
|
+
)
|
|
23
|
+
|
|
24
|
+
try:
|
|
25
|
+
response = client.chat.completions.create(
|
|
26
|
+
model="gpt-4o",
|
|
27
|
+
messages=[
|
|
28
|
+
{"role": "system", "content": system_prompt},
|
|
29
|
+
{"role": "user", "content": prompt}
|
|
30
|
+
],
|
|
31
|
+
response_format = {"type": "json_object"},
|
|
32
|
+
temperature = 0.5,
|
|
33
|
+
timeout=10
|
|
34
|
+
)
|
|
35
|
+
result_container[0] = response
|
|
36
|
+
event.set()
|
|
37
|
+
except Exception as e:
|
|
38
|
+
result_container[1] = e
|
|
39
|
+
event.set()
|
|
40
|
+
|
|
41
|
+
def send_to_openai(system_prompt, prompt, max_retries=3, manual_timeout=15):
|
|
42
|
+
|
|
43
|
+
"""Handle time-outs by OpenAI."""
|
|
44
|
+
|
|
45
|
+
base_delay = 1 # Base delay in seconds
|
|
46
|
+
max_delay = 480 # Max delay in seconds
|
|
47
|
+
retries = 0
|
|
48
|
+
|
|
49
|
+
while retries < max_retries:
|
|
50
|
+
#print("Making API call...")
|
|
51
|
+
|
|
52
|
+
result_container = [None, None]
|
|
53
|
+
done_event = threading.Event()
|
|
54
|
+
|
|
55
|
+
api_thread = threading.Thread(target=api_call_wrapper, args=(system_prompt, prompt, result_container, done_event))
|
|
56
|
+
api_thread.start()
|
|
57
|
+
api_thread.join(timeout=manual_timeout)
|
|
58
|
+
|
|
59
|
+
if done_event.is_set():
|
|
60
|
+
response, error = result_container
|
|
61
|
+
if response:
|
|
62
|
+
clean_response = response.choices[0].message.content.strip().lower()
|
|
63
|
+
return clean_response
|
|
64
|
+
elif isinstance(error, requests.exceptions.Timeout):
|
|
65
|
+
retries += 1
|
|
66
|
+
delay = base_delay * (2 ** retries) # Exponential backoff
|
|
67
|
+
delay = min(max_delay, delay) + random.uniform(0, 0.1*base_delay) # Add some jitter
|
|
68
|
+
time.sleep(delay)
|
|
69
|
+
print("Timeout exception")
|
|
70
|
+
else:
|
|
71
|
+
raise error # Some other exception occurred
|
|
72
|
+
else:
|
|
73
|
+
# The manual timeout was triggered
|
|
74
|
+
retries += 1
|
|
75
|
+
print(f"Manual timeout triggered after {manual_timeout} seconds")
|
|
76
|
+
|
|
77
|
+
return None # Default return value if all retries fail
|
|
78
|
+
|
|
79
|
+
def contains_whole_words(sentence, analysis_type, terms_to_check):
|
|
80
|
+
|
|
81
|
+
"""Check whether the sentence contains any words in a specified list and report the count."""
|
|
82
|
+
|
|
83
|
+
#terms_to_check = []
|
|
84
|
+
|
|
85
|
+
#uncertainty_terms = ['can', 'could', 'may', 'might', 'would', 'unlikely', 'possibly', 'probably', 'likely', 'often', 'sometimes', 'potential', 'potentially', 'presumably','perhaps','seem', 'seems', 'appear', 'appears', 'apparently', 'suggest', 'suggests', 'imply', 'implies', 'suppose', 'supposedly', 'hypothesis', 'hypothesize', 'hypothesized', 'hypothesise', 'hypothesised', 'postulate', 'we expect', 'offer', 'offers', 'enables']
|
|
86
|
+
#opinion_terms = ['should', 'must', 'ought', 'is imperative', 'are imperative', 'be imperative', 'is crucial', 'are crucial', 'be crucial', 'is essential', 'are essential', 'be essenetial', 'is needed', 'are needed', 'be needed', 'needs to', 'need to', 'is necessary', 'are necessary', 'be necessary', 'is required', 'are required', 'be required', 'is critical', 'are critical', 'be critical', 'argue', 'argued', 'important to', 'promising']
|
|
87
|
+
#association_terms = ['correlation', 'correlated with', 'correlated to', 'association', 'associated with', 'associated to', 'related with', 'related to', 'relates to', 'linkage', 'linked with', 'linked to']
|
|
88
|
+
|
|
89
|
+
#if analysis_type == "uncertainty":
|
|
90
|
+
# terms_to_check = uncertainty_terms
|
|
91
|
+
|
|
92
|
+
#elif analysis_type == "opinion":
|
|
93
|
+
# terms_to_check = opinion_terms
|
|
94
|
+
|
|
95
|
+
#elif analysis_type == "association":
|
|
96
|
+
# terms_to_check = association_terms
|
|
97
|
+
|
|
98
|
+
sentence_text = sentence.lower()
|
|
99
|
+
count = 0
|
|
100
|
+
|
|
101
|
+
for word in terms_to_check:
|
|
102
|
+
matches = re.findall(rf'\b{re.escape(word)}\b', sentence_text)
|
|
103
|
+
count += len(matches)
|
|
104
|
+
|
|
105
|
+
return count
|
|
106
|
+
|
|
107
|
+
def get_decision_for_sentence(sentence, context_window, analysis_type, terms_to_check):
|
|
108
|
+
|
|
109
|
+
"""Get a decision for the current sentence, based on the current type of analysis."""
|
|
110
|
+
|
|
111
|
+
if analysis_type in classification_analyses:
|
|
112
|
+
with open(f'./FlashQDA/Prompts/{analysis_type}.txt', 'r') as file:
|
|
113
|
+
system_prompt = "You are a helpful assistant that classifies sentences. Follow the user's instructions carefully. Respond using JSON."
|
|
114
|
+
prompt = file.read()
|
|
115
|
+
prompt = prompt.format(sentence=sentence, context_window=context_window)
|
|
116
|
+
decision = send_to_openai(system_prompt, prompt)
|
|
117
|
+
elif analysis_type in extraction_analyses:
|
|
118
|
+
with open(f'./FlashQDA/Prompts/{analysis_type}.txt', 'r') as file:
|
|
119
|
+
system_prompt = "You are a helpful assistant that extracts and lists causal relationships from sentences. Follow the user's instructions carefully. Respond using JSON."
|
|
120
|
+
prompt = file.read()
|
|
121
|
+
prompt = prompt.format(sentence=sentence, context_window=context_window)
|
|
122
|
+
decision = send_to_openai(system_prompt, prompt)
|
|
123
|
+
elif analysis_type == 'list_of_terms':
|
|
124
|
+
decision = contains_whole_words(sentence, analysis_type, terms_to_check)
|
|
125
|
+
#elif analysis_type == 'opinion':
|
|
126
|
+
# decision = contains_whole_words(sentence, analysis_type, terms_to_check)
|
|
127
|
+
#elif analysis_type == 'association':
|
|
128
|
+
# decision = contains_whole_words(sentence, analysis_type, terms_to_check)
|
|
129
|
+
return decision
|
|
130
|
+
|
|
131
|
+
def count_decisions_for_sentence(analysis_type, decisions):
|
|
132
|
+
|
|
133
|
+
"""For the current sentence and the current type of analysis, count the number of decisions that match the specified criteria."""
|
|
134
|
+
|
|
135
|
+
if analysis_type == "tenses":
|
|
136
|
+
decisions_count = "{},{},{},{},{},{},{},{},{},{},{},{}".format(
|
|
137
|
+
decisions.count('simple past'),
|
|
138
|
+
decisions.count('past perfect'),
|
|
139
|
+
decisions.count('past continuous'),
|
|
140
|
+
decisions.count('past perfect continuous'),
|
|
141
|
+
decisions.count('simple present') + decisions.count('present simple'),
|
|
142
|
+
decisions.count('present continuous'),
|
|
143
|
+
decisions.count('present perfect'),
|
|
144
|
+
decisions.count('present perfect continuous'),
|
|
145
|
+
decisions.count('simple future'),
|
|
146
|
+
decisions.count('future perfect'),
|
|
147
|
+
decisions.count('future continuous'),
|
|
148
|
+
decisions.count('future perfect continuous'))
|
|
149
|
+
|
|
150
|
+
elif analysis_type == "relationships_classify":
|
|
151
|
+
decisions_count = "{},{},{}".format(
|
|
152
|
+
decisions.count('causal'),
|
|
153
|
+
decisions.count('correlational'),
|
|
154
|
+
decisions.count('none'))
|
|
155
|
+
|
|
156
|
+
elif analysis_type == "relationships_extract":
|
|
157
|
+
decisions_count = decisions
|
|
158
|
+
|
|
159
|
+
elif analysis_type == "list_of_terms":
|
|
160
|
+
decisions_count = decisions[0]
|
|
161
|
+
|
|
162
|
+
return decisions_count
|
|
163
|
+
|
|
164
|
+
def calculate_subscore_for_sentence(analysis_type, decisions):
|
|
165
|
+
|
|
166
|
+
"""For the current sentence and the current type of analysis, calculate a subscore."""
|
|
167
|
+
|
|
168
|
+
if len(decisions) > 0:
|
|
169
|
+
if analysis_type == "tenses":
|
|
170
|
+
subscore = (decisions.count('simple present') +
|
|
171
|
+
decisions.count('present simple') +
|
|
172
|
+
decisions.count('present continuous')
|
|
173
|
+
) / len(decisions)
|
|
174
|
+
|
|
175
|
+
elif analysis_type == "relationships_classify":
|
|
176
|
+
subscore = decisions.count('causal') / len(decisions)
|
|
177
|
+
|
|
178
|
+
elif analysis_type == "relationships_extract":
|
|
179
|
+
subscore = None
|
|
180
|
+
|
|
181
|
+
elif analysis_type == "list_of_terms":
|
|
182
|
+
subscore = decisions[0] if decisions[0] <= 1 else 1
|
|
183
|
+
else:
|
|
184
|
+
subscore = 0
|
|
185
|
+
|
|
186
|
+
return subscore
|
|
187
|
+
|
|
188
|
+
def analyze_the_current_sentence(sentence, context_window, analysis_type, terms_to_check, query_count):
|
|
189
|
+
|
|
190
|
+
"""Analyze the current sentence based on the specified analysis type."""
|
|
191
|
+
|
|
192
|
+
decisions = []
|
|
193
|
+
|
|
194
|
+
# Classify the current sentence by tense or relationship type (causal, correlational, none)
|
|
195
|
+
if analysis_type in classification_analyses:
|
|
196
|
+
type = "tenses" if analysis_type == "tenses" else "relationships"
|
|
197
|
+
for _ in range(query_count):
|
|
198
|
+
decision = get_decision_for_sentence(sentence["sentence"], context_window, analysis_type, terms_to_check)
|
|
199
|
+
data = json.loads(decision)
|
|
200
|
+
for key, value in data[type].items():
|
|
201
|
+
decisions.append(value)
|
|
202
|
+
|
|
203
|
+
# Extract causal relationships from the current sentence
|
|
204
|
+
elif analysis_type in extraction_analyses:
|
|
205
|
+
decision = get_decision_for_sentence(sentence["sentence"], context_window, analysis_type, terms_to_check)
|
|
206
|
+
decisions = decision
|
|
207
|
+
|
|
208
|
+
# Detect keywords in the current sentence
|
|
209
|
+
elif analysis_type == 'list_of_terms':
|
|
210
|
+
decision = get_decision_for_sentence(sentence["sentence"], context_window, analysis_type, terms_to_check)
|
|
211
|
+
decisions.append(decision)
|
|
212
|
+
|
|
213
|
+
decisions_count = count_decisions_for_sentence(analysis_type, decisions)
|
|
214
|
+
subscore = calculate_subscore_for_sentence(analysis_type, decisions)
|
|
215
|
+
|
|
216
|
+
return decisions_count, subscore
|
|
217
|
+
|
|
218
|
+
def initialize_files_analysis(save_name, analysis_type, terms_to_check):
|
|
219
|
+
base_path = os.path.join('./Results', save_name)
|
|
220
|
+
results_file = os.path.join(base_path, f'{save_name}_sentences.csv')
|
|
221
|
+
if analysis_type == "list_of_terms":
|
|
222
|
+
temp_file = os.path.join(base_path, f'{save_name}_{analysis_type}_{terms_to_check[0]}_temp.csv')
|
|
223
|
+
log_file = os.path.join(base_path, f'{save_name}_{analysis_type}_{terms_to_check[0]}_log.csv')
|
|
224
|
+
else:
|
|
225
|
+
temp_file = os.path.join(base_path, f'{save_name}_{analysis_type}_temp.csv')
|
|
226
|
+
log_file = os.path.join(base_path, f'{save_name}_{analysis_type}_log.csv')
|
|
227
|
+
|
|
228
|
+
os.makedirs(base_path, exist_ok=True)
|
|
229
|
+
|
|
230
|
+
# Ensure results file exists or create an empty one
|
|
231
|
+
if not os.path.exists(results_file):
|
|
232
|
+
pd.DataFrame().to_csv(results_file, index=False)
|
|
233
|
+
|
|
234
|
+
# Ensure log file exists or create with default values
|
|
235
|
+
if not os.path.exists(log_file):
|
|
236
|
+
with open(log_file, 'w') as log:
|
|
237
|
+
writer = csv.writer(log)
|
|
238
|
+
writer.writerow(['start_time', 'end_time', 'last_processed_document', 'last_processed_sentence'])
|
|
239
|
+
writer.writerow([datetime.now().strftime('%Y-%m-%d %H:%M:%S'), 0, 1, -1])
|
|
240
|
+
|
|
241
|
+
# Initialize temp file with empty DataFrame if it doesn't exist
|
|
242
|
+
if not os.path.exists(temp_file):
|
|
243
|
+
pd.DataFrame(columns=["document_id", "filename", "sentence_id", "sentence"]).to_csv(temp_file, index=False)
|
|
244
|
+
|
|
245
|
+
# Read existing results from temp file if it exists, otherwise initialize as empty list
|
|
246
|
+
try:
|
|
247
|
+
existing_results = pd.read_csv(temp_file).to_dict('records')
|
|
248
|
+
except EmptyDataError:
|
|
249
|
+
existing_results = []
|
|
250
|
+
|
|
251
|
+
return temp_file, log_file, existing_results
|
|
252
|
+
|
|
253
|
+
def get_start_ids(log_file):
|
|
254
|
+
try:
|
|
255
|
+
with open(log_file, 'r') as log:
|
|
256
|
+
reader = csv.reader(log)
|
|
257
|
+
next(reader) # Skip header
|
|
258
|
+
row = next(reader)
|
|
259
|
+
start_time = row[0]
|
|
260
|
+
end_time = row[1]
|
|
261
|
+
start_document_id, start_sentence_id = map(int, row[2:])
|
|
262
|
+
except EmptyDataError:
|
|
263
|
+
start_time, end_time, start_document_id, start_sentence_id = 0, 0, 1, -1 # Default values
|
|
264
|
+
|
|
265
|
+
return start_time, end_time, start_document_id, start_sentence_id
|
|
266
|
+
|
|
267
|
+
def update_log(log_file, start_time, end_time, document_id, sentence_id):
|
|
268
|
+
with open(log_file, 'w') as log:
|
|
269
|
+
writer = csv.writer(log)
|
|
270
|
+
writer.writerow(['start_time', 'end_time', 'last_processed_document', 'last_processed_sentence'])
|
|
271
|
+
writer.writerow([start_time, end_time, document_id, sentence_id])
|
|
272
|
+
|
|
273
|
+
def append_to_context(context_window, sentence, context_length):
|
|
274
|
+
if context_length > 0:
|
|
275
|
+
if len(context_window) == context_length:
|
|
276
|
+
context_window.pop(0)
|
|
277
|
+
context_window.append(sentence)
|
|
278
|
+
return context_window
|
|
279
|
+
|
|
280
|
+
def check_filter_for_analysis(sentence, context_window, analysis_type, filter, filter_key, filter_cutoff, terms_to_check, query_count):
|
|
281
|
+
decisions_count, subscore = None, None
|
|
282
|
+
|
|
283
|
+
if filter:
|
|
284
|
+
if filter == 'score':
|
|
285
|
+
success = sentence.get(filter, None)
|
|
286
|
+
else:
|
|
287
|
+
success = sentence.get(f'subscore_{filter}', None)
|
|
288
|
+
if success not in ['', 'nan', None]:
|
|
289
|
+
success = float(success)
|
|
290
|
+
if success >= filter_cutoff:
|
|
291
|
+
decisions_count, subscore = analyze_the_current_sentence(sentence, context_window, analysis_type, terms_to_check, query_count)
|
|
292
|
+
else:
|
|
293
|
+
decisions_count, subscore = analyze_the_current_sentence(sentence, context_window, analysis_type, terms_to_check, query_count)
|
|
294
|
+
|
|
295
|
+
return decisions_count, subscore
|
|
296
|
+
|
|
297
|
+
def handle_classified_sentence(sentence, analysis_type, decisions_count, subscore, existing_results, terms_to_check):
|
|
298
|
+
|
|
299
|
+
classified_sentence = {}
|
|
300
|
+
|
|
301
|
+
for key, value in sentence.items():
|
|
302
|
+
if key not in classified_sentence:
|
|
303
|
+
classified_sentence[key] = value
|
|
304
|
+
|
|
305
|
+
if analysis_type in classification_analyses:
|
|
306
|
+
classified_sentence.update({
|
|
307
|
+
f"decisions_{analysis_type}": decisions_count,
|
|
308
|
+
f"subscore_{analysis_type}": subscore
|
|
309
|
+
})
|
|
310
|
+
elif analysis_type == "list_of_terms":
|
|
311
|
+
classified_sentence.update({
|
|
312
|
+
f"decisions_{analysis_type}_{terms_to_check[0]}": decisions_count,
|
|
313
|
+
f"subscore_{analysis_type}_{terms_to_check[0]}": subscore
|
|
314
|
+
})
|
|
315
|
+
else:
|
|
316
|
+
if decisions_count:
|
|
317
|
+
i = 1
|
|
318
|
+
data = json.loads(decisions_count)
|
|
319
|
+
for relationship in data["relationships"]:
|
|
320
|
+
classified_sentence.update({
|
|
321
|
+
"cause": relationship["cause"],
|
|
322
|
+
"effect": relationship["effect"],
|
|
323
|
+
"relationship_id": i
|
|
324
|
+
})
|
|
325
|
+
existing_results.append(classified_sentence.copy())
|
|
326
|
+
i += 1
|
|
327
|
+
return
|
|
328
|
+
else:
|
|
329
|
+
classified_sentence.update({"cause": None, "effect": None})
|
|
330
|
+
|
|
331
|
+
existing_results.append(classified_sentence)
|
|
332
|
+
|
|
333
|
+
def document_analysis(log_file, filter, filter_key, filter_cutoff, terms_to_check):
|
|
334
|
+
with open(log_file, 'a', newline='') as log:
|
|
335
|
+
writer = csv.writer(log)
|
|
336
|
+
writer.writerow(['filter', 'filter_key', 'filter_cutoff', 'terms_to_check'])
|
|
337
|
+
writer.writerow([filter, filter_key, filter_cutoff, terms_to_check])
|
|
338
|
+
|
|
339
|
+
def analyze_all_sentences(sentences, analysis_type, save_name, context_length=0, filter=None, filter_key=0, filter_cutoff=0, terms_to_check=[], query_count=3):
|
|
340
|
+
temp_file, log_file, existing_results = initialize_files_analysis(save_name, analysis_type, terms_to_check)
|
|
341
|
+
start_time, end_time, start_document_id, start_sentence_id = get_start_ids(log_file)
|
|
342
|
+
#existing_results = pd.read_csv(temp_file).to_dict('records') if os.path.exists(temp_file) else []
|
|
343
|
+
|
|
344
|
+
print(f"Starting '{analysis_type}' analysis with '{filter}' filter (filter key: {filter_key}, filter cutoff: {filter_cutoff}) at document {start_document_id}, sentence {start_sentence_id}")
|
|
345
|
+
|
|
346
|
+
context_window = []
|
|
347
|
+
for sentence in tqdm(sentences):
|
|
348
|
+
document_id = int(sentence["document_id"])
|
|
349
|
+
sentence_id = int(sentence["sentence_id"])
|
|
350
|
+
|
|
351
|
+
# Skip the document if all sentences have been analyzed (if not, skip all sentences that have been analyzed)
|
|
352
|
+
if document_id < start_document_id or (document_id == start_document_id and sentence_id <= start_sentence_id):
|
|
353
|
+
continue
|
|
354
|
+
|
|
355
|
+
# Reset context window at the start of each document
|
|
356
|
+
if sentence_id == 1:
|
|
357
|
+
context_window = []
|
|
358
|
+
# Todo: Recreate context window if analysis interrupted and restarted midway
|
|
359
|
+
|
|
360
|
+
decisions_count, subscore = check_filter_for_analysis(sentence, context_window, analysis_type, filter, filter_key, filter_cutoff, terms_to_check, query_count)
|
|
361
|
+
context_window = append_to_context(context_window, sentence["sentence"], context_length)
|
|
362
|
+
handle_classified_sentence(sentence, analysis_type, decisions_count, subscore, existing_results, terms_to_check)
|
|
363
|
+
end_time = datetime.now().strftime('%Y-%m-%d %H:%M:%S')
|
|
364
|
+
try:
|
|
365
|
+
pd.DataFrame(existing_results).to_csv(temp_file, index=False)
|
|
366
|
+
update_log(log_file, start_time, end_time, document_id, sentence_id)
|
|
367
|
+
except KeyboardInterrupt:
|
|
368
|
+
print(f"Saving interrupted at document {document_id}, sentence {sentence_id}. Continuing until saving completed.")
|
|
369
|
+
pd.DataFrame(existing_results).to_csv(temp_file, index=False)
|
|
370
|
+
update_log(log_file, start_time, end_time, document_id, sentence_id)
|
|
371
|
+
raise
|
|
372
|
+
|
|
373
|
+
results_folder = os.path.join('./Results', save_name)
|
|
374
|
+
pd.DataFrame(existing_results).to_csv(os.path.join(results_folder, f'{save_name}_sentences.csv'), index=False)
|
|
375
|
+
document_analysis(log_file, filter, filter_key, filter_cutoff, terms_to_check)
|
|
376
|
+
|
|
377
|
+
return existing_results
|
|
378
|
+
|
|
379
|
+
def score_sentences(sentences, subscores, weights, save_name):
|
|
380
|
+
"""Calculate score for the list of sentences."""
|
|
381
|
+
TEMP_FILE = os.path.join('./Results/' + save_name + '_scores_temp.csv')
|
|
382
|
+
|
|
383
|
+
scored_sentences = []
|
|
384
|
+
|
|
385
|
+
for sentence in sentences:
|
|
386
|
+
score = 0
|
|
387
|
+
for i in range(len(subscores)):
|
|
388
|
+
for key in sentence.keys():
|
|
389
|
+
if key == 'subscore_' + subscores[i]:
|
|
390
|
+
if sentence[key] not in ['', 'nan']:
|
|
391
|
+
score += float(sentence[key]) * weights[i]
|
|
392
|
+
break
|
|
393
|
+
|
|
394
|
+
scored_sentence = {
|
|
395
|
+
"score": round(score * 100, 2)
|
|
396
|
+
}
|
|
397
|
+
|
|
398
|
+
# Dynamically add keys from the CSV file to classified_sentence
|
|
399
|
+
for key, value in sentence.items():
|
|
400
|
+
if key not in scored_sentence:
|
|
401
|
+
scored_sentence[key] = value
|
|
402
|
+
|
|
403
|
+
scored_sentences.append(scored_sentence) # Append modified dictionary to new list
|
|
404
|
+
|
|
405
|
+
# Save the list of sentences as a csv file
|
|
406
|
+
with open(TEMP_FILE, 'w') as f:
|
|
407
|
+
pd.DataFrame(scored_sentences).to_csv(f, index=False)
|
|
408
|
+
results_folder = './Results/' + save_name + '/'
|
|
409
|
+
save_to_csv(scored_sentences, results_folder, save_name + '_sentences.csv')
|
|
410
|
+
os.remove(TEMP_FILE)
|
|
411
|
+
|
|
412
|
+
def retrieve_concepts(project_name, file_name):
|
|
413
|
+
df = pd.read_csv(os.path.join("./Results", project_name, file_name))
|
|
414
|
+
|
|
415
|
+
causes = df["cause"].tolist()
|
|
416
|
+
effects = df["effect"].tolist()
|
|
417
|
+
|
|
418
|
+
all_concepts = causes + effects
|
|
419
|
+
valid_concepts = [concept for concept in all_concepts if pd.notna(concept)]
|
|
420
|
+
|
|
421
|
+
concepts = list(set(valid_concepts))
|
|
422
|
+
|
|
423
|
+
return concepts
|
|
424
|
+
|
|
425
|
+
def create_matrix(concepts):
|
|
426
|
+
size = len(concepts)
|
|
427
|
+
|
|
428
|
+
matrix = np.zeros((size, size), dtype = int)
|
|
429
|
+
matrix = pd.DataFrame(matrix, index = concepts, columns = concepts)
|
|
430
|
+
|
|
431
|
+
return matrix
|
|
432
|
+
|
|
433
|
+
def initialize_files_comparison(project_name, file_name):
|
|
434
|
+
base_path = "./Results"
|
|
435
|
+
results_file = os.path.join(base_path, project_name, project_name + "_comparison.csv")
|
|
436
|
+
temp_file = os.path.join(base_path, project_name, project_name + "_comparison_temp.csv")
|
|
437
|
+
log_file = os.path.join(base_path, project_name, project_name + "_comparison_log.csv")
|
|
438
|
+
|
|
439
|
+
os.makedirs(base_path, exist_ok=True)
|
|
440
|
+
|
|
441
|
+
# Ensure results file exists or create an empty one
|
|
442
|
+
if not os.path.exists(results_file):
|
|
443
|
+
|
|
444
|
+
# Create a list of concepts from causes and effects
|
|
445
|
+
concepts = retrieve_concepts(project_name, file_name)
|
|
446
|
+
|
|
447
|
+
# Create a symmetrical matrix from the list of concepts
|
|
448
|
+
matrix_df = create_matrix(concepts)
|
|
449
|
+
matrix_df.to_csv(results_file)
|
|
450
|
+
|
|
451
|
+
# Ensure log file exists or create with default values
|
|
452
|
+
if not os.path.exists(log_file):
|
|
453
|
+
with open(log_file, 'w') as log:
|
|
454
|
+
writer = csv.writer(log)
|
|
455
|
+
writer.writerow(['start_time', 'end_time', 'last_processed_concept_1', 'last_processed_concept_2', 'concepts'])
|
|
456
|
+
writer.writerow([datetime.now().strftime('%Y-%m-%d %H:%M:%S'), 0, 0, 0, len(concepts)])
|
|
457
|
+
|
|
458
|
+
# Initialize temp file with empty DataFrame if it doesn't exist
|
|
459
|
+
if not os.path.exists(temp_file):
|
|
460
|
+
matrix_df.to_csv(temp_file)
|
|
461
|
+
else:
|
|
462
|
+
matrix_df = pd.read_csv(temp_file, index_col = 0)
|
|
463
|
+
concepts = matrix_df.index.tolist()
|
|
464
|
+
|
|
465
|
+
return results_file, temp_file, log_file, matrix_df, concepts
|
|
466
|
+
|
|
467
|
+
def get_start_ids_comparison(log_file):
|
|
468
|
+
try:
|
|
469
|
+
with open(log_file, 'r') as log:
|
|
470
|
+
reader = csv.reader(log)
|
|
471
|
+
next(reader) # Skip header
|
|
472
|
+
row = next(reader)
|
|
473
|
+
start_time = row[0]
|
|
474
|
+
end_time = row[1]
|
|
475
|
+
start_concept_1_id, start_concept_2_id, num_concepts = map(int, row[2:])
|
|
476
|
+
except EmptyDataError:
|
|
477
|
+
start_time, end_time, start_concept_1_id, start_concept_2_id, num_concepts = 0, 0, 0, 0, 0 # Default values
|
|
478
|
+
|
|
479
|
+
return start_time, end_time, start_concept_1_id, start_concept_2_id, num_concepts
|
|
480
|
+
|
|
481
|
+
def update_log_comparison(log_file, start_time, end_time, concept_1_id, concept_2_id, num_concepts):
|
|
482
|
+
with open(log_file, 'w') as log:
|
|
483
|
+
writer = csv.writer(log)
|
|
484
|
+
writer.writerow(['start_time', 'end_time', 'last_processed_concept_1', 'last_processed_concept_2', 'num_concepts'])
|
|
485
|
+
writer.writerow([start_time, end_time, concept_1_id, concept_2_id, num_concepts])
|
|
486
|
+
|
|
487
|
+
def handle_concept_pair(concept_1, concept_2, query_count):
|
|
488
|
+
|
|
489
|
+
decisions = []
|
|
490
|
+
|
|
491
|
+
# Compare the pair of concepts
|
|
492
|
+
for _ in range(query_count):
|
|
493
|
+
|
|
494
|
+
analysis_type = "compare_concepts"
|
|
495
|
+
with open(f'./FlashQDA/Prompts/{analysis_type}.txt', 'r') as file:
|
|
496
|
+
system_prompt = "You are a helpful assistant that compares concepts. Follow the user's instructions carefully. Respond using JSON."
|
|
497
|
+
prompt = file.read()
|
|
498
|
+
prompt = prompt.format(concept_1 = concept_1, concept_2 = concept_2)
|
|
499
|
+
decision = send_to_openai(system_prompt, prompt)
|
|
500
|
+
|
|
501
|
+
data = json.loads(decision)
|
|
502
|
+
for key, value in data.items():
|
|
503
|
+
decisions.append(value)
|
|
504
|
+
|
|
505
|
+
#print(decisions)
|
|
506
|
+
|
|
507
|
+
similarity = (decisions.count("equivalent") * 1 +
|
|
508
|
+
decisions.count("very similar") * 0.75 +
|
|
509
|
+
decisions.count("somewhat similar") * 0.50 +
|
|
510
|
+
decisions.count("not very similar") * 0.25 +
|
|
511
|
+
decisions.count("unrelated") * 0) / len(decisions)
|
|
512
|
+
|
|
513
|
+
return similarity
|
|
514
|
+
|
|
515
|
+
def compare_concepts(project_name, file_name, query_count=3):
|
|
516
|
+
|
|
517
|
+
# Compare each pair of concepts
|
|
518
|
+
results_file, temp_file, log_file, matrix_df, concepts = initialize_files_comparison(project_name, file_name)
|
|
519
|
+
start_time, end_time, start_concept_1_id, start_concept_2_id, num_concepts = get_start_ids_comparison(log_file)
|
|
520
|
+
print(f"Starting concept comparison at concept {start_concept_1_id}, concept {start_concept_2_id}")
|
|
521
|
+
|
|
522
|
+
for i in tqdm(range(num_concepts)):
|
|
523
|
+
|
|
524
|
+
if i >= start_concept_1_id:
|
|
525
|
+
|
|
526
|
+
for j in tqdm(range(num_concepts)):
|
|
527
|
+
|
|
528
|
+
if j >= start_concept_2_id:
|
|
529
|
+
|
|
530
|
+
if i == j:
|
|
531
|
+
continue
|
|
532
|
+
if i > j:
|
|
533
|
+
matrix_df.iloc[i, j] = matrix_df.iloc[j, i]
|
|
534
|
+
else:
|
|
535
|
+
similarity = handle_concept_pair(concepts[i], concepts[j], query_count)
|
|
536
|
+
matrix_df.iloc[i, j] = similarity
|
|
537
|
+
|
|
538
|
+
end_time = datetime.now().strftime('%Y-%m-%d %H:%M:%S')
|
|
539
|
+
|
|
540
|
+
try:
|
|
541
|
+
pd.DataFrame(matrix_df).to_csv(temp_file)
|
|
542
|
+
update_log_comparison(log_file, start_time, end_time, i, j, num_concepts)
|
|
543
|
+
except KeyboardInterrupt:
|
|
544
|
+
print(f"Saving interrupted at concept {i}, concept {j}. Continuing until saving completed.")
|
|
545
|
+
pd.DataFrame(matrix_df).to_csv(temp_file)
|
|
546
|
+
update_log_comparison(log_file, start_time, end_time, i, j, num_concepts)
|
|
547
|
+
raise
|
|
548
|
+
|
|
549
|
+
start_concept_2_id += 1
|
|
550
|
+
|
|
551
|
+
start_concept_2_id = 0
|
|
552
|
+
start_concept_1_id += 1
|
|
553
|
+
|
|
554
|
+
pd.DataFrame(matrix_df).to_csv(results_file)
|
|
555
|
+
print(f"Finished comparing concepts at concept {i+1}, concept {j+1}. Results saved at {results_file}.")
|