ctxprune 0.1.0__py3-none-any.whl

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.
ctxprune/__init__.py ADDED
@@ -0,0 +1,6 @@
1
+ """ctxprune: context compression for AI agents."""
2
+
3
+ from .compress import Compressor
4
+
5
+ __all__ = ["Compressor"]
6
+ __version__ = "0.1.0"
ctxprune/align.py ADDED
@@ -0,0 +1,93 @@
1
+ """Turn a teacher's deletion-only compression into per-atom keep labels.
2
+
3
+ The teacher is asked to only delete atoms, so its output should be an
4
+ ordered subsequence of the source atoms. We find the longest common
5
+ subsequence (exact match first, casefold second) and label matched source
6
+ atoms as kept. Teacher atoms that match nothing are "variations" (the
7
+ teacher rewrote or invented text); a high variation rate means the sample
8
+ is untrustworthy and gets filtered.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ from dataclasses import dataclass
14
+
15
+ from .atoms import Atom, atomize, is_protected, norm
16
+
17
+
18
+ @dataclass
19
+ class Alignment:
20
+ keep: list[bool]
21
+ n_src: int
22
+ n_comp: int
23
+ n_matched: int
24
+ variation_rate: float # teacher atoms with no source match
25
+ comp_rate: float # kept / source, in atoms
26
+ protected_total: int
27
+ protected_kept: int
28
+
29
+
30
+ def _lcs_pairs(a: list[str], b: list[str]) -> list[tuple[int, int]]:
31
+ n, m = len(a), len(b)
32
+ if n == 0 or m == 0:
33
+ return []
34
+ # dp[i][j] = LCS length of a[i:], b[j:]; rows built bottom-up.
35
+ dp = [[0] * (m + 1) for _ in range(n + 1)]
36
+ for i in range(n - 1, -1, -1):
37
+ ai, row, nxt = a[i], dp[i], dp[i + 1]
38
+ for j in range(m - 1, -1, -1):
39
+ if ai == b[j]:
40
+ row[j] = nxt[j + 1] + 1
41
+ else:
42
+ x, y = nxt[j], row[j + 1]
43
+ row[j] = x if x >= y else y
44
+ pairs, i, j = [], 0, 0
45
+ while i < n and j < m:
46
+ if a[i] == b[j]:
47
+ pairs.append((i, j))
48
+ i += 1
49
+ j += 1
50
+ elif dp[i + 1][j] >= dp[i][j + 1]:
51
+ i += 1
52
+ else:
53
+ j += 1
54
+ return pairs
55
+
56
+
57
+ def align(src_atoms: list[Atom], comp_text: str) -> Alignment:
58
+ src = [a.text for a in src_atoms]
59
+ comp = [a.text for a in atomize(comp_text) if a.text]
60
+ keep = [False] * len(src)
61
+ matched_comp = [False] * len(comp)
62
+
63
+ # Pass 1: exact. Pass 2: casefold over what is left, in the gaps between
64
+ # pass-1 anchors so order is preserved.
65
+ pairs = _lcs_pairs(src, comp)
66
+ for i, j in pairs:
67
+ keep[i] = True
68
+ matched_comp[j] = True
69
+ anchors = [(-1, -1)] + pairs + [(len(src), len(comp))]
70
+ for (i0, j0), (i1, j1) in zip(anchors, anchors[1:]):
71
+ si = [i for i in range(i0 + 1, i1) if not keep[i]]
72
+ cj = [j for j in range(j0 + 1, j1) if not matched_comp[j]]
73
+ if not si or not cj:
74
+ continue
75
+ sub = _lcs_pairs([norm(src[i]) for i in si], [norm(comp[j]) for j in cj])
76
+ for a, b in sub:
77
+ keep[si[a]] = True
78
+ matched_comp[cj[b]] = True
79
+
80
+ real = [i for i, t in enumerate(src) if t]
81
+ n_src = len(real)
82
+ n_matched = sum(matched_comp)
83
+ prot = [i for i in real if is_protected(src[i])]
84
+ return Alignment(
85
+ keep=keep,
86
+ n_src=n_src,
87
+ n_comp=len(comp),
88
+ n_matched=n_matched,
89
+ variation_rate=(1 - n_matched / len(comp)) if comp else 0.0,
90
+ comp_rate=(sum(keep[i] for i in real) / n_src) if n_src else 0.0,
91
+ protected_total=len(prot),
92
+ protected_kept=sum(keep[i] for i in prot),
93
+ )
ctxprune/atoms.py ADDED
@@ -0,0 +1,106 @@
1
+ """Atom segmentation: the unit the compressor keeps or drops.
2
+
3
+ An atom is either a "word" that glues identifier-ish characters together
4
+ (paths, URLs, versions, IPs, snake_case, UUIDs) or a single punctuation
5
+ character. Each atom remembers the whitespace that preceded it, so
6
+ ``detok(atoms, keep)`` rebuilds kept text exactly as it appeared in the
7
+ source, without the "0. 21" / "ord _ 8f3a" artifacts LLMLingua-2 produces.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import re
13
+ from dataclasses import dataclass
14
+
15
+ # Scripts written without spaces (Han, kana) get one atom per character, so
16
+ # a deletion inside a sentence is still a deletion and not a "rewrite".
17
+ _CJK = "\u3040-\u30ff\u3400-\u4dbf\u4e00-\u9fff\uf900-\ufaff"
18
+ _W = rf"[^\W{_CJK}]" # word char that is not CJK
19
+ # A word may start with one path/flag sigil, then word chars, with
20
+ # identifier punctuation allowed only *between* word chars.
21
+ _WORD = rf"(?:--?|[/~.$@#])?{_W}(?:(?:{_W}|[.:/@%+#~=?&'\-])*{_W})?"
22
+ _ATOM_RE = re.compile(rf"(\s*)({_WORD}|[{_CJK}]|[^\s])", re.UNICODE)
23
+
24
+ # All-caps English that is prose, not an identifier (license headers, shouting).
25
+ _CAPS_WORDS = set(
26
+ "THE AND FOR ANY NOT BUT ARE WAS YOU ALL WITH FROM THIS THAT WITHOUT WARRANTIES WARRANTY SOFTWARE KIND "
27
+ "IMPLIED INCLUDING LIMITED COPYRIGHT OTHER LIABILITY USE BASIS CONDITIONS SUCH EVENT SHALL HOLDERS "
28
+ "CONTRIBUTORS DAMAGES PURPOSE PARTICULAR FITNESS MERCHANTABILITY EXPRESS OUT CONNECTION ARISING "
29
+ "WHETHER CONTRACT TORT OTHERWISE PROVIDED LICENSE NOTICE NOTE IMPORTANT".split()
30
+ )
31
+ _HEX_RE = re.compile(r"^[0-9a-fA-F]{7,}$")
32
+ _CAMEL_RE = re.compile(r"[a-z][A-Z]|[A-Z]{2}[a-z]")
33
+
34
+
35
+ @dataclass(frozen=True)
36
+ class Atom:
37
+ ws: str # whitespace preceding the atom in the source
38
+ text: str
39
+
40
+
41
+ def atomize(text: str) -> list[Atom]:
42
+ atoms: list[Atom] = []
43
+ pos = 0
44
+ for m in _ATOM_RE.finditer(text):
45
+ if m.start() != pos: # unreachable by construction; guard anyway
46
+ raise ValueError(f"atomizer skipped {text[pos:m.start()]!r}")
47
+ atoms.append(Atom(m.group(1), m.group(2)))
48
+ pos = m.end()
49
+ tail = text[pos:]
50
+ if tail.strip():
51
+ raise ValueError(f"atomizer left non-space tail {tail!r}")
52
+ if tail: # trailing whitespace rides on a sentinel-free empty atom
53
+ atoms.append(Atom(tail, ""))
54
+ return atoms
55
+
56
+
57
+ def detok(atoms: list[Atom], keep: list[bool], collapse_ws: bool = True) -> str:
58
+ """Rebuild text from kept atoms.
59
+
60
+ Kept atoms keep their own leading whitespace. When an atom's predecessor
61
+ was dropped, the dropped run's whitespace is collapsed to the "largest"
62
+ whitespace seen in that run (a newline wins over a space), so line
63
+ structure survives deletions.
64
+ """
65
+ out: list[str] = []
66
+ pending = ""
67
+ for a, k in zip(atoms, keep):
68
+ if k:
69
+ ws = a.ws
70
+ if pending and collapse_ws:
71
+ ws = _merge_ws(pending, a.ws)
72
+ if not out:
73
+ ws = ws if not collapse_ws else ""
74
+ out.append(ws + a.text)
75
+ pending = ""
76
+ else:
77
+ pending = _merge_ws(pending, a.ws) if pending else (a.ws or " ")
78
+ return "".join(out)
79
+
80
+
81
+ def _merge_ws(a: str, b: str) -> str:
82
+ if "\n" in a or "\n" in b:
83
+ return "\n"
84
+ return " " if (a or b) else ""
85
+
86
+
87
+ def is_protected(atom: str) -> bool:
88
+ """Atoms an agent is likely to need verbatim later (IDs, numbers, paths)."""
89
+ if not atom or not any(c.isalnum() for c in atom):
90
+ return False
91
+ if any(c.isdigit() for c in atom):
92
+ return True
93
+ core = atom.lstrip("-/~.$@#")
94
+ if any(c in core for c in "_/.:@="):
95
+ return True
96
+ if _HEX_RE.match(core):
97
+ return True
98
+ if len(core) >= 3 and core.isupper() and core not in _CAPS_WORDS:
99
+ return True
100
+ if _CAMEL_RE.search(core):
101
+ return True
102
+ return atom != core # flags like --format, paths like /tmp
103
+
104
+
105
+ def norm(atom: str) -> str:
106
+ return atom.casefold()
ctxprune/chunking.py ADDED
@@ -0,0 +1,104 @@
1
+ """Split documents into chunks that fit the compressor's 512-token window.
2
+
3
+ Chunks break on line boundaries where possible; a single line longer than
4
+ the budget is split between atoms. Token counts use the target encoder's
5
+ tokenizer, so every training chunk fits one forward pass, matching how
6
+ llmlingua feeds the model (``max_seq_len = 512`` minus CLS/SEP).
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import hashlib
12
+ import random
13
+ from collections.abc import Callable
14
+
15
+ from .atoms import atomize
16
+
17
+ DEFAULT_TOKENIZER = "jhu-clsp/mmBERT-small"
18
+
19
+
20
+ def load_counter(name: str = DEFAULT_TOKENIZER) -> Callable[[list[str]], list[int]]:
21
+ from transformers import AutoTokenizer
22
+
23
+ tok = AutoTokenizer.from_pretrained(name)
24
+
25
+ def count(texts: list[str]) -> list[int]:
26
+ return [len(ids) for ids in tok(texts, add_special_tokens=False)["input_ids"]]
27
+
28
+ return count
29
+
30
+
31
+ def _split_long_line(line: str, budget: int, count) -> list[str]:
32
+ atoms = atomize(line)
33
+ pieces = [a.ws + a.text for a in atoms]
34
+ sizes = count(pieces)
35
+ out, cur, cur_n = [], [], 0
36
+ for p, n in zip(pieces, sizes):
37
+ if cur and cur_n + n > budget:
38
+ out.append("".join(cur))
39
+ cur, cur_n = [], 0
40
+ cur.append(p)
41
+ cur_n += n
42
+ if cur:
43
+ out.append("".join(cur))
44
+ return out
45
+
46
+
47
+ def chunk_text(text: str, count, max_tokens: int = 500, min_tokens: int = 48) -> list[str]:
48
+ lines = text.splitlines(keepends=True)
49
+ if not lines:
50
+ return []
51
+ sizes = count(lines)
52
+ units: list[tuple[str, int]] = []
53
+ for line, n in zip(lines, sizes):
54
+ if n <= max_tokens:
55
+ units.append((line, n))
56
+ else:
57
+ for piece in _split_long_line(line, max_tokens, count):
58
+ units.append((piece, count([piece])[0]))
59
+ chunks, cur, cur_n = [], [], 0
60
+ for u, n in units:
61
+ if cur and cur_n + n > max_tokens:
62
+ chunks.append("".join(cur))
63
+ cur, cur_n = [], 0
64
+ cur.append(u)
65
+ cur_n += n
66
+ if cur:
67
+ chunks.append("".join(cur))
68
+ # Line-sum can undercount at joins; re-check and drop tiny tails.
69
+ final = []
70
+ for c, n in zip(chunks, count(chunks)):
71
+ if n < min_tokens or not c.strip():
72
+ continue
73
+ if n > max_tokens:
74
+ final.extend(p for p in _split_long_line(c, max_tokens, count) if p.strip())
75
+ else:
76
+ final.append(c)
77
+ return final
78
+
79
+
80
+ def split_of(group: str, val_pct: int = 2, test_pct: int = 2) -> str:
81
+ b = int(hashlib.sha1(group.encode()).hexdigest(), 16) % 100
82
+ if b < test_pct:
83
+ return "test"
84
+ if b < test_pct + val_pct:
85
+ return "val"
86
+ return "train"
87
+
88
+
89
+ def chunk_doc(doc: dict, count, rng: random.Random, max_tokens: int = 500,
90
+ max_chunks_per_doc: int = 2) -> list[dict]:
91
+ chunks = chunk_text(doc["text"], count, max_tokens=max_tokens)
92
+ if len(chunks) > max_chunks_per_doc:
93
+ idx = sorted(rng.sample(range(len(chunks)), max_chunks_per_doc))
94
+ chunks = [chunks[i] for i in idx]
95
+ out = []
96
+ for c in chunks:
97
+ meta = {k: v for k, v in doc.items() if k != "text"}
98
+ out.append({
99
+ "id": f"{doc['source']}:{hashlib.sha1(c.encode()).hexdigest()[:16]}",
100
+ "text": c,
101
+ "split": split_of(doc["group"]),
102
+ **meta,
103
+ })
104
+ return out
ctxprune/cli.py ADDED
@@ -0,0 +1,40 @@
1
+ """ctxprune command line: compress a file (or stdin) and print the result.
2
+
3
+ ctxprune order.json --rate 0.5
4
+ cat build.log | ctxprune --rate 0.33 --force-protected
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import argparse
10
+ import sys
11
+
12
+ DEFAULT_MODEL = "darioooooo0o/ctxprune-small"
13
+
14
+
15
+ def main(argv: list[str] | None = None) -> None:
16
+ ap = argparse.ArgumentParser(prog="ctxprune", description="Compress text for an LLM by deleting low-value words.")
17
+ ap.add_argument("file", nargs="?", help="input file (default: stdin)")
18
+ ap.add_argument("--rate", type=float, default=0.5, help="fraction of the text to keep (default 0.5)")
19
+ ap.add_argument("--force-protected", action="store_true", help="never drop identifiers, numbers or paths")
20
+ ap.add_argument("--model", default=DEFAULT_MODEL)
21
+ ap.add_argument("--backend", choices=["onnx", "torch"], default=None,
22
+ help="default: onnx if onnxruntime is installed, else torch")
23
+ args = ap.parse_args(argv)
24
+
25
+ text = open(args.file, encoding="utf-8").read() if args.file else sys.stdin.read()
26
+ backend = args.backend
27
+ if backend is None:
28
+ try:
29
+ import onnxruntime # noqa: F401
30
+ backend = "onnx"
31
+ except ImportError:
32
+ backend = "torch"
33
+ from .compress import Compressor
34
+
35
+ c = Compressor(args.model, backend=backend)
36
+ sys.stdout.write(c.compress(text, rate=args.rate, force_protected=args.force_protected)["text"] + "\n")
37
+
38
+
39
+ if __name__ == "__main__":
40
+ main()
ctxprune/compress.py ADDED
@@ -0,0 +1,182 @@
1
+ """Inference: score atoms with the trained token classifier and keep the best.
2
+
3
+ from ctxprune.compress import Compressor
4
+ c = Compressor("path/or/hub-id")
5
+ c.compress(text, rate=0.5)["text"]
6
+
7
+ `rate` is the fraction of tokens to keep (like LLMLingua-2's `rate`).
8
+ Identifiers are never split or re-spaced, because selection happens on whole
9
+ atoms and text is rebuilt from the source's own whitespace.
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ import math
15
+
16
+ from .atoms import atomize, detok, is_protected
17
+ from .tokenlabels import atom_spans, char_to_atom, special_wrap, token_atom_sets, token_atoms
18
+
19
+ STRUCTURE = set('{}[]",:')
20
+
21
+
22
+ class Compressor:
23
+ def __init__(self, model: str, device: str | None = None, max_len: int = 512, batch_size: int = 16,
24
+ backend: str = "torch", onnx_file: str = "onnx/model.onnx"):
25
+ """backend="torch" (any device) or "onnx" (CPU, no torch needed; onnx_file may be
26
+ "onnx/model_quantized.onnx" for the int8 build)."""
27
+ if backend == "onnx":
28
+ import os
29
+
30
+ os.environ.setdefault("TRANSFORMERS_VERBOSITY", "error") # silence "PyTorch was not found"
31
+ from transformers import AutoTokenizer
32
+
33
+ self.tok = AutoTokenizer.from_pretrained(model)
34
+ self.backend = backend
35
+ if backend == "onnx":
36
+ import os
37
+
38
+ import onnxruntime as ort
39
+ from huggingface_hub import hf_hub_download
40
+
41
+ path = os.path.join(model, onnx_file) if os.path.isdir(model) else hf_hub_download(model, onnx_file)
42
+ self.session = ort.InferenceSession(path, providers=["CPUExecutionProvider"])
43
+ self.device = "cpu"
44
+ else:
45
+ import torch
46
+ from transformers import AutoModelForTokenClassification
47
+
48
+ self.torch = torch
49
+ if device is None:
50
+ device = "xpu" if hasattr(torch, "xpu") and torch.xpu.is_available() else (
51
+ "cuda" if torch.cuda.is_available() else "cpu")
52
+ self.device = device
53
+ self.model = AutoModelForTokenClassification.from_pretrained(model).to(device).eval()
54
+ self.max_len = max_len
55
+ self.batch_size = batch_size
56
+ self.pre, self.post = special_wrap(self.tok)
57
+
58
+ def atom_scores(self, text: str) -> tuple[list, list[float], list[int]]:
59
+ """Keep-probability and token count for every atom of `text`."""
60
+ atoms = atomize(text)
61
+ full, spans = atom_spans(atoms)
62
+ owner = char_to_atom(spans, len(full))
63
+ enc = self.tok(full, add_special_tokens=False, return_offsets_mapping=True)
64
+ ids, offs = enc["input_ids"], enc["offset_mapping"]
65
+ tok_atom = token_atoms(offs, owner)
66
+
67
+ # Windows of max_len - 2 tokens, cut where a new atom starts.
68
+ win = self.max_len - len(self.pre) - len(self.post)
69
+ windows, start = [], 0
70
+ while start < len(ids):
71
+ end = min(start + win, len(ids))
72
+ if end < len(ids):
73
+ cut = end
74
+ while cut > start + 1 and tok_atom[cut] == tok_atom[cut - 1]:
75
+ cut -= 1
76
+ end = cut if cut > start + 1 else end
77
+ windows.append((start, end))
78
+ start = end
79
+
80
+ probs = [0.0] * len(ids)
81
+ for b in range(0, len(windows), self.batch_size):
82
+ batch = windows[b:b + self.batch_size]
83
+ seqs = [self.pre + ids[s:e] + self.post for s, e in batch]
84
+ width = max(map(len, seqs))
85
+ pad = self.tok.pad_token_id
86
+ ids_rows = [q + [pad] * (width - len(q)) for q in seqs]
87
+ mask_rows = [[1] * len(q) + [0] * (width - len(q)) for q in seqs]
88
+ p = self._keep_probs(ids_rows, mask_rows)
89
+ for (s, e), row in zip(batch, p):
90
+ k = len(self.pre)
91
+ probs[s:e] = row[k:k + (e - s)]
92
+
93
+ # Score: mean prob over every token touching the atom. Budget: each token is
94
+ # charged once, to the first atom it touches.
95
+ sums = [0.0] * len(atoms)
96
+ hits = [0] * len(atoms)
97
+ counts = [0] * len(atoms)
98
+ for a, sets, p in zip(tok_atom, token_atom_sets(offs, owner), probs):
99
+ if a >= 0:
100
+ counts[a] += 1
101
+ for b in sets:
102
+ sums[b] += p
103
+ hits[b] += 1
104
+ scores = [sums[i] / hits[i] if hits[i] else 0.0 for i in range(len(atoms))]
105
+ return atoms, scores, counts
106
+
107
+ def _keep_probs(self, ids_rows: list[list[int]], mask_rows: list[list[int]]) -> list[list[float]]:
108
+ if self.backend == "onnx":
109
+ import numpy as np
110
+
111
+ logits = self.session.run(["logits"], {"input_ids": np.array(ids_rows, dtype=np.int64),
112
+ "attention_mask": np.array(mask_rows, dtype=np.int64)})[0]
113
+ z = logits - logits.max(-1, keepdims=True)
114
+ e = np.exp(z)
115
+ return (e[..., 1] / e.sum(-1)).tolist()
116
+ t = self.torch
117
+ input_ids = t.tensor(ids_rows, device=self.device)
118
+ mask = t.tensor(mask_rows, device=self.device)
119
+ with t.no_grad():
120
+ logits = self.model(input_ids=input_ids, attention_mask=mask).logits.float()
121
+ return logits.softmax(-1)[..., 1].cpu().tolist()
122
+
123
+ def compress(self, text: str, rate: float | None = 0.5, threshold: float | None = None,
124
+ force_protected: bool = False, keep_structure: bool = False, unit="chars",
125
+ measure=None) -> dict:
126
+ """Keep the highest-scoring atoms until `rate` of the text is kept.
127
+
128
+ unit="chars" measures the budget in characters (one per whitespace gap), which
129
+ tracks a downstream LLM's token count without depending on our tokenizer: the
130
+ encoder splits digits one by one, so in "tokens" identifiers look expensive.
131
+ unit may also be a callable mapping a list of strings to their costs.
132
+ measure: optional callable str -> int (e.g. the target LLM's token count). The
133
+ cut is then binary-searched so measure(output) <= rate * measure(text) exactly,
134
+ which makes `rate` comparable across compressors with different tokenizers.
135
+ """
136
+ atoms, scores, tok_counts = self.atom_scores(text)
137
+ if callable(unit):
138
+ counts = list(unit([(" " if a.ws else "") + a.text for a in atoms]))
139
+ elif unit == "chars":
140
+ counts = [len(a.text) + (1 if a.ws else 0) for a in atoms]
141
+ elif unit == "tokens":
142
+ counts = tok_counts
143
+ else:
144
+ raise ValueError(f"unit must be 'chars' or 'tokens', not {unit!r}")
145
+ forced = [
146
+ bool(a.text) and ((force_protected and is_protected(a.text)) or (keep_structure and a.text in STRUCTURE))
147
+ for a in atoms
148
+ ]
149
+ if threshold is not None:
150
+ keep = [f or s >= threshold for f, s in zip(forced, scores)]
151
+ elif measure is not None and rate is not None:
152
+ order = [i for i in sorted(range(len(atoms)), key=lambda i: -scores[i])
153
+ if atoms[i].text and not forced[i]]
154
+ target = rate * measure(text)
155
+
156
+ def with_top(k: int) -> list[bool]:
157
+ kk = list(forced)
158
+ for i in order[:k]:
159
+ kk[i] = True
160
+ return kk
161
+
162
+ lo, hi = 0, len(order) # largest k whose output fits the target
163
+ while lo < hi:
164
+ mid = (lo + hi + 1) // 2
165
+ if measure(detok(atoms, with_top(mid))) <= target:
166
+ lo = mid
167
+ else:
168
+ hi = mid - 1
169
+ keep = with_top(lo)
170
+ else:
171
+ total = sum(counts)
172
+ budget = math.ceil((rate if rate is not None else 1.0) * total)
173
+ keep = list(forced)
174
+ used = sum(c for c, f in zip(counts, forced) if f)
175
+ for i in sorted(range(len(atoms)), key=lambda i: -scores[i]):
176
+ if used >= budget:
177
+ break
178
+ if not keep[i] and atoms[i].text:
179
+ keep[i] = True
180
+ used += counts[i]
181
+ out = detok(atoms, keep)
182
+ return {"text": out, "kept": sum(c for c, k in zip(counts, keep) if k), "total": sum(counts)}
ctxprune/idcheck.py ADDED
@@ -0,0 +1,59 @@
1
+ """Identifier-survival metric: did the IDs, numbers and paths make it through?
2
+
3
+ A protected atom survives only if it appears intact as an atom of the
4
+ compressed text, so "0.21" -> "0. 21" counts as lost: that is what the
5
+ downstream LLM actually sees.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import json
11
+
12
+ from .atoms import atomize, is_protected
13
+
14
+
15
+ def protected_atoms(text: str) -> set[str]:
16
+ return {a.text for a in atomize(text) if is_protected(a.text)}
17
+
18
+
19
+ def survival(original: str, compressed: str, count_tokens=None) -> dict:
20
+ want = protected_atoms(original)
21
+ have = {a.text for a in atomize(compressed)}
22
+ lost = sorted(want - have)
23
+ r = {
24
+ "protected": len(want),
25
+ "lost": len(lost),
26
+ "lost_examples": lost[:10],
27
+ "char_ratio": len(compressed) / max(1, len(original)),
28
+ }
29
+ if count_tokens is not None:
30
+ r["llm_token_ratio"] = count_tokens(compressed) / max(1, count_tokens(original))
31
+ try:
32
+ json.loads(original)
33
+ except (json.JSONDecodeError, ValueError):
34
+ r["json"] = None
35
+ else:
36
+ try:
37
+ json.loads(compressed)
38
+ r["json"] = True
39
+ except (json.JSONDecodeError, ValueError):
40
+ r["json"] = False
41
+ return r
42
+
43
+
44
+ def summarize(rows: list[dict]) -> dict:
45
+ prot = sum(r["protected"] for r in rows)
46
+ lost = sum(r["lost"] for r in rows)
47
+ js = [r["json"] for r in rows if r.get("json") is not None]
48
+ summary = {
49
+ "n": len(rows),
50
+ "protected": prot,
51
+ "lost": lost,
52
+ "lost_rate": lost / prot if prot else 0.0,
53
+ "json_inputs": len(js),
54
+ "json_still_valid": sum(js),
55
+ "mean_char_ratio": sum(r["char_ratio"] for r in rows) / len(rows) if rows else 0.0,
56
+ }
57
+ if rows and "llm_token_ratio" in rows[0]:
58
+ summary["mean_llm_token_ratio"] = sum(r["llm_token_ratio"] for r in rows) / len(rows)
59
+ return summary
@@ -0,0 +1,90 @@
1
+ """Drop-in compressors for LangChain and LlamaIndex retrieval pipelines.
2
+
3
+ from ctxprune.integrations import CtxpruneDocumentCompressor # LangChain
4
+ retriever = ContextualCompressionRetriever(base_compressor=CtxpruneDocumentCompressor(rate=0.5),
5
+ base_retriever=base)
6
+
7
+ from ctxprune.integrations import CtxpruneNodePostprocessor # LlamaIndex
8
+ engine = index.as_query_engine(node_postprocessors=[CtxpruneNodePostprocessor(rate=0.5)])
9
+
10
+ Both compress each document's text in place and leave metadata alone. Classes are built
11
+ on first access so neither framework is required to import ctxprune.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ from functools import lru_cache
17
+
18
+ from .cli import DEFAULT_MODEL
19
+
20
+
21
+ @lru_cache(maxsize=4)
22
+ def _compressor(model: str, backend: str | None):
23
+ from .compress import Compressor
24
+
25
+ if backend is None:
26
+ try:
27
+ import onnxruntime # noqa: F401
28
+ backend = "onnx"
29
+ except ImportError:
30
+ backend = "torch"
31
+ return Compressor(model, backend=backend)
32
+
33
+
34
+ def _langchain():
35
+ from langchain_core.documents import Document
36
+ from langchain_core.documents.compressor import BaseDocumentCompressor
37
+
38
+ class CtxpruneDocumentCompressor(BaseDocumentCompressor):
39
+ """LangChain document compressor: keeps `rate` of each document's text."""
40
+
41
+ model: str = DEFAULT_MODEL
42
+ backend: str | None = None
43
+ rate: float = 0.5
44
+ force_protected: bool = False
45
+
46
+ def compress_documents(self, documents, query, callbacks=None):
47
+ c = _compressor(self.model, self.backend)
48
+ return [Document(page_content=c.compress(d.page_content, rate=self.rate,
49
+ force_protected=self.force_protected)["text"],
50
+ metadata=d.metadata, id=d.id)
51
+ for d in documents]
52
+
53
+ return CtxpruneDocumentCompressor
54
+
55
+
56
+ def _llamaindex():
57
+ from llama_index.core.postprocessor.types import BaseNodePostprocessor
58
+
59
+ class CtxpruneNodePostprocessor(BaseNodePostprocessor):
60
+ """LlamaIndex node postprocessor: keeps `rate` of each retrieved node's text."""
61
+
62
+ model: str = DEFAULT_MODEL
63
+ backend: str | None = None
64
+ rate: float = 0.5
65
+ force_protected: bool = False
66
+
67
+ @classmethod
68
+ def class_name(cls) -> str:
69
+ return "CtxpruneNodePostprocessor"
70
+
71
+ def _postprocess_nodes(self, nodes, query_bundle=None):
72
+ c = _compressor(self.model, self.backend)
73
+ for n in nodes:
74
+ n.node.set_content(c.compress(n.node.get_content(), rate=self.rate,
75
+ force_protected=self.force_protected)["text"])
76
+ return nodes
77
+
78
+ return CtxpruneNodePostprocessor
79
+
80
+
81
+ _FACTORIES = {"CtxpruneDocumentCompressor": _langchain, "CtxpruneNodePostprocessor": _llamaindex}
82
+ _built: dict[str, type] = {}
83
+
84
+
85
+ def __getattr__(name: str):
86
+ if name not in _FACTORIES:
87
+ raise AttributeError(name)
88
+ if name not in _built:
89
+ _built[name] = _FACTORIES[name]()
90
+ return _built[name]