TD/BU integration changes

This commit is contained in:
Katon Minhas
2024-06-03 14:53:40 -07:00
committed by Michael McGuinness
parent fa6ffb8dde
commit fc5ebe889b
6 changed files with 207 additions and 138 deletions
+134
View File
@@ -0,0 +1,134 @@
import os
import pandas as pd
import config
import preprocess
import table_funcs
import prompt_funcs
import prompts
import claude_funcs
def clean_td(td):
td_clean = []
for d in td:
new_d = {}
for k, v in d.items():
if 'DATE' in k:
new_d[k] = v if isinstance(v, list) else [v]
elif k not in ['page_num', 'Filename']:
if isinstance(v, str) and ',' in v:
new_d[k] = [item.strip() for item in v.split(',')]
elif v == 'N/A':
new_d[k] = []
else:
new_d[k] = [v] if isinstance(v, str) else v
else:
new_d[k] = v
td_clean.append(new_d)
return td_clean
def get_unique_keys(dicts):
keys = set()
for d in dicts:
keys.update(d.keys())
return keys
def get_unique_date_fields(td_results):
unique_date_fields = {}
for td in td_results:
for key, value in td.items():
if 'DATE' in key:
if key not in unique_date_fields:
unique_date_fields[key] = set()
unique_date_fields[key].update(value if isinstance(value, list) else [value])
for key in unique_date_fields:
unique_date_fields[key] = list(unique_date_fields[key])
return unique_date_fields
def merge_results(td_results, bu_results, text_dict):
all_keys = get_unique_keys(td_results) | get_unique_keys(bu_results)
print("All Keys:", all_keys)
date_fields = get_unique_date_fields(td_results)
print("Date Fields:", date_fields)
merged_results = []
for bu in bu_results:
merged = {key: bu.get(key, "") for key in all_keys}
td_on_page = [td for td in td_results if td['page_num'] == bu['page_num']]
for td in td_on_page:
for key, value in td.items():
if not merged[key]:
merged[key] = value
elif isinstance(value, list) and value and not isinstance(merged[key], list):
merged[key] = value
elif isinstance(value, list) and value:
merged[key].extend(value)
for date_key, date_values in date_fields.items():
if date_key not in merged or not merged[date_key]:
merged[date_key] = date_values
else:
merged[date_key].extend([val for val in date_values if val not in merged[date_key]])
for key in all_keys:
if isinstance(merged[key], list) and len(merged[key]) > 1:
page_num = bu['page_num']
page_text = text_dict.get(page_num, "")
prompt = prompts.GENERATE_PROMPT(bu, td_on_page, page_text, field_name=key, values=merged[key])
response = claude_funcs.invoke_claude_3(prompt, max_tokens=4000)
try:
selected_value = response.strip()
merged[key] = selected_value
except Exception as e:
print(f"Error processing LLM response for key {key}: {e}")
merged[key] = ', '.join(merged[key])
merged_results.append(merged)
return merged_results
def process_file(file_object):
filename, contract_text = file_object
if config.VERBOSE: print(f"Processing {filename}...")
#### PREPROCESS ####
contract_text = preprocess.clean_newlines(contract_text)
text_dict = preprocess.split_text(contract_text)
text_dict = table_funcs.align_and_format_tables(text_dict)
text_dict = preprocess.highlight_rates(text_dict)
# Run Top Down
td_results = prompt_funcs.run_top_down(filename, text_dict) # Returns list of dictionaries for each page
print("TD Results:", td_results)
# Run Bottom Up
bu_results = prompt_funcs.run_bottom_up(filename, text_dict) # Returns list of dictionaries
print("BU Results:", bu_results)
# Combine
combined_results = merge_results(clean_td(td_results), bu_results, text_dict)
print("Combined Results:", combined_results)
# Convert to DataFrame
combined_df = pd.DataFrame(combined_results)
# Post-process combined results
# post_processed_combined_df = postprocess.postprocess_results(combined_df)
# print("Post-processed Results:", post_processed_combined_df)
# Create directories
base_filename = os.path.splitext(filename)[0]
output_dir = os.path.join(config.OUTPUT_FOLDER, base_filename)
os.makedirs(output_dir, exist_ok=True)
# Save results
pd.DataFrame(td_results).to_csv(os.path.join(output_dir, 'td_results.csv'), index=False)
pd.DataFrame(bu_results).to_csv(os.path.join(output_dir, 'bu_results.csv'), index=False)
combined_df.to_csv(os.path.join(output_dir, 'combined_results.csv'), index=False)
# post_processed_combined_df.to_csv(os.path.join(output_dir, 'combined_results.csv'), index=False)