-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathfind_threads_v2.py
More file actions
330 lines (258 loc) · 11.4 KB
/
Copy pathfind_threads_v2.py
File metadata and controls
330 lines (258 loc) · 11.4 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
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
"""
Find invisible threads using topic extraction + embedding.
The key insight: embed TOPICS, not full insight text.
This captures conceptual similarity rather than vocabulary overlap.
Step 1: Extract TOPIC/THEME from each insight using LLM
Step 2: Embed topics (not full insights)
Step 3: Build similarity graph on topic embeddings
Step 4: Use Louvain to find communities
Step 5: Filter by min_episodes (true invisible threads span multiple conversations)
Usage:
modal run insights_first/find_threads_v2.py --input insights_first/data/modal_extraction_20260120_024600.json
"""
import modal
import json
import numpy as np
from datetime import datetime
from collections import defaultdict
app = modal.App("thread-finder-v2")
MODEL_ID = "Qwen/Qwen2.5-7B-Instruct"
model_volume = modal.Volume.from_name("qwen-model-cache", create_if_missing=True)
vllm_image = (
modal.Image.debian_slim(python_version="3.11")
.pip_install(
"vllm>=0.6.0",
"torch",
"transformers",
"huggingface_hub",
"sentence-transformers",
"networkx",
"python-louvain",
"scikit-learn",
)
)
# Topic extraction prompt - designed to capture the CONCEPT, not vocabulary
TOPIC_EXTRACTION_PROMPT = """Extract the core TOPIC and CLAIM from this insight.
## Rules:
1. TOPIC should be the underlying concept/principle, not the specific example
2. CLAIM should be the actionable takeaway, stated generically
3. Use common business vocabulary (not speaker-specific terms)
## Examples:
Insight: "Companies should fire their best customers first when pivoting, because they'll hold you back with feature requests for the old product"
→ TOPIC: "Managing customer relationships during strategic pivots"
→ CLAIM: "Prioritize future direction over existing customer demands when pivoting"
Insight: "The best PMs spend 80% of their time on problems they're NOT going to solve"
→ TOPIC: "Product manager time allocation and prioritization"
→ CLAIM: "Effective prioritization requires spending more time deciding what NOT to do"
Insight: "Duolingo's streak system is optimized for daily engagement by making the cost of breaking a streak feel significant"
→ TOPIC: "Gamification and habit formation in products"
→ CLAIM: "Loss aversion mechanics drive user engagement better than rewards"
## Current Insight:
{insight}
## Respond with JSON:
{{
"topic": "The underlying business concept in 5-10 words",
"claim": "The actionable principle in one sentence",
"category": "One of: hiring, firing, product, growth, leadership, culture, strategy, pricing, customer, team, metrics, communication, decision-making, other"
}}"""
@app.cls(
gpu="A10G",
image=vllm_image,
volumes={"/model-cache": model_volume},
timeout=600,
scaledown_window=300,
)
class TopicExtractor:
@modal.enter()
def load_model(self):
from vllm import LLM, SamplingParams
self.llm = LLM(
model=MODEL_ID,
download_dir="/model-cache",
trust_remote_code=True,
max_model_len=4096,
gpu_memory_utilization=0.9,
)
self.sampling_params = SamplingParams(temperature=0.2, max_tokens=300)
@modal.method()
def extract_topics_batch(self, insights: list[tuple]) -> list[dict]:
"""Extract topic and claim from a batch of insights."""
import re
prompts = [TOPIC_EXTRACTION_PROMPT.format(insight=text) for text, idx in insights]
outputs = self.llm.generate(prompts, self.sampling_params)
results = []
for (text, idx), output in zip(insights, outputs):
response = output.outputs[0].text.strip()
try:
json_match = re.search(r'\{[^{}]*\}', response, re.DOTALL)
if json_match:
result = json.loads(json_match.group())
else:
result = {"topic": text[:50], "claim": text[:100], "category": "other"}
except:
result = {"topic": text[:50], "claim": text[:100], "category": "other"}
result["idx"] = idx
result["insight_text"] = text
results.append(result)
return results
def build_topic_graph(topics: list[str], embeddings: np.ndarray, threshold: float = 0.55):
"""Build similarity graph based on topic embeddings."""
import networkx as nx
from sklearn.metrics.pairwise import cosine_similarity
sim_matrix = cosine_similarity(embeddings)
G = nx.Graph()
G.add_nodes_from(range(len(topics)))
edge_count = 0
for i in range(len(topics)):
for j in range(i + 1, len(topics)):
if sim_matrix[i, j] >= threshold:
G.add_edge(i, j, weight=sim_matrix[i, j])
edge_count += 1
return G, sim_matrix
def find_communities(G, min_size: int = 3, resolution: float = 1.5):
"""Find communities using Louvain algorithm."""
import community as community_louvain
# Only consider nodes with edges
nodes_with_edges = [n for n in G.nodes() if G.degree(n) > 0]
if len(nodes_with_edges) < min_size:
return []
subgraph = G.subgraph(nodes_with_edges)
partition = community_louvain.best_partition(subgraph, resolution=resolution, random_state=42)
# Group by community
communities = defaultdict(set)
for node, comm_id in partition.items():
communities[comm_id].add(node)
# Filter by min_size
valid = [c for c in communities.values() if len(c) >= min_size]
return sorted(valid, key=len, reverse=True)
@app.local_entrypoint()
def main(
input: str,
threshold: float = 0.55,
min_size: int = 3,
min_episodes: int = 3,
resolution: float = 1.5,
batch_size: int = 50,
output: str = None,
):
"""Find invisible threads using topic-based clustering."""
from sentence_transformers import SentenceTransformer
# Load insights
print(f"Loading insights from {input}...")
with open(input, 'r', encoding='utf-8') as f:
data = json.load(f)
insights = [r['insight'] for r in data['results'] if r['has_insight'] and r['insight']]
print(f"Loaded {len(insights)} insights")
# === STEP 1: Extract topics ===
print(f"\n{'='*60}")
print("STEP 1: Extracting topics from insights")
print('='*60)
extractor = TopicExtractor()
insight_inputs = [(i['insight_text'], idx) for idx, i in enumerate(insights)]
batches = [insight_inputs[i:i+batch_size] for i in range(0, len(insight_inputs), batch_size)]
import time
start = time.time()
all_topics = []
for i, batch in enumerate(batches):
results = extractor.extract_topics_batch.remote(batch)
all_topics.extend(results)
print(f" Batch {i+1}/{len(batches)} complete ({len(all_topics)} topics extracted)")
elapsed = time.time() - start
print(f"Topic extraction complete in {elapsed:.1f}s")
# Show sample topics
print(f"\nSample extracted topics:")
for t in all_topics[:5]:
print(f" [{t.get('category', '?')}] {t.get('topic', '?')}")
# === STEP 2: Embed topics ===
print(f"\n{'='*60}")
print("STEP 2: Embedding topics")
print('='*60)
model = SentenceTransformer("all-MiniLM-L6-v2")
# Embed the topic + claim (not the full insight)
topic_texts = [f"{t.get('topic', '')} - {t.get('claim', '')}" for t in all_topics]
print(f"Embedding {len(topic_texts)} topic descriptions...")
embeddings = model.encode(topic_texts, show_progress_bar=True)
# === STEP 3: Build graph and find communities ===
print(f"\n{'='*60}")
print(f"STEP 3: Finding threads (threshold={threshold})")
print('='*60)
G, sim_matrix = build_topic_graph(topic_texts, embeddings, threshold=threshold)
nodes_with_edges = sum(1 for n in G.nodes() if G.degree(n) > 0)
edges = G.number_of_edges()
print(f"Graph: {len(topic_texts)} nodes, {edges} edges, {nodes_with_edges} connected")
communities = find_communities(G, min_size=min_size, resolution=resolution)
print(f"Found {len(communities)} communities with >= {min_size} insights")
# === STEP 4: Filter by min_episodes ===
print(f"\n{'='*60}")
print(f"STEP 4: Filtering by episode coverage (>= {min_episodes} episodes)")
print('='*60)
threads = []
skipped = 0
for i, comm in enumerate(communities):
# Get insights in this community
thread_insights = [insights[all_topics[idx]['idx']] for idx in comm]
thread_topics = [all_topics[idx] for idx in comm]
# Count unique episodes
episodes = list(set(ins['document_id'] for ins in thread_insights))
if len(episodes) < min_episodes:
skipped += 1
continue
# Calculate coherence
comm_embeddings = embeddings[list(comm)]
upper_tri = sim_matrix[np.ix_(list(comm), list(comm))][np.triu_indices(len(comm), k=1)]
coherence = float(np.mean(upper_tri)) if len(upper_tri) > 0 else 1.0
# Find dominant category
categories = [t.get('category', 'other') for t in thread_topics]
from collections import Counter
dominant_cat = Counter(categories).most_common(1)[0][0]
# Get representative topics
rep_topics = list(set(t.get('topic', '') for t in thread_topics))[:5]
threads.append({
'thread_id': len(threads),
'size': len(comm),
'num_episodes': len(episodes),
'episodes': episodes,
'coherence': round(coherence, 3),
'category': dominant_cat,
'representative_topics': rep_topics,
'insights': thread_insights,
'topics': thread_topics,
'indices': list(comm),
})
print(f"Valid threads: {len(threads)} (skipped {skipped} with < {min_episodes} episodes)")
# Sort by episode coverage
threads.sort(key=lambda x: (x['num_episodes'], x['size']), reverse=True)
# Summary
print(f"\n{'='*60}")
print("RESULTS")
print('='*60)
total_in_threads = sum(t['size'] for t in threads)
print(f"Threads found: {len(threads)}")
print(f"Insights in threads: {total_in_threads}/{len(insights)} ({total_in_threads/len(insights)*100:.0f}%)")
print(f"\nTop threads:")
for t in threads[:10]:
print(f" Thread {t['thread_id']}: {t['size']} insights, {t['num_episodes']} episodes, [{t['category']}]")
print(f" Topics: {', '.join(t['representative_topics'][:2])}")
# Save output
if output is None:
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
output = f"insights_first/data/threads_v2_{timestamp}.json"
output_data = {
'metadata': {
'timestamp': datetime.now().isoformat(),
'input_file': input,
'approach': 'topic-based embedding',
'threshold': threshold,
'min_size': min_size,
'min_episodes': min_episodes,
'resolution': resolution,
'total_insights': len(insights),
'insights_in_threads': total_in_threads,
'threads_found': len(threads),
},
'threads': {f"Thread {t['thread_id']}": t for t in threads},
}
with open(output, 'w', encoding='utf-8') as f:
json.dump(output_data, f, indent=2, ensure_ascii=False)
print(f"\nSaved to: {output}")
return output_data