-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathexecutive_summarization.py
More file actions
160 lines (127 loc) · 6.38 KB
/
Copy pathexecutive_summarization.py
File metadata and controls
160 lines (127 loc) · 6.38 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
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
import sqlite3
import os
import torch
from tqdm import tqdm
from transformers import pipeline
from sentence_transformers import SentenceTransformer
# --- Configuration ---
DATABASE_FILE = "econsultation.db"
# We only need the summarizer and key points models for this final version
SUMMARIZER_MODEL_NAME = "facebook/bart-large-cnn"
KEY_POINTS_MODEL_NAME = "google/flan-t5-base"
def run_section_analysis(conn, summarizer, key_points_extractor):
"""
Part 1: Generates a summary paragraph and bulleted key points for
each section based on all of its comments.
"""
cursor = conn.cursor()
print("--- Starting Part 1: Section-Wise Executive Analysis ---")
cursor.execute("SELECT section_id FROM sections")
all_section_ids = [row['section_id'] for row in cursor.fetchall()]
section_updates = []
for section_id in tqdm(all_section_ids, desc="Analyzing Sections"):
# Check if already processed to allow re-running the script
cursor.execute("SELECT section_ai_key_points FROM sections WHERE section_id = ?", (section_id,))
if cursor.fetchone()['section_ai_key_points']:
print(f"\nSkipping Section ID: {section_id} - already analyzed.")
continue
# Gather all relevant comments for the section
cursor.execute("""
SELECT comment_text FROM comments
WHERE section_id = ? AND comment_text IS NOT NULL AND LENGTH(comment_text) > 20
""", (section_id,))
comments = [row['comment_text'] for row in cursor.fetchall()]
if len(comments) < 2:
print(f"\nSkipping Section ID: {section_id} (not enough comments for a meaningful summary).")
continue
print(f"\nProcessing Section ID: {section_id} with {len(comments)} comments...")
combined_text = "\n\n".join(comments)
try:
# Generate the executive summary paragraph for the whole section
summary_paragraph = summarizer(
combined_text, max_length=256, min_length=64, do_sample=False, truncation=True
)[0]['summary_text']
# Use the instruction model to extract clean bullet points from that summary
key_point_prompt = f"Extract the key points as a bulleted list from the following text:\n{summary_paragraph}"
key_points = key_points_extractor(
key_point_prompt, max_length=256, truncation=True
)[0]['generated_text']
section_updates.append((summary_paragraph, key_points, section_id))
except Exception as e:
print(f" - ERROR processing Section ID {section_id}: {e}")
if section_updates:
update_query = "UPDATE sections SET section_ai_summary = ?, section_ai_key_points = ? WHERE section_id = ?"
cursor.executemany(update_query, section_updates)
conn.commit()
print(f"\nPart 1 Complete: Successfully updated {len(section_updates)} sections with executive analysis.")
else:
print("\nPart 1 Complete: No new sections to update.")
def run_draft_analysis_simplified(conn, summarizer):
"""
Part 2 (Simplified): Generates a single summary paragraph for each draft by
rolling up the section-level summaries.
"""
cursor = conn.cursor()
print("\n--- Starting Part 2: Draft-Wise Roll-up Analysis ---")
cursor.execute("SELECT draft_id FROM drafts")
all_draft_ids = [row['draft_id'] for row in cursor.fetchall()]
draft_updates = []
for draft_id in tqdm(all_draft_ids, desc="Analyzing Drafts"):
# Check if already processed
cursor.execute("SELECT draft_ai_summary FROM drafts WHERE draft_id = ?", (draft_id,))
if cursor.fetchone()['draft_ai_summary']:
print(f"\nSkipping Draft ID: {draft_id} - already analyzed.")
continue
# Gather the summaries from the child sections
cursor.execute("""
SELECT section_ai_summary FROM sections
WHERE draft_id = ? AND section_ai_summary IS NOT NULL AND section_ai_summary != ''
""", (draft_id,))
section_summaries = [row['section_ai_summary'] for row in cursor.fetchall()]
if not section_summaries:
print(f"\nSkipping Draft ID: {draft_id} (no section summaries to roll up).")
continue
print(f"\nProcessing Draft ID: {draft_id} by rolling up {len(section_summaries)} section summaries...")
try:
# Create a "summary of summaries"
combined_summaries = "\n\n".join(section_summaries)
draft_summary_paragraph = summarizer(
combined_summaries, max_length=400, min_length=100, do_sample=False, truncation=True
)[0]['summary_text']
draft_updates.append((draft_summary_paragraph, draft_id))
except Exception as e:
print(f" - ERROR processing Draft ID {draft_id}: {e}")
if draft_updates:
# Note: We are only updating one column now.
update_query = "UPDATE drafts SET draft_ai_summary = ? WHERE draft_id = ?"
cursor.executemany(update_query, draft_updates)
conn.commit()
print(f"\nPart 2 Complete: Successfully updated {len(draft_updates)} drafts.")
else:
print("\nPart 2 Complete: No new drafts to update.")
def main():
"""
Main orchestrator for the entire Phase 2 process.
"""
device = "cuda:0" if torch.cuda.is_available() else "cpu"
print(f"Using device: {device}")
conn = None
try:
conn = sqlite3.connect(DATABASE_FILE)
conn.row_factory = sqlite3.Row
print("Successfully connected to the database.")
print("Loading AI models... (This may take several minutes)")
summarizer = pipeline("summarization", model=SUMMARIZER_MODEL_NAME, device=device)
key_points_extractor = pipeline("text2text-generation", model=KEY_POINTS_MODEL_NAME, device=device)
print("All models loaded successfully.\n")
# --- RUN THE PROCESS IN THE CORRECT ORDER ---
run_section_analysis(conn, summarizer, key_points_extractor)
run_draft_analysis_simplified(conn, summarizer)
except Exception as e:
print(f"A critical error occurred in the main process: {e}")
finally:
if conn:
conn.close()
print("\nPhase 2 analysis is complete. Database connection closed.")
if __name__ == '__main__':
main()