cat-stack 2.4.0__tar.gz → 2.5.0__tar.gz

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.
Files changed (45) hide show
  1. {cat_stack-2.4.0 → cat_stack-2.5.0}/PKG-INFO +1 -1
  2. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/__about__.py +1 -1
  3. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/collapse_themes.py +208 -6
  4. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/explore.py +1 -1
  5. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/extract.py +105 -4
  6. {cat_stack-2.4.0 → cat_stack-2.5.0}/.gitignore +0 -0
  7. {cat_stack-2.4.0 → cat_stack-2.5.0}/LICENSE +0 -0
  8. {cat_stack-2.4.0 → cat_stack-2.5.0}/README.md +0 -0
  9. {cat_stack-2.4.0 → cat_stack-2.5.0}/pyproject.toml +0 -0
  10. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/cat_stack/__init__.py +0 -0
  11. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/__init__.py +0 -0
  12. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/_batch.py +0 -0
  13. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/_category_analysis.py +0 -0
  14. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/_chunked.py +0 -0
  15. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/_embeddings.py +0 -0
  16. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/_formatter.py +0 -0
  17. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/_pilot_test.py +0 -0
  18. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/_prompts.py +0 -0
  19. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/_providers.py +0 -0
  20. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/_review_ui.py +0 -0
  21. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/_tiebreaker.py +0 -0
  22. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/_utils.py +0 -0
  23. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/_web_fetch.py +0 -0
  24. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/_wrapper_helpers.py +0 -0
  25. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/calls/CoVe.py +0 -0
  26. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/calls/__init__.py +0 -0
  27. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/calls/image_CoVe.py +0 -0
  28. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/calls/image_stepback.py +0 -0
  29. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/calls/pdf_CoVe.py +0 -0
  30. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/calls/pdf_stepback.py +0 -0
  31. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/calls/stepback.py +0 -0
  32. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/calls/top_n.py +0 -0
  33. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/classify.py +0 -0
  34. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/image_functions.py +0 -0
  35. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/images/circle.png +0 -0
  36. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/images/cube.png +0 -0
  37. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/images/diamond.png +0 -0
  38. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/images/overlapping_pentagons.png +0 -0
  39. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/images/rectangles.png +0 -0
  40. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/model_reference_list.py +0 -0
  41. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/pdf_functions.py +0 -0
  42. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/prompt_tune.py +0 -0
  43. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/summarize.py +0 -0
  44. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/text_functions.py +0 -0
  45. {cat_stack-2.4.0 → cat_stack-2.5.0}/src/catstack/text_functions_ensemble.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: cat-stack
3
- Version: 2.4.0
3
+ Version: 2.5.0
4
4
  Summary: Domain-agnostic text, image, PDF, and DOCX classification engine powered by LLMs
5
5
  Project-URL: Documentation, https://github.com/chrissoria/cat-stack#readme
6
6
  Project-URL: Issues, https://github.com/chrissoria/cat-stack/issues
@@ -1,7 +1,7 @@
1
1
  # SPDX-FileCopyrightText: 2025-present Christopher Soria <chrissoria@berkeley.edu>
2
2
  #
3
3
  # SPDX-License-Identifier: GPL-3.0-or-later
4
- __version__ = "2.4.0"
4
+ __version__ = "2.5.0"
5
5
  __author__ = "Chris Soria"
6
6
  __email__ = "chrissoria@berkeley.edu"
7
7
  __title__ = "cat-stack"
@@ -122,13 +122,18 @@ def _collapse_batch(client, batch, description, creativity, mode="unique"):
122
122
 
123
123
  mode="unique": extract unique categories only (remove restatements, keep
124
124
  distinct ones) — gentle, near-idempotent, guaranteed to only remove.
125
+ mode="prune": drop CONCEPTUAL duplicates (labels naming the same concept in
126
+ different words), keeping one representative per concept VERBATIM — never
127
+ renames or merges merely-related labels. Stronger than "unique" (catches
128
+ same-concept-different-words) but, unlike "merge", never collapses distinct
129
+ concepts. Output forced to a subset of the input.
125
130
  mode="merge": aggressively consolidate related labels into broader concepts
