f7a9d80162
Feature/dtc clean * dtc v1 initial commit * second layer of prompt * Added MAX_PAGES_TO_CHECK and page looping to get answer from llm * Edited contract_type prompt to get cleaner output with no additional text * Updated document_type_classification prompt * Added non-contract type classifcation * Added dtc_v1_test.py for testing only * Updated prompt * Updated prompt * at scale changes * add save feature for individual file output * update dockerfile * clean dtc to prepare for merge * clean up dtc script * Merge remote-tracking branch 'origin/main' into feature/dtc-clean * reafactor code * Merge remote-tracking branch 'origin/main' into feature/dtc-clean Approved-by: Katon Minhas
329 lines
12 KiB
Python
329 lines
12 KiB
Python
import concurrent.futures
|
|
import logging
|
|
import os
|
|
from datetime import datetime
|
|
from threading import Lock
|
|
|
|
import pandas as pd
|
|
|
|
from src.investment.preprocessing_funcs import (
|
|
clean_newlines,
|
|
split_text,
|
|
clean_law_symbols,
|
|
)
|
|
from src.prompts.fieldset import Field
|
|
from src.utils import io_utils, llm_utils
|
|
from src import config
|
|
|
|
# Load prompts from JSON using Field class (with thread-safe caching)
|
|
# Load prompts at module level (cached for reuse across threads)
|
|
DOCUMENT_TYPE_CLASSIFICATION_PROMPT = Field.load_from_file(
|
|
config.DTC_PROMPTS_JSON_PATH, "document_type_classification"
|
|
).prompt
|
|
|
|
CONTRACT_TYPE_PROMPT = Field.load_from_file(
|
|
config.DTC_PROMPTS_JSON_PATH, "contract_type"
|
|
).prompt
|
|
|
|
NON_CONTRACT_TYPE_PROMPT = Field.load_from_file(
|
|
config.DTC_PROMPTS_JSON_PATH, "non_contract_type"
|
|
).prompt
|
|
|
|
# Thread-safe counter for progress tracking
|
|
progress_lock = Lock()
|
|
progress_counter = {"completed": 0, "total": 0}
|
|
|
|
|
|
def call_llm(filename, context, question):
|
|
"""Call Claude through llm_utils.invoke_claude (unified LLM interface)
|
|
|
|
Args:
|
|
filename: Name of the file being processed (for logging/tracking)
|
|
context: Document text context
|
|
question: The prompt/question to ask the LLM
|
|
|
|
Returns:
|
|
str: The LLM's response, or "ERROR" if the call fails
|
|
"""
|
|
# Prepare the full prompt with context
|
|
prompt = f"Document text:\n{context}\n\n{question}"
|
|
|
|
try:
|
|
# Use the unified llm_utils interface (supports both local and EC2 modes)
|
|
answer = llm_utils.invoke_claude(
|
|
prompt=prompt,
|
|
model_id="sonnet_latest", # Uses config.MODEL_ALIASES["sonnet_latest"]
|
|
filename=filename,
|
|
max_tokens=8192,
|
|
)
|
|
return answer.strip()
|
|
|
|
except Exception as e:
|
|
logging.error(f"Error calling Claude for {filename}: {str(e)}")
|
|
return "ERROR"
|
|
|
|
|
|
def process_single_file(filename, contract_text, json_folder):
|
|
"""Process a single file to determine if it's a contract and classify it.
|
|
|
|
Args:
|
|
filename: Name of the file being processed
|
|
contract_text: Raw text content of the file
|
|
json_folder: Path to folder for saving individual JSON results
|
|
|
|
Returns:
|
|
tuple: (filename, result_dict)
|
|
"""
|
|
try:
|
|
contract_text = clean_newlines(contract_text)
|
|
contract_text = clean_law_symbols(contract_text)
|
|
text_dict = split_text(contract_text)
|
|
|
|
# Loop through n pages
|
|
is_contract = "NO"
|
|
contract_type = "N/A"
|
|
non_contract_type = "N/A"
|
|
found_on_page = "N/A"
|
|
|
|
for i in range(config.DTC_MAX_PAGES_TO_CHECK):
|
|
page_key = str(i + 1)
|
|
|
|
# Handle edge case: document has fewer pages than config.DTC_MAX_PAGES_TO_CHECK
|
|
if page_key not in text_dict:
|
|
break
|
|
|
|
context = text_dict[page_key]
|
|
|
|
# Ask if this page is a contract
|
|
is_contract_answer = call_llm(
|
|
filename,
|
|
context,
|
|
DOCUMENT_TYPE_CLASSIFICATION_PROMPT,
|
|
)
|
|
|
|
if is_contract_answer.strip().upper() == "YES":
|
|
# If contract found, get contract type
|
|
type_of_contract_answer = call_llm(
|
|
filename, context, CONTRACT_TYPE_PROMPT
|
|
)
|
|
is_contract = "YES"
|
|
contract_type = type_of_contract_answer
|
|
found_on_page = i + 1
|
|
break # Stop checking further pages once contract is found
|
|
|
|
# If not a contract, determine non-contract document type
|
|
if is_contract == "NO":
|
|
# Use the first page for non-contract classification
|
|
if "1" in text_dict:
|
|
context = text_dict["1"]
|
|
non_contract_type_answer = call_llm(
|
|
filename, context, NON_CONTRACT_TYPE_PROMPT
|
|
)
|
|
non_contract_type = non_contract_type_answer
|
|
found_on_page = 1
|
|
|
|
result = {
|
|
"is_contract": is_contract,
|
|
"contract_type": contract_type,
|
|
"non_contract_type": non_contract_type,
|
|
"found_on_page": found_on_page,
|
|
}
|
|
|
|
# Save result to JSON immediately
|
|
io_utils.save_result_to_json(filename, result, json_folder)
|
|
|
|
# Update progress counter
|
|
with progress_lock:
|
|
progress_counter["completed"] += 1
|
|
completed = progress_counter["completed"]
|
|
total = progress_counter["total"]
|
|
if (
|
|
completed % 5 == 0 or completed == total
|
|
): # Log every 5 files or at completion
|
|
logging.info(
|
|
f"Progress: {completed}/{total} files processed ({completed/total*100:.1f}%)"
|
|
)
|
|
|
|
return filename, result
|
|
|
|
except Exception as e:
|
|
logging.error(f"Error processing {filename}: {str(e)}")
|
|
error_result = {
|
|
"is_contract": "ERROR",
|
|
"contract_type": "ERROR",
|
|
"non_contract_type": "ERROR",
|
|
"found_on_page": "ERROR",
|
|
}
|
|
# Save error result to JSON
|
|
io_utils.save_result_to_json(filename, error_result, json_folder)
|
|
return filename, error_result
|
|
|
|
|
|
def main(input_dict, run_timestamp):
|
|
"""Run document type classification on a dictionary of files.
|
|
|
|
This function is designed to be called from investment/main.py as a pre-filter.
|
|
It classifies all files and returns a dict mapping filenames to classification results.
|
|
|
|
Args:
|
|
input_dict: Dictionary mapping filename -> file_text
|
|
run_timestamp: Timestamp string for organizing output files (format: run_YYYYMMDD_HH-MM_BATCHID)
|
|
|
|
Returns:
|
|
dict: Mapping of filename -> classification result dict with keys:
|
|
- is_contract: "YES" or "NO"
|
|
- contract_type: Type if contract, otherwise "N/A"
|
|
- non_contract_type: Type if non-contract, otherwise "N/A"
|
|
- found_on_page: Page number where classification was determined
|
|
|
|
Side effects:
|
|
- Saves results to S3 (if config.WRITE_TO_S3) or local filesystem
|
|
- Saves individual JSON files to config.DTC_JSON_OUTPUT_FOLDER
|
|
"""
|
|
logging.info(f"Starting DTC classification for {len(input_dict)} files...")
|
|
logging.info(f"Max pages to check: {config.DTC_MAX_PAGES_TO_CHECK}")
|
|
logging.info(f"Results will be saved to: {config.DTC_OUTPUT_FILE}")
|
|
|
|
# Initialize progress counter
|
|
progress_counter["total"] = len(input_dict)
|
|
progress_counter["completed"] = 0
|
|
|
|
# Create output folder
|
|
os.makedirs(config.DTC_JSON_OUTPUT_FOLDER, exist_ok=True)
|
|
logging.info(f"JSON output folder: {config.DTC_JSON_OUTPUT_FOLDER}")
|
|
|
|
# Record start time
|
|
start_time = datetime.now()
|
|
|
|
answer_dict = {}
|
|
|
|
# Process files concurrently
|
|
with concurrent.futures.ThreadPoolExecutor(max_workers=10) as executor:
|
|
future_to_file = {}
|
|
for filename, contract_text in input_dict.items():
|
|
future = executor.submit(
|
|
process_single_file,
|
|
filename,
|
|
contract_text,
|
|
config.DTC_JSON_OUTPUT_FOLDER,
|
|
)
|
|
future_to_file[future] = filename
|
|
|
|
# Collect results as they complete
|
|
for future in concurrent.futures.as_completed(future_to_file):
|
|
filename = future_to_file[future]
|
|
try:
|
|
result_filename, result_data = future.result()
|
|
answer_dict[result_filename] = result_data
|
|
except Exception as e:
|
|
logging.error(f"Exception occurred for {filename}: {str(e)}")
|
|
answer_dict[filename] = {
|
|
"is_contract": "ERROR",
|
|
"contract_type": "ERROR",
|
|
"non_contract_type": "ERROR",
|
|
"found_on_page": "ERROR",
|
|
}
|
|
|
|
# Record end time
|
|
end_time = datetime.now()
|
|
duration = (end_time - start_time).total_seconds()
|
|
|
|
logging.info(f"DTC classification complete in {duration:.2f} seconds")
|
|
logging.info(f"Average time per file: {duration/len(input_dict):.2f} seconds")
|
|
|
|
# Create DataFrame and save results
|
|
df = pd.DataFrame([{"FILE_NAME": k, **v} for k, v in answer_dict.items()])
|
|
|
|
# Log summary statistics
|
|
contracts_found = len(df[df["is_contract"] == "YES"])
|
|
non_contracts = len(df[df["is_contract"] == "NO"])
|
|
errors = len(df[df["is_contract"] == "ERROR"])
|
|
|
|
logging.info(
|
|
f"DTC Summary: {contracts_found} contracts, {non_contracts} non-contracts, {errors} errors"
|
|
)
|
|
|
|
# Create summary DataFrame with aggregate statistics
|
|
summary_data = {
|
|
"total_files": [len(df)],
|
|
"contracts": [contracts_found],
|
|
"non_contracts": [non_contracts],
|
|
"errors": [errors],
|
|
"contract_rate": [
|
|
f"{(contracts_found/len(df)*100):.1f}%" if len(df) > 0 else "0%"
|
|
],
|
|
"run_timestamp": [run_timestamp],
|
|
"batch_id": [config.BATCH_ID],
|
|
}
|
|
|
|
# Add contract type breakdown
|
|
if contracts_found > 0:
|
|
contract_types = df[df["is_contract"] == "YES"]["contract_type"].value_counts()
|
|
summary_data["most_common_contract_type"] = [
|
|
contract_types.index[0] if len(contract_types) > 0 else "N/A"
|
|
]
|
|
summary_data["contract_type_breakdown"] = [contract_types.to_dict()]
|
|
else:
|
|
summary_data["most_common_contract_type"] = ["N/A"]
|
|
summary_data["contract_type_breakdown"] = [{}]
|
|
|
|
# Add non-contract type breakdown
|
|
if non_contracts > 0:
|
|
non_contract_types = df[df["is_contract"] == "NO"][
|
|
"non_contract_type"
|
|
].value_counts()
|
|
summary_data["most_common_non_contract_type"] = [
|
|
non_contract_types.index[0] if len(non_contract_types) > 0 else "N/A"
|
|
]
|
|
summary_data["non_contract_type_breakdown"] = [non_contract_types.to_dict()]
|
|
else:
|
|
summary_data["most_common_non_contract_type"] = ["N/A"]
|
|
summary_data["non_contract_type_breakdown"] = [{}]
|
|
|
|
summary_df = pd.DataFrame(summary_data)
|
|
|
|
# Save detailed results using same pattern as investment pipeline
|
|
if config.WRITE_TO_S3:
|
|
io_utils.write_s3(df, "", run_timestamp, "dtc")
|
|
logging.info(
|
|
f"DTC results uploaded to S3: {config.BATCH_ID}/{run_timestamp}/{config.BATCH_ID}-DTC.csv"
|
|
)
|
|
|
|
# Save summary
|
|
io_utils.write_s3(summary_df, "", run_timestamp, "dtc_summary")
|
|
logging.info(
|
|
f"DTC summary uploaded to S3: {config.BATCH_ID}/{run_timestamp}/{config.BATCH_ID}-DTC-SUMMARY.csv"
|
|
)
|
|
else:
|
|
io_utils.write_local(df, "", run_timestamp, "dtc")
|
|
logging.info(
|
|
f"DTC results saved locally: {config.CONSOLIDATED_OUTPUT_DIRECTORY}/{run_timestamp}/{config.BATCH_ID}-DTC.csv"
|
|
)
|
|
|
|
# Save summary
|
|
io_utils.write_local(summary_df, "", run_timestamp, "dtc_summary")
|
|
logging.info(
|
|
f"DTC summary saved locally: {config.CONSOLIDATED_OUTPUT_DIRECTORY}/{run_timestamp}/{config.BATCH_ID}-DTC-SUMMARY.csv"
|
|
)
|
|
|
|
return answer_dict
|
|
|
|
|
|
if __name__ == "__main__":
|
|
logging.basicConfig(
|
|
level=logging.INFO,
|
|
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
|
|
)
|
|
|
|
# Generate timestamp for this run
|
|
run_timestamp = f"run_{datetime.now().strftime('%Y%m%d_%H-%M')}_{config.BATCH_ID}"
|
|
|
|
# Load input files
|
|
logging.info("Loading input files...")
|
|
input_dict = io_utils.read_input()
|
|
|
|
# Run classification
|
|
results = main(input_dict, run_timestamp)
|
|
|
|
logging.info(f"Classification complete. Processed {len(results)} files.")
|