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 +6 -0
- ctxprune/align.py +93 -0
- ctxprune/atoms.py +106 -0
- ctxprune/chunking.py +104 -0
- ctxprune/cli.py +40 -0
- ctxprune/compress.py +182 -0
- ctxprune/idcheck.py +59 -0
- ctxprune/integrations.py +90 -0
- ctxprune/sources.py +256 -0
- ctxprune/teacher.py +241 -0
- ctxprune/tokenlabels.py +71 -0
- ctxprune-0.1.0.dist-info/METADATA +143 -0
- ctxprune-0.1.0.dist-info/RECORD +16 -0
- ctxprune-0.1.0.dist-info/WHEEL +4 -0
- ctxprune-0.1.0.dist-info/entry_points.txt +2 -0
- ctxprune-0.1.0.dist-info/licenses/LICENSE +202 -0
ctxprune/__init__.py
ADDED
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
|
ctxprune/integrations.py
ADDED
|
@@ -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]
|