126
131
  while retaining meaningful distinctions — for a final compression step.
127
132
 
128
133
  Strict numbered-list prompt + strict parsing, so the reply is always a clean
129
134
  list and any stray prose is ignored. Guardrails: a failed call returns the
130
- batch unchanged (no data loss); in "unique" mode the output is forced to be a
131
- subset of the input (monotone, drift-free).
135
+ batch unchanged (no data loss); in "unique"/"prune" modes the output is forced
136
+ to be a subset of the input (monotone, drift-free).
132
137
  """
133
138
  items_blob = "; ".join(batch)
134
139
  context = f' about: "{description}"' if description else ""
@@ -150,6 +155,23 @@ def _collapse_batch(client, batch, description, creativity, mode="unique"):
150
155
  "2. Education\n"
151
156
  "3. Religion"
152
157
  )
158
+ elif mode == "prune":
159
+ prompt = (
160
+ f"You are given a list of category labels{context}. Some labels are DUPLICATES — they "
161
+ "name the SAME concept as another label using different words. Return the list with "
162
+ "duplicates removed. Rules: keep ONE label per distinct concept, copied EXACTLY as "
163
+ "written (verbatim); do NOT rename, rephrase, merge, or broaden any label; keep two "
164
+ "labels SEPARATE whenever they name genuinely different concepts, even if related; if "
165
+ "no labels duplicate each other, return them ALL unchanged. "
166
+ f"Labels are separated by semicolons within triple backticks: ```{items_blob}```.\n\n"
167
+ "Return ONLY a numbered list, using the labels exactly as they appear. Each line must "
168
+ "follow this exact format, with no other text before or after the list:\n"
169
+ "N. label\n\n"
170
+ "Example:\n"
171
+ "1. Employment\n"
172
+ "2. Education\n"
173
+ "3. Religion"
174
+ )
153
175
  else:
154
176
  prompt = (
155
177
  f"You are given a list of category labels{context}. "
@@ -186,8 +208,16 @@ def _collapse_batch(client, batch, description, creativity, mode="unique"):
186
208
  if label:
187
209
  out.append(label)
188
210
 
189
- if mode == "unique":
190
- # Contraction guarantee: extract-unique must only REMOVE, never add or
211
+ if not out:
212
+ # Unparseable reply (no numbered list found): keep the batch unchanged
213
+ # rather than silently dropping every label in it. Mirrors the API-error
214
+ # fallback above; matters most in "merge" mode, which has no
215
+ # subset-of-input guarantee to fall back on.
216
+ sys.stderr.write("[collapse_themes] unparseable reply — keeping batch unchanged\n")
217
+ return [str(x).strip().lower() for x in batch]
218
+
219
+ if mode in ("unique", "prune"):
220
+ # Contraction guarantee: extract-unique/prune must only REMOVE, never add or
191
221
  # mutate. Keep only outputs that map back to an input label (by normalized
192
222
  # key), as the original input string. Makes every pass monotone and
193
223
  # drift-free, immune to intermittent model rephrasing/splitting.
@@ -205,6 +235,72 @@ def _collapse_batch(client, batch, description, creativity, mode="unique"):
205
235
  return out
206
236
 
207
237
 
238
+ def _select_top_n(client, counts, n, description, creativity, dedupe_threshold):
239
+ """The "nuclear" final step: one global LLM call that consolidates the whole
240
+ surviving list into at most `n` categories, guided by generation counts
241
+ (higher = the theme recurred more often during extraction). Overlapping
242
+ labels are merged rather than dropped; only rare, unabsorbable themes fall
243
+ away. The result is guaranteed <= n: over-length replies are truncated, and
244
+ a failed or unparseable call falls back to the top-n labels by count
245
+ (deterministic)."""
246
+ ordered = sorted(counts, key=counts.get, reverse=True)
247
+ fallback = [str(x).strip().lower() for x in ordered[:n]]
248
+ blob = "; ".join(f"{lbl} ({counts[lbl]})" for lbl in ordered)
249
+ context = f' about: "{description}"' if description else ""
250
+ prompt = (
251
+ f"You are given category labels{context}. Each label is followed by a count in "
252
+ "parentheses: how many times it was independently generated during extraction — "
253
+ "higher counts mean the theme is more common in the data. Consolidate this list "
254
+ f"into EXACTLY {n} final categories that best summarize the data. Favor "
255
+ "high-count themes; merge overlapping or related labels into one clear "
256
+ "representative label rather than dropping them; drop only rare themes that "
257
+ f"cannot reasonably fold into any of the {n}. Labels are separated by semicolons "
258
+ f"within triple backticks: ```{blob}```\n\n"
259
+ f"Return ONLY a numbered list of exactly {n} category labels, without counts. "
260
+ "Each line must follow this exact format, with no other text before or after "
261
+ "the list:\n"
262
+ "N. label\n\n"
263
+ "Example:\n"
264
+ "1. Employment\n"
265
+ "2. Education\n"
266
+ "3. Religion"
267
+ )
268
+ reply, error = client.complete(
269
+ messages=[{"role": "user", "content": prompt}],
270
+ creativity=creativity,
271
+ force_json=False,
272
+ )
273
+ if error:
274
+ sys.stderr.write(f"[collapse_themes] top_n call failed: {error} — "
275
+ "falling back to top-n by count\n")
276
+ return fallback
277
+ out = []
278
+ for line in (reply or "").splitlines():
279
+ m = _LINE_PAT.match(line.strip())
280
+ if m:
281
+ label = _clean_label(m.group(1)).strip(" ;.,")
282
+ if label:
283
+ out.append(label)
284
+ out = _jw_dedupe(out, dedupe_threshold)
285
+ if not out:
286
+ sys.stderr.write("[collapse_themes] top_n reply unparseable — "
287
+ "falling back to top-n by count\n")
288
+ return fallback
289
+ return out[:n]
290
+
291
+
292
+ def _count_guidance(current, input_data):
293
+ """{label: count} evidence for the top_n prompt: each surviving label gets its
294
+ aggregate generation count from the ORIGINAL input (matched by normalized key)
295
+ when available, else its count in the current list."""
296
+ cur = _to_counts(current)
297
+ orig = {}
298
+ for k, v in _to_counts(input_data).items():
299
+ key = _norm_key(k)
300
+ orig[key] = orig.get(key, 0) + int(v)
301
+ return {lbl: max(int(c), orig.get(_norm_key(lbl), 0)) for lbl, c in cur.items()}
302
+
303
+
208
304
  def _to_counts(input_data):
209
305
  """Coerce the accepted input forms into a {category: count} dict."""
210
306
  if isinstance(input_data, pd.DataFrame):
@@ -288,6 +384,9 @@ def collapse_themes(
288
384
  embedding_merge_threshold=0.92,
289
385
  shuffle=True,
290
386
  final_consolidation=0.82,
387
+ top_n=None,
388
+ prune=False,
389
+ prune_threshold=50,
291
390
  user_model="gpt-4o",
292
391
  model_source="auto",
293
392
  unique_model=None,
@@ -332,7 +431,9 @@ def collapse_themes(
332
431
  which distinctions matter.
333
432
  passes (int | str): Number of collapse iterations, or "auto" to iterate
334
433
  until the deterministic quality benchmark peaks (the recommended mode
335
- for a final taxonomy — pair with aggressive=True). Default 1.
434
+ for a final taxonomy — pair with aggressive=True). Either way,
435
+ iteration stops early once two consecutive passes leave the list
436
+ unchanged. Default 1.
336
437
  max_passes (int): Cap on iterations when passes="auto". Default 10.
337
438
  batch_size (int): Themes per LLM chunk (ceil(n / batch_size) calls per
338
439
  pass). Default 40.
@@ -345,12 +446,29 @@ def collapse_themes(
345
446
  0.92. None or >=1.0 skips embeddings.
346
447
  shuffle (bool): Randomize order each pass so batch composition varies.
347
448
  Default True (improves convergence stability).
449
+ prune (bool): If True, use the prune strategy instead of the merge/unique
450
+ passes: drop only CONCEPTUAL duplicates (one representative per concept,
451
+ kept verbatim), never renaming or merging merely-related labels, so a
452
+ short clean list is left intact while a long list is deduplicated. Lists
453
+ at/below prune_threshold go straight to a single global prune; longer
454
+ lists are reduced by batched prune first, then a final global prune.
455
+ prune_threshold (int): Max list length that goes directly to a single global
456
+ prune call (no batching). Default 50.
348
457
  final_consolidation (float): Cosine threshold for one greedy embedding
349
458
  re-merge over the whole result after all passes, collapsing cross-batch
350
459
  lexical-sibling duplicates that batched passes (and the auto loop) cannot
351
460
  reach. Default 0.82 — deterministic and tuned to land just above the true
352
461
  concept count (errs toward keeping categories; over-segmentation is
353
462
  preferred over over-consolidation). False/None skips.
463
+ top_n (int): Optional "nuclear" final step. After all passes and the final
464
+ consolidation, ONE global LLM call consolidates the surviving list into
465
+ at most `top_n` categories, guided by each label's generation count
466
+ (frequent themes favored; overlapping labels merged into one
467
+ representative rather than dropped). None (default) changes nothing
468
+ about existing behavior. The result is guaranteed to have <= top_n
469
+ labels: a failed or unparseable call falls back deterministically to
470
+ the top_n labels by count. Use when a compact fixed-size taxonomy is
471
+ required and coverage of rare themes is knowingly traded away.
354
472
  user_model (str): Model name for the merge phase. Default "gpt-4o". Use a
355
473
  capable model — small models can degenerate into repetition.
356
474
  model_source (str): Provider — "auto", "openai", "huggingface", etc.
@@ -378,7 +496,7 @@ def collapse_themes(
378
496
  list[str]: The collapsed category list after `passes` iterations.
379
497
 
380
498
  Examples:
381
- >>> import cat_stack as cat
499
+ >>> import catstack as cat
382
500
  >>> themes = cat.explore(df['responses'], description="Why did you move?",
383
501
  ... api_key=key)
384
502
  >>> # Recommended: aggressive merge, auto-stop at the quality peak
@@ -419,6 +537,63 @@ def collapse_themes(
419
537
  def _pass(items, p):
420
538
  return _run(client, items, mode, p)
421
539
 
540
+ # ── Prune strategy ───────────────────────────────────────────────────────
541
+ # Drop CONCEPTUAL duplicates only (keep one representative per concept,
542
+ # verbatim) — never rename or merge merely-related labels, so it cannot
543
+ # over-consolidate a short, already-clean list. Length-routed: a list at or
544
+ # below prune_threshold goes straight to a single global "master list" prune;
545
+ # a longer list is first reduced by the same prune primitive in shuffled
546
+ # batches (until it stops shrinking or fits the threshold), then finished with
547
+ # the global master prune.
548
+ if prune:
549
+ items = list(_to_counts(input_data).keys())
550
+ items = _jw_dedupe(items, dedupe_threshold) # gentle, no rename
551
+ items = _embedding_merge(items, embedding_merge_threshold)
552
+ rng = random.Random(random_state)
553
+ guard = 0
554
+ while len(items) > prune_threshold and guard < max_passes:
555
+ guard += 1
556
+ if shuffle:
557
+ rng.shuffle(items)
558
+ batches = [items[i:i + prune_threshold]
559
+ for i in range(0, len(items), prune_threshold)]
560
+ if max_workers and max_workers > 1 and len(batches) > 1:
561
+ from concurrent.futures import ThreadPoolExecutor, as_completed
562
+
563
+ results = [None] * len(batches)
564
+ with ThreadPoolExecutor(max_workers=max_workers) as ex:
565
+ futures = {
566
+ ex.submit(_collapse_batch, client, b, description,
567
+ creativity, "prune"): i
568
+ for i, b in enumerate(batches)
569
+ }
570
+ for fut in as_completed(futures):
571
+ results[futures[fut]] = fut.result()
572
+ out = [label for r in results for label in (r or [])]
573
+ else:
574
+ out = []
575
+ for b in batches:
576
+ out += _collapse_batch(client, b, description, creativity,
577
+ mode="prune")
578
+ out = _jw_dedupe(out, dedupe_threshold)
579
+ if progress_callback:
580
+ progress_callback(guard, max_passes, "collapse_themes:prune")
581
+ if len(out) >= len(items): # no further reduction this round
582
+ items = out
583
+ break
584
+ items = out
585
+ if len(items) > 1: # final global master-list prune
586
+ items = _jw_dedupe(
587
+ _collapse_batch(client, items, description, creativity, mode="prune"),
588
+ dedupe_threshold)
589
+ if top_n and len(items) > int(top_n):
590
+ items = _select_top_n(client, _count_guidance(items, input_data),
591
+ int(top_n), description, creativity, dedupe_threshold)
592
+ if filename:
593
+ pd.DataFrame({"category": items}).to_csv(filename, index=False)
594
+ print(f"Collapsed categories saved to {filename}")
595
+ return items
596
+
422
597
  current = input_data
423
598
 
424
599
  # Phase 1 (optional): cheap unique-keeping thin. When unique_model is set, run
@@ -442,6 +617,7 @@ def collapse_themes(
442
617
  show_progress_bar=False,
443
618
  )
444
619
  best, best_q = None, -1.0
620
+ prev_key, stable = None, 0
445
621
  for p in range(max_passes):
446
622
  current = _pass(current, p)
447
623
  q = _quality(current, raw_embs)
@@ -450,12 +626,26 @@ def collapse_themes(
450
626
  if q < best_q:
451
627
  break # quality dropped -> the previous pass was the peak
452
628
  best, best_q = current, q
629
+ key = tuple(sorted(current))
630
+ stable = stable + 1 if key == prev_key else 0
631
+ prev_key = key
632
+ if stable >= 2:
633
+ break # converged: two consecutive passes changed nothing
453
634
  current = best if best is not None else current
454
635
  else:
636
+ prev_key, stable = None, 0
455
637
  for p in range(int(passes)):
456
638
  current = _pass(current, p)
457
639
  if progress_callback:
458
640
  progress_callback(p + 1, int(passes), "collapse_themes")
641
+ # Early stop at a fixed point: shuffling gives labels one fresh batch
642
+ # composition to merge under; if two consecutive passes both change
643
+ # nothing, further passes are near-certain no-ops — stop burning calls.
644
+ key = tuple(sorted(current))
645
+ stable = stable + 1 if key == prev_key else 0
646
+ prev_key = key
647
+ if stable >= 2:
648
+ break
459
649
 
460
650
  # Final global consolidation. Batched passes (and the auto loop) can only merge
461
651
  # labels that share a batch, so cross-batch lexical siblings — e.g. "tension" vs
@@ -470,8 +660,20 @@ def collapse_themes(
470
660
  # categories — over-segmentation is the preferred failure mode, not
471
661
  # over-consolidation. Set final_consolidation=False to skip.
472
662
  if final_consolidation and len(current) > 1:
663
+ # Order by original generation frequency first, so the greedy merge keeps
664
+ # the most frequently generated variant as each cluster's representative
665
+ # instead of whichever label the last (shuffled) pass happened to emit
666
+ # first. Labels the merge phase renamed have no original count and sort
667
+ # last, preserving their relative order (sort is stable).
668
+ orig_counts = {_norm_key(k): v for k, v in _to_counts(input_data).items()}
669
+ current = sorted(current, key=lambda c: orig_counts.get(_norm_key(c), 0),
670
+ reverse=True)
473
671
  current = _embedding_merge(current, final_consolidation)
474
672
 
673
+ if top_n and len(_to_counts(current)) > int(top_n):
674
+ current = _select_top_n(client, _count_guidance(current, input_data),
675
+ int(top_n), description, creativity, dedupe_threshold)
676
+
475
677
  if filename:
476
678
  pd.DataFrame({"category": current}).to_csv(filename, index=False)
477
679
  print(f"Collapsed categories saved to {filename}")
@@ -77,7 +77,7 @@ def explore(
77
77
  every iteration. Length ≈ iterations × divisions × categories_per_chunk.
78
78
 
79
79
  Examples:
80
- >>> import cat_stack as cat
80
+ >>> import catstack as cat
81
81
  >>>
82
82
  >>> raw_categories = cat.explore(
83
83
  ... input_data=df['responses'],
@@ -36,6 +36,8 @@ from .pdf_functions import (
36
36
  explore_pdf_categories,
37
37
  )
38
38
 
39
+ from .collapse_themes import collapse_themes
40
+
39
41
 
40
42
  def extract(
41
43
  input_data,
@@ -61,6 +63,9 @@ def extract(
61
63
  auto_download: bool = False,
62
64
  input_mode=None,
63
65
  domain: str = "neutral",
66
+ engine: str = "collapse",
67
+ max_workers: int = 1,
68
+ collapse_kwargs: dict = None,
64
69
  ):
65
70
  """
66
71
  Unified category extraction function for text, image, and PDF inputs.
@@ -111,15 +116,31 @@ def extract(
111
116
  limits. Default 0.0 (no delay).
112
117
  auto_download (bool): If True, automatically download missing Ollama
113
118
  models without prompting. Default False.
119
+ engine (str): Consolidation engine for text input. "collapse" (default):
120
+ run raw extraction (as explore() does) and consolidate the FULL
121
+ inventory with collapse_themes() — semantic pre-clean, quality-
122
+ controlled passes, then a count-guided reduction to at most
123
+ `max_categories`. "legacy": the pre-2.5 single merge call, which
124
+ truncates the inventory to the top max_categories*3 labels by
125
+ exact-string count before merging; kept for reproducing older runs.
126
+ max_workers (int): Parallel API calls for both the extraction chunks and
127
+ the consolidation batches (text engine="collapse" only; extraction
128
+ also honors it under "legacy"). Default 1.
129
+ collapse_kwargs (dict): Optional overrides forwarded to collapse_themes()
130
+ when engine="collapse" — e.g. {"prune": True} or
131
+ {"passes": 2, "aggressive": False}. Defaults applied first:
132
+ passes="auto", aggressive=True. `top_n` cannot be overridden here;
133
+ it is always max_categories.
114
134
 
115
135
  Returns:
116
136
  dict with keys:
117
137
  - counts_df: DataFrame of categories with counts
118
138
  - top_categories: List of top category names
119
- - raw_top_text: Raw model output from final merge step
139
+ - raw_top_text: Raw model output from final merge step ("" when
140
+ engine="collapse", which has no single merge reply)
120
141
 
121
142
  Examples:
122
- >>> import cat_stack as cat
143
+ >>> import catstack as cat
123
144
  >>>
124
145
  >>> # Extract categories from text responses
125
146
  >>> results = cat.extract(
@@ -171,7 +192,41 @@ def extract(
171
192
  resolved_description = description or ""
172
193
 
173
194
  if input_type == "text":
174
- return explore_common_categories(
195
+ if engine not in ("collapse", "legacy"):
196
+ raise ValueError(f"engine must be 'collapse' or 'legacy', got '{engine}'")
197
+
198
+ if engine == "legacy":
199
+ return explore_common_categories(
200
+ input_data=input_data,
201
+ api_key=api_key,
202
+ survey_question=resolved_description,
203
+ max_categories=max_categories,
204
+ categories_per_chunk=categories_per_chunk,
205
+ divisions=divisions,
206
+ user_model=user_model,
207
+ creativity=creativity,
208
+ specificity=specificity,
209
+ research_question=research_question,
210
+ filename=filename,
211
+ model_source=model_source,
212
+ iterations=iterations,
213
+ random_state=random_state,
214
+ focus=focus,
215
+ progress_callback=progress_callback,
216
+ chunk_delay=chunk_delay,
217
+ auto_download=auto_download,
218
+ max_workers=max_workers,
219
+ domain=domain,
220
+ )
221
+
222
+ # engine="collapse": raw extraction (what explore() does), then consolidate
223
+ # the FULL inventory with collapse_themes(). Unlike the legacy merge, no
224
+ # label is truncated away before consolidation, pre-cleaning is semantic
225
+ # (Jaro-Winkler + embeddings) rather than exact-string, and the final
226
+ # count-guided top_n step guarantees at most max_categories categories.
227
+ import pandas as pd
228
+
229
+ raw_items = explore_common_categories(
175
230
  input_data=input_data,
176
231
  api_key=api_key,
177
232
  survey_question=resolved_description,
@@ -182,17 +237,63 @@ def extract(
182
237
  creativity=creativity,
183
238
  specificity=specificity,
184
239
  research_question=research_question,
185
- filename=filename,
240
+ filename=None,
186
241
  model_source=model_source,
187
242
  iterations=iterations,
188
243
  random_state=random_state,
189
244
  focus=focus,
190
245
  progress_callback=progress_callback,
246
+ return_raw=True,
191
247
  chunk_delay=chunk_delay,
192
248
  auto_download=auto_download,
249
+ max_workers=max_workers,
193
250
  domain=domain,
194
251
  )
195
252
 
253
+ # Frequency inventory in the same shape the legacy engine returned.
254
+ def _normalize(cat):
255
+ return "/".join(sorted(t.strip().lower() for t in str(cat).split("/")))
256
+
257
+ flat = [str(x).strip() for x in raw_items if str(x).strip()]
258
+ if not flat:
259
+ raise ValueError("No categories were extracted from the model responses.")
260
+ inv = pd.DataFrame(flat, columns=["Category"])
261
+ inv["normalized"] = inv["Category"].map(_normalize)
262
+ counts_df = (
263
+ inv.groupby("normalized")
264
+ .agg(Category=("Category", lambda x: x.value_counts().index[0]),
265
+ counts=("Category", "size"))
266
+ .sort_values("counts", ascending=False)
267
+ .reset_index(drop=True)
268
+ )
269
+
270
+ ck = dict(passes="auto", aggressive=True)
271
+ ck.update(collapse_kwargs or {})
272
+ ck["top_n"] = int(max_categories) # the required N — not overridable
273
+ top = collapse_themes(
274
+ raw_items,
275
+ api_key=api_key,
276
+ description=resolved_description,
277
+ user_model=user_model,
278
+ model_source=model_source,
279
+ creativity=0 if creativity is None else creativity,
280
+ max_workers=max_workers,
281
+ random_state=random_state,
282
+ progress_callback=progress_callback,
283
+ **ck,
284
+ )
285
+
286
+ if filename:
287
+ pd.DataFrame({"rank": range(1, len(top) + 1), "category": top}).to_csv(
288
+ filename, index=False)
289
+ print(f"Top {len(top)} categories saved to {filename}")
290
+
291
+ return {
292
+ "counts_df": counts_df,
293
+ "top_categories": top,
294
+ "raw_top_text": "",
295
+ }
296
+
196
297
  elif input_type == "image":
197
298
  return explore_image_categories(
198
299
  image_input=input_data,
File without changes
File without changes
File without changes
File without changes