quire-grammar 0.2.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.
- quire_grammar/__init__.py +17 -0
- quire_grammar/__main__.py +57 -0
- quire_grammar/check.py +56 -0
- quire_grammar/edits.py +241 -0
- quire_grammar/model.py +245 -0
- quire_grammar/py.typed +0 -0
- quire_grammar/registry.py +45 -0
- quire_grammar/rules/__init__.py +11 -0
- quire_grammar/rules/agreement.py +99 -0
- quire_grammar/rules/confusables.py +200 -0
- quire_grammar/rules/mechanics.py +142 -0
- quire_grammar/rules/punctuation.py +130 -0
- quire_grammar/rules/style.py +144 -0
- quire_grammar/segment.py +35 -0
- quire_grammar/types.py +41 -0
- quire_grammar-0.2.0.dist-info/METADATA +121 -0
- quire_grammar-0.2.0.dist-info/RECORD +21 -0
- quire_grammar-0.2.0.dist-info/WHEEL +5 -0
- quire_grammar-0.2.0.dist-info/entry_points.txt +2 -0
- quire_grammar-0.2.0.dist-info/licenses/LICENSE +21 -0
- quire_grammar-0.2.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
"""quire-grammar — an offline English grammar and style checker.
|
|
2
|
+
|
|
3
|
+
Public surface (stable):
|
|
4
|
+
|
|
5
|
+
check(text, *, enable=None, disable=None, min_confidence=0.0) -> list[Correction]
|
|
6
|
+
|
|
7
|
+
Everything else is an implementation detail. The phase-2 local model plugs in
|
|
8
|
+
behind ``check`` without changing this signature.
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
from __future__ import annotations
|
|
12
|
+
|
|
13
|
+
from .check import check
|
|
14
|
+
from .types import CATEGORIES, Correction
|
|
15
|
+
|
|
16
|
+
__version__ = "0.0.1"
|
|
17
|
+
__all__ = ["CATEGORIES", "Correction", "__version__", "check"]
|
|
@@ -0,0 +1,57 @@
|
|
|
1
|
+
"""quire-grammar CLI:
|
|
2
|
+
|
|
3
|
+
quire-grammar "Its a nice day and and the the sky is is blue."
|
|
4
|
+
echo "text" | quire-grammar -
|
|
5
|
+
quire-grammar --json chapter-3.txt
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import argparse
|
|
11
|
+
import json
|
|
12
|
+
import sys
|
|
13
|
+
|
|
14
|
+
from . import __version__, check
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def _line_col(text: str, pos: int) -> tuple[int, int]:
|
|
18
|
+
head = text[:pos]
|
|
19
|
+
return head.count("\n") + 1, pos - (head.rfind("\n") + 1) + 1
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def main(argv: list[str] | None = None) -> int:
|
|
23
|
+
ap = argparse.ArgumentParser(prog="quire-grammar", description=__doc__)
|
|
24
|
+
ap.add_argument("source", help='text, a file path, or "-" for stdin')
|
|
25
|
+
ap.add_argument("--json", action="store_true", help="machine-readable output")
|
|
26
|
+
ap.add_argument("--min-confidence", type=float, default=0.0)
|
|
27
|
+
ap.add_argument("--disable", default="", help="comma-separated rule ids / categories")
|
|
28
|
+
ap.add_argument("--version", action="version", version=f"%(prog)s {__version__}")
|
|
29
|
+
args = ap.parse_args(argv)
|
|
30
|
+
|
|
31
|
+
if args.source == "-":
|
|
32
|
+
text = sys.stdin.read()
|
|
33
|
+
else:
|
|
34
|
+
try:
|
|
35
|
+
with open(args.source, encoding="utf-8") as f:
|
|
36
|
+
text = f.read()
|
|
37
|
+
except OSError:
|
|
38
|
+
text = args.source
|
|
39
|
+
|
|
40
|
+
hits = check(text, disable=[s for s in args.disable.split(",") if s],
|
|
41
|
+
min_confidence=args.min_confidence)
|
|
42
|
+
|
|
43
|
+
if args.json:
|
|
44
|
+
print(json.dumps([vars(h) | {"meta": h.meta} for h in hits], indent=2))
|
|
45
|
+
return 1 if hits else 0
|
|
46
|
+
|
|
47
|
+
for h in hits:
|
|
48
|
+
line, col = _line_col(text, h.start)
|
|
49
|
+
fix = f" → {h.suggestions[0]!r}" if h.suggestions else ""
|
|
50
|
+
print(f"{line}:{col} [{h.category}/{h.rule_id}] {h.message}{fix}")
|
|
51
|
+
if not hits:
|
|
52
|
+
print("no issues found")
|
|
53
|
+
return 1 if hits else 0
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
if __name__ == "__main__":
|
|
57
|
+
raise SystemExit(main())
|
quire_grammar/check.py
ADDED
|
@@ -0,0 +1,56 @@
|
|
|
1
|
+
"""The one call the rest of the world makes.
|
|
2
|
+
|
|
3
|
+
from quire_grammar import check
|
|
4
|
+
for c in check("Its a nice day and and the the sky is is blue."):
|
|
5
|
+
print(c.rule_id, c.original, "->", c.suggestions)
|
|
6
|
+
|
|
7
|
+
``enable`` / ``disable`` accept rule ids *or* category names, so a UI can
|
|
8
|
+
switch off "style" wholesale or turn on a single opt-in rule.
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
from __future__ import annotations
|
|
12
|
+
|
|
13
|
+
from collections.abc import Callable, Iterable
|
|
14
|
+
|
|
15
|
+
from . import rules as _rules # noqa: F401 — importing registers every rule
|
|
16
|
+
from .registry import all_rules
|
|
17
|
+
from .types import Correction
|
|
18
|
+
|
|
19
|
+
# Phase 2: when the optional model is installed, its predictions are merged
|
|
20
|
+
# in here (same Correction shape, category "agreement" etc.), de-duplicated
|
|
21
|
+
# against the rule hits, and gated by a confidence threshold.
|
|
22
|
+
_model_predict: Callable[[str], list[Correction]] | None
|
|
23
|
+
try: # pragma: no cover
|
|
24
|
+
from .model import predict as _model_predict
|
|
25
|
+
except Exception: # pragma: no cover
|
|
26
|
+
_model_predict = None
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def check(
|
|
30
|
+
text: str,
|
|
31
|
+
*,
|
|
32
|
+
enable: Iterable[str] | None = None,
|
|
33
|
+
disable: Iterable[str] | None = None,
|
|
34
|
+
min_confidence: float = 0.0,
|
|
35
|
+
) -> list[Correction]:
|
|
36
|
+
on = set(enable or ())
|
|
37
|
+
off = set(disable or ())
|
|
38
|
+
|
|
39
|
+
out: list[Correction] = []
|
|
40
|
+
for r in all_rules():
|
|
41
|
+
keys = {r.id, r.category}
|
|
42
|
+
active = bool(keys & on) or (r.enabled_by_default and not (keys & off))
|
|
43
|
+
if active:
|
|
44
|
+
out.extend(r.fn(text))
|
|
45
|
+
|
|
46
|
+
if _model_predict is not None and "model" not in off: # pragma: no cover
|
|
47
|
+
rule_spans = [(c.start, c.end) for c in out]
|
|
48
|
+
for c in _model_predict(text):
|
|
49
|
+
# a rule hit wins on any span overlap — it's the higher-precision
|
|
50
|
+
# signal; the model only adds what the rules missed
|
|
51
|
+
if not any(c.start < re_ and rs < c.end for rs, re_ in rule_spans):
|
|
52
|
+
out.append(c)
|
|
53
|
+
|
|
54
|
+
out = [c for c in out if c.confidence >= min_confidence]
|
|
55
|
+
out.sort(key=lambda c: (c.start, c.end, c.rule_id))
|
|
56
|
+
return out
|
quire_grammar/edits.py
ADDED
|
@@ -0,0 +1,241 @@
|
|
|
1
|
+
"""The GECToR-style edit-tag scheme — shared by the phase-2 model's training
|
|
2
|
+
pipeline and its runtime decoder, so it lives in the shipped package.
|
|
3
|
+
|
|
4
|
+
The model reads a sentence, splits it into word tokens, and predicts one
|
|
5
|
+
**edit tag** per token. Applying the tags and re-joining rebuilds the
|
|
6
|
+
corrected sentence.
|
|
7
|
+
|
|
8
|
+
$KEEP leave the token alone
|
|
9
|
+
$DELETE drop the token
|
|
10
|
+
$REPLACE_{t} swap the token for t
|
|
11
|
+
$APPEND_{t} insert t (one or more words) after the token
|
|
12
|
+
$CASE_CAPITAL Titlecase the token
|
|
13
|
+
$CASE_LOWER lowercase the token
|
|
14
|
+
$CASE_UPPER UPPERCASE the token
|
|
15
|
+
|
|
16
|
+
Kept small on purpose — precision over coverage. Multi-edit tokens (GECToR
|
|
17
|
+
uses iterative rounds) are rare in native-writer text; ``$APPEND_{t}`` carries
|
|
18
|
+
the whole inserted run so one pass covers them. Anything the scheme cannot
|
|
19
|
+
express round-trips wrong, and :func:`encode`'s caller drops the pair.
|
|
20
|
+
|
|
21
|
+
Zero dependencies — part of the rules-layer contract.
|
|
22
|
+
"""
|
|
23
|
+
|
|
24
|
+
from __future__ import annotations
|
|
25
|
+
|
|
26
|
+
import re
|
|
27
|
+
from dataclasses import dataclass
|
|
28
|
+
|
|
29
|
+
KEEP = "$KEEP"
|
|
30
|
+
DELETE = "$DELETE"
|
|
31
|
+
_REPLACE = "$REPLACE_"
|
|
32
|
+
_APPEND = "$APPEND_"
|
|
33
|
+
CASE_CAPITAL = "$CASE_CAPITAL"
|
|
34
|
+
CASE_LOWER = "$CASE_LOWER"
|
|
35
|
+
CASE_UPPER = "$CASE_UPPER"
|
|
36
|
+
|
|
37
|
+
_CASE_TAGS = (CASE_CAPITAL, CASE_LOWER, CASE_UPPER)
|
|
38
|
+
|
|
39
|
+
#: the tags that are always in the vocabulary, at fixed indices
|
|
40
|
+
BASE_TAGS: tuple[str, ...] = (KEEP, DELETE, *_CASE_TAGS)
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def replace_tag(token: str) -> str:
|
|
44
|
+
return _REPLACE + token
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def append_tag(run: str) -> str:
|
|
48
|
+
return _APPEND + run
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def is_replace(tag: str) -> bool:
|
|
52
|
+
return tag.startswith(_REPLACE)
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def is_append(tag: str) -> bool:
|
|
56
|
+
return tag.startswith(_APPEND)
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def tag_value(tag: str) -> str:
|
|
60
|
+
"""The payload of a ``$REPLACE_`` / ``$APPEND_`` tag."""
|
|
61
|
+
return tag.split("_", 1)[1]
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
# --------------------------------------------------------------------------
|
|
65
|
+
# tokenisation — word tokens with punctuation split off, plus a deterministic
|
|
66
|
+
# detokeniser. Not exact for pathological spacing; callers verify and drop.
|
|
67
|
+
# --------------------------------------------------------------------------
|
|
68
|
+
|
|
69
|
+
_TOKEN_RE = re.compile(
|
|
70
|
+
r"""
|
|
71
|
+
\d[\d,]*(?:\.\d+)? # 12 1,000 3.5
|
|
72
|
+
| [A-Za-z]+(?:['’][A-Za-z]+)* # word, with internal apostrophes
|
|
73
|
+
| \.\.\.|[.!?]+ # sentence punctuation, runs kept together
|
|
74
|
+
| --+|[-–—] # dashes / hyphen
|
|
75
|
+
| [^\s\w] # any other single punctuation mark
|
|
76
|
+
""",
|
|
77
|
+
re.VERBOSE,
|
|
78
|
+
)
|
|
79
|
+
|
|
80
|
+
# marks that hug the previous token (no space before)
|
|
81
|
+
_CLOSE = set(".,;:!?)]}%…") | {"...", "n't", "'s", "'re", "'ve", "'ll", "'m", "'d"}
|
|
82
|
+
# marks that hug the next token (no space after)
|
|
83
|
+
_OPEN = set("([{$")
|
|
84
|
+
_CONTRACTION = re.compile(r"^['’](s|re|ve|ll|m|d)$|^n['’]t$", re.I)
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def tokenize(text: str) -> list[str]:
|
|
88
|
+
return _TOKEN_RE.findall(text)
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def tokenize_spans(text: str) -> list[tuple[str, int, int]]:
|
|
92
|
+
"""``(token, start, end)`` — same tokens as :func:`tokenize`, with their
|
|
93
|
+
character offsets in ``text`` (the runtime decoder needs them to place
|
|
94
|
+
``Correction`` spans)."""
|
|
95
|
+
return [(m.group(0), m.start(), m.end()) for m in _TOKEN_RE.finditer(text)]
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
def detokenize(tokens: list[str]) -> str:
|
|
99
|
+
"""Join word tokens back into a sentence. Deterministic; assumes the
|
|
100
|
+
tokens describe ordinary prose."""
|
|
101
|
+
out: list[str] = []
|
|
102
|
+
quote_open = False
|
|
103
|
+
for tok in tokens:
|
|
104
|
+
if not out:
|
|
105
|
+
out.append(tok)
|
|
106
|
+
continue
|
|
107
|
+
prev = out[-1]
|
|
108
|
+
glue = (
|
|
109
|
+
tok in _CLOSE
|
|
110
|
+
or _CONTRACTION.match(tok)
|
|
111
|
+
or prev in _OPEN
|
|
112
|
+
or prev[-1:] in _OPEN
|
|
113
|
+
or (tok == '"' and quote_open)
|
|
114
|
+
or (prev == '"' and not quote_open)
|
|
115
|
+
or (tok in "-–—")
|
|
116
|
+
or (prev in "-–—")
|
|
117
|
+
or (tok in "'’" and quote_open)
|
|
118
|
+
)
|
|
119
|
+
out.append(tok if glue else " " + tok)
|
|
120
|
+
if tok == '"':
|
|
121
|
+
quote_open = not quote_open
|
|
122
|
+
return "".join(out)
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
# --------------------------------------------------------------------------
|
|
126
|
+
# alignment — token-level Levenshtein, then read off the edit tags
|
|
127
|
+
# --------------------------------------------------------------------------
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
def _case_transform(src: str, dst: str) -> str | None:
|
|
131
|
+
if src == dst or src.lower() != dst.lower():
|
|
132
|
+
return None
|
|
133
|
+
if dst == src.capitalize():
|
|
134
|
+
return CASE_CAPITAL
|
|
135
|
+
if dst == src.lower():
|
|
136
|
+
return CASE_LOWER
|
|
137
|
+
if dst == src.upper():
|
|
138
|
+
return CASE_UPPER
|
|
139
|
+
return None
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
def _flush_append(tags: list[str], src_pos: int, pending: list[str]) -> None:
|
|
143
|
+
if pending:
|
|
144
|
+
tags[src_pos - 1] = append_tag(" ".join(pending))
|
|
145
|
+
pending.clear()
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
def _tags_from_alignment(src: list[str], tgt: list[str]) -> list[str] | None:
|
|
149
|
+
n, m = len(src), len(tgt)
|
|
150
|
+
d = [[0] * (m + 1) for _ in range(n + 1)]
|
|
151
|
+
for i in range(1, n + 1):
|
|
152
|
+
d[i][0] = i
|
|
153
|
+
for j in range(1, m + 1):
|
|
154
|
+
d[0][j] = j
|
|
155
|
+
for i in range(1, n + 1):
|
|
156
|
+
for j in range(1, m + 1):
|
|
157
|
+
cost = 0 if src[i - 1] == tgt[j - 1] else 1
|
|
158
|
+
d[i][j] = min(d[i - 1][j] + 1, d[i][j - 1] + 1, d[i - 1][j - 1] + cost)
|
|
159
|
+
|
|
160
|
+
tags: list[str] = [KEEP] * n
|
|
161
|
+
pending_append: list[str] = []
|
|
162
|
+
i, j = n, m
|
|
163
|
+
lead_insert: list[str] = []
|
|
164
|
+
while i > 0 or j > 0:
|
|
165
|
+
if i > 0 and j > 0 and src[i - 1] == tgt[j - 1] and d[i][j] == d[i - 1][j - 1]:
|
|
166
|
+
_flush_append(tags, i, pending_append)
|
|
167
|
+
i, j = i - 1, j - 1
|
|
168
|
+
elif i > 0 and j > 0 and d[i][j] == d[i - 1][j - 1] + 1:
|
|
169
|
+
ct = _case_transform(src[i - 1], tgt[j - 1])
|
|
170
|
+
if pending_append:
|
|
171
|
+
tags[i - 1] = replace_tag(" ".join([tgt[j - 1], *pending_append]))
|
|
172
|
+
pending_append.clear()
|
|
173
|
+
else:
|
|
174
|
+
tags[i - 1] = ct or replace_tag(tgt[j - 1])
|
|
175
|
+
i, j = i - 1, j - 1
|
|
176
|
+
elif j > 0 and d[i][j] == d[i][j - 1] + 1:
|
|
177
|
+
pending_append.insert(0, tgt[j - 1])
|
|
178
|
+
j -= 1
|
|
179
|
+
if i == 0:
|
|
180
|
+
lead_insert = list(pending_append)
|
|
181
|
+
else:
|
|
182
|
+
tags[i - 1] = DELETE
|
|
183
|
+
_flush_append(tags, i, pending_append)
|
|
184
|
+
i -= 1
|
|
185
|
+
if lead_insert:
|
|
186
|
+
return None
|
|
187
|
+
return tags
|
|
188
|
+
|
|
189
|
+
|
|
190
|
+
# --------------------------------------------------------------------------
|
|
191
|
+
# public: encode / apply / verify
|
|
192
|
+
# --------------------------------------------------------------------------
|
|
193
|
+
|
|
194
|
+
|
|
195
|
+
@dataclass(frozen=True)
|
|
196
|
+
class Tagged:
|
|
197
|
+
tokens: tuple[str, ...]
|
|
198
|
+
tags: tuple[str, ...]
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
def encode(corrupt: str, clean: str) -> Tagged | None:
|
|
202
|
+
"""``(corrupt, clean)`` → token tags, or ``None`` if the scheme can't
|
|
203
|
+
represent the edit or the result won't round-trip."""
|
|
204
|
+
src, tgt = tokenize(corrupt), tokenize(clean)
|
|
205
|
+
if not src:
|
|
206
|
+
return None
|
|
207
|
+
tags = _tags_from_alignment(src, tgt)
|
|
208
|
+
if tags is None:
|
|
209
|
+
return None
|
|
210
|
+
tagged = Tagged(tuple(src), tuple(tags))
|
|
211
|
+
if apply(tagged.tokens, tagged.tags) != detokenize(tgt):
|
|
212
|
+
return None
|
|
213
|
+
return tagged
|
|
214
|
+
|
|
215
|
+
|
|
216
|
+
def apply(tokens: tuple[str, ...] | list[str],
|
|
217
|
+
tags: tuple[str, ...] | list[str]) -> str:
|
|
218
|
+
out: list[str] = []
|
|
219
|
+
for tok, tag in zip(tokens, tags, strict=True):
|
|
220
|
+
if tag == KEEP:
|
|
221
|
+
out.append(tok)
|
|
222
|
+
elif tag == DELETE:
|
|
223
|
+
continue
|
|
224
|
+
elif tag == CASE_CAPITAL:
|
|
225
|
+
out.append(tok.capitalize())
|
|
226
|
+
elif tag == CASE_LOWER:
|
|
227
|
+
out.append(tok.lower())
|
|
228
|
+
elif tag == CASE_UPPER:
|
|
229
|
+
out.append(tok.upper())
|
|
230
|
+
elif is_replace(tag):
|
|
231
|
+
out.extend(tokenize(tag_value(tag)))
|
|
232
|
+
elif is_append(tag):
|
|
233
|
+
out.append(tok)
|
|
234
|
+
out.extend(tokenize(tag_value(tag)))
|
|
235
|
+
else: # unknown tag → treat as KEEP (a miss, never a crash)
|
|
236
|
+
out.append(tok)
|
|
237
|
+
return detokenize(out)
|
|
238
|
+
|
|
239
|
+
|
|
240
|
+
def verify(corrupt: str, clean: str, tagged: Tagged) -> bool:
|
|
241
|
+
return apply(tagged.tokens, tagged.tags) == detokenize(tokenize(clean))
|
quire_grammar/model.py
ADDED
|
@@ -0,0 +1,245 @@
|
|
|
1
|
+
"""Phase 2 — the optional local model behind :func:`quire_grammar.check`.
|
|
2
|
+
|
|
3
|
+
``check()`` imports :func:`predict` from here inside a ``try``/``except``. When
|
|
4
|
+
the ``ml`` extra (``onnxruntime`` + ``tokenizers``) isn't installed, or the
|
|
5
|
+
model file isn't bundled in ``models/``, :func:`predict` returns ``[]`` and
|
|
6
|
+
the checker runs rules-only — exactly how Quire treats a missing WordNet or
|
|
7
|
+
spell dictionary. It never raises for a missing model.
|
|
8
|
+
|
|
9
|
+
The model is a GECToR-style tagger: word tokens in, one
|
|
10
|
+
:mod:`quire_grammar.edits` tag per token out. This module runs the ONNX
|
|
11
|
+
graph, reads the argmax tag for each word, and turns the non-``$KEEP`` ones
|
|
12
|
+
into :class:`~quire_grammar.types.Correction` spans over the original text.
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
from __future__ import annotations
|
|
16
|
+
|
|
17
|
+
import importlib.resources
|
|
18
|
+
import json
|
|
19
|
+
from dataclasses import dataclass
|
|
20
|
+
from typing import Any
|
|
21
|
+
|
|
22
|
+
from .edits import (
|
|
23
|
+
CASE_CAPITAL,
|
|
24
|
+
CASE_LOWER,
|
|
25
|
+
CASE_UPPER,
|
|
26
|
+
KEEP,
|
|
27
|
+
is_append,
|
|
28
|
+
is_replace,
|
|
29
|
+
tag_value,
|
|
30
|
+
tokenize_spans,
|
|
31
|
+
)
|
|
32
|
+
from .edits import (
|
|
33
|
+
apply as _apply_tags,
|
|
34
|
+
)
|
|
35
|
+
from .types import AGREEMENT, CONFUSABLE, MECHANICS, PUNCTUATION, Correction
|
|
36
|
+
|
|
37
|
+
_FILES = ("tagger.onnx", "tokenizer.json", "tag_vocab.json")
|
|
38
|
+
|
|
39
|
+
# Drop any predicted edit below this softmax mass. Tuned on held-out clean
|
|
40
|
+
# Gutenberg prose: 0.9 keeps the false-flag rate on correct text near 1.3% on
|
|
41
|
+
# x86 Linux while F0.5 on the synthetic test set holds at ~0.88.
|
|
42
|
+
#
|
|
43
|
+
# CAVEAT: int8 ONNX inference is NOT bit-portable. On a genuinely ambiguous
|
|
44
|
+
# input one word scored 0.51 on x86 Linux and 0.94 on Apple Silicon — no
|
|
45
|
+
# fixed threshold is stable across CPU architectures near a decision
|
|
46
|
+
# boundary. The model is advisory, not autocorrect, and the deterministic
|
|
47
|
+
# rules layer is the backbone; still, treat model output as "correct on this
|
|
48
|
+
# machine" not "the same everywhere". v3: try QDQ static quantisation, or
|
|
49
|
+
# drop the ambiguous confusable pairs (breath/breathe, affect/effect, …)
|
|
50
|
+
# that the narrow rules already cover.
|
|
51
|
+
# Precision over recall — raise it, never lower it, without re-checking.
|
|
52
|
+
_MIN_PROB = 0.9
|
|
53
|
+
_CONFIDENCE_SCALE = 0.9 # a model hit stays below a mechanical rule's 0.9
|
|
54
|
+
_MAX_TOKENS = 128
|
|
55
|
+
|
|
56
|
+
# homophone pairs the model may reach for — flagged as CONFUSABLE, not
|
|
57
|
+
# AGREEMENT, when it swaps one for the other
|
|
58
|
+
_CONFUSABLE_PAIRS: frozenset[frozenset[str]] = frozenset(
|
|
59
|
+
frozenset(p) for p in (
|
|
60
|
+
("its", "it's"), ("there", "their"), ("there", "they're"),
|
|
61
|
+
("their", "they're"), ("your", "you're"), ("then", "than"),
|
|
62
|
+
("affect", "effect"), ("lose", "loose"), ("to", "too"),
|
|
63
|
+
("passed", "past"), ("lead", "led"), ("whose", "who's"),
|
|
64
|
+
("weather", "whether"), ("breath", "breathe"),
|
|
65
|
+
)
|
|
66
|
+
)
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
@dataclass(frozen=True)
|
|
70
|
+
class _Model:
|
|
71
|
+
session: Any
|
|
72
|
+
tokenizer: Any
|
|
73
|
+
id2tag: dict[int, str]
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
_UNSET: Any = object()
|
|
77
|
+
_cache: _Model | None = _UNSET
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
def _resolve_dir() -> Any:
|
|
81
|
+
"""The directory holding the model files, or ``None``. Checked in order:
|
|
82
|
+
the companion ``quire-grammar-model`` package (the ``[model]`` extra),
|
|
83
|
+
then a ``models/`` folder inside this package (a local build). Both go
|
|
84
|
+
through ``importlib.resources`` so they resolve inside a frozen app."""
|
|
85
|
+
try:
|
|
86
|
+
import quire_grammar_model
|
|
87
|
+
|
|
88
|
+
d = quire_grammar_model.path()
|
|
89
|
+
if all((d / f).is_file() for f in _FILES):
|
|
90
|
+
return d
|
|
91
|
+
except ImportError:
|
|
92
|
+
pass
|
|
93
|
+
try:
|
|
94
|
+
d = importlib.resources.files("quire_grammar") / "models"
|
|
95
|
+
if all((d / f).is_file() for f in _FILES):
|
|
96
|
+
return d
|
|
97
|
+
except (ModuleNotFoundError, FileNotFoundError):
|
|
98
|
+
pass
|
|
99
|
+
return None
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
def _load() -> _Model | None:
|
|
103
|
+
try:
|
|
104
|
+
import onnxruntime
|
|
105
|
+
from tokenizers import Tokenizer
|
|
106
|
+
except ImportError:
|
|
107
|
+
return None
|
|
108
|
+
d = _resolve_dir()
|
|
109
|
+
if d is None:
|
|
110
|
+
return None
|
|
111
|
+
try:
|
|
112
|
+
# str() of a Traversable is a real path for an unpacked package —
|
|
113
|
+
# and a 66 MB model is never shipped zipped
|
|
114
|
+
session = onnxruntime.InferenceSession(
|
|
115
|
+
str(d / "tagger.onnx"), providers=["CPUExecutionProvider"])
|
|
116
|
+
tokenizer = Tokenizer.from_file(str(d / "tokenizer.json"))
|
|
117
|
+
tags = json.loads((d / "tag_vocab.json").read_text(encoding="utf-8"))["tags"]
|
|
118
|
+
return _Model(session, tokenizer, dict(enumerate(tags)))
|
|
119
|
+
except Exception: # a corrupt bundle must not take the checker down
|
|
120
|
+
return None
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
def _model() -> _Model | None:
|
|
124
|
+
global _cache
|
|
125
|
+
if _cache is _UNSET:
|
|
126
|
+
_cache = _load()
|
|
127
|
+
return _cache
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
def reset_cache() -> None:
|
|
131
|
+
"""Forget the loaded model — for tests, or after swapping the bundle."""
|
|
132
|
+
global _cache
|
|
133
|
+
_cache = _UNSET
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
def is_available() -> bool:
|
|
137
|
+
return _model() is not None
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
def _softmax(row: Any) -> Any:
|
|
141
|
+
import numpy as np
|
|
142
|
+
|
|
143
|
+
shifted = row - row.max()
|
|
144
|
+
ex = np.exp(shifted)
|
|
145
|
+
return ex / ex.sum()
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
def _category(original: str, replacement: str) -> str:
|
|
149
|
+
if frozenset((original.lower(), replacement.lower())) in _CONFUSABLE_PAIRS:
|
|
150
|
+
return CONFUSABLE
|
|
151
|
+
if replacement and not replacement[0].isalnum():
|
|
152
|
+
return PUNCTUATION
|
|
153
|
+
return AGREEMENT
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
def _to_correction(text: str, span: tuple[str, int, int], tag: str,
|
|
157
|
+
conf: float) -> Correction | None:
|
|
158
|
+
word, start, end = span
|
|
159
|
+
if is_replace(tag):
|
|
160
|
+
repl = tag_value(tag)
|
|
161
|
+
return Correction(
|
|
162
|
+
start, end, word, _category(word, repl), "model:replace",
|
|
163
|
+
f"“{word}” here reads as “{repl}”.", (repl,), conf)
|
|
164
|
+
if is_append(tag):
|
|
165
|
+
become = _apply_tags((word,), (tag,))
|
|
166
|
+
inserted = tag_value(tag)
|
|
167
|
+
cat = PUNCTUATION if inserted[:1] and not inserted[0].isalnum() else AGREEMENT
|
|
168
|
+
return Correction(
|
|
169
|
+
start, end, word, cat, "model:append",
|
|
170
|
+
f"“{word}” may want “{become}” here.", (become,), conf)
|
|
171
|
+
if tag in (CASE_CAPITAL, CASE_LOWER, CASE_UPPER):
|
|
172
|
+
fixed = {CASE_CAPITAL: word.capitalize(), CASE_LOWER: word.lower(),
|
|
173
|
+
CASE_UPPER: word.upper()}[tag]
|
|
174
|
+
if fixed == word:
|
|
175
|
+
return None
|
|
176
|
+
return Correction(
|
|
177
|
+
start, end, word, MECHANICS, "model:case",
|
|
178
|
+
f"“{word}” may need to be “{fixed}” here.", (fixed,), conf)
|
|
179
|
+
if tag == "$DELETE":
|
|
180
|
+
stop = end
|
|
181
|
+
while stop < len(text) and text[stop] == " ":
|
|
182
|
+
stop += 1
|
|
183
|
+
return Correction(
|
|
184
|
+
start, stop, text[start:stop], _category(word, ""), "model:delete",
|
|
185
|
+
f"“{word}” may be extra here.", ("",), conf)
|
|
186
|
+
return None
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
def decode(text: str, per_word: list[tuple[str, float] | None]) -> list[Correction]:
|
|
190
|
+
"""Turn per-word ``(tag, probability)`` predictions — aligned to
|
|
191
|
+
``tokenize_spans(text)`` — into ``Correction`` spans. Pure: the numpy /
|
|
192
|
+
ONNX half of :func:`predict` hands off here, and the tests drive it
|
|
193
|
+
directly."""
|
|
194
|
+
spans = tokenize_spans(text)
|
|
195
|
+
out: list[Correction] = []
|
|
196
|
+
for span, pred in zip(spans, per_word, strict=False):
|
|
197
|
+
if pred is None:
|
|
198
|
+
continue
|
|
199
|
+
tag, prob = pred
|
|
200
|
+
if tag == KEEP or prob < _MIN_PROB:
|
|
201
|
+
continue
|
|
202
|
+
corr = _to_correction(text, span, tag, round(prob * _CONFIDENCE_SCALE, 3))
|
|
203
|
+
if corr is not None:
|
|
204
|
+
out.append(corr)
|
|
205
|
+
return out
|
|
206
|
+
|
|
207
|
+
|
|
208
|
+
def predict(text: str) -> list[Correction]:
|
|
209
|
+
"""Model corrections for ``text``. ``[]`` when the model isn't available."""
|
|
210
|
+
model = _model()
|
|
211
|
+
if model is None or not text.strip():
|
|
212
|
+
return []
|
|
213
|
+
|
|
214
|
+
import numpy as np
|
|
215
|
+
|
|
216
|
+
spans = tokenize_spans(text)
|
|
217
|
+
if not spans:
|
|
218
|
+
return []
|
|
219
|
+
words = [w for w, _, _ in spans]
|
|
220
|
+
|
|
221
|
+
enc = model.tokenizer.encode(words, is_pretokenized=True)
|
|
222
|
+
ids = list(enc.ids)[:_MAX_TOKENS]
|
|
223
|
+
word_ids = list(enc.word_ids)[:_MAX_TOKENS]
|
|
224
|
+
|
|
225
|
+
feed = {
|
|
226
|
+
"input_ids": np.asarray([ids], dtype=np.int64),
|
|
227
|
+
"attention_mask": np.ones((1, len(ids)), dtype=np.int64),
|
|
228
|
+
}
|
|
229
|
+
logits = model.session.run(None, feed)[0][0] # (seq, n_tags)
|
|
230
|
+
|
|
231
|
+
first_pos: dict[int, int] = {}
|
|
232
|
+
for sub_i, wid in enumerate(word_ids):
|
|
233
|
+
if wid is not None and wid not in first_pos:
|
|
234
|
+
first_pos[wid] = sub_i
|
|
235
|
+
|
|
236
|
+
per_word: list[tuple[str, float] | None] = []
|
|
237
|
+
for wi in range(len(spans)):
|
|
238
|
+
pos = first_pos.get(wi)
|
|
239
|
+
if pos is None: # word past the token limit
|
|
240
|
+
per_word.append(None)
|
|
241
|
+
continue
|
|
242
|
+
probs = _softmax(logits[pos])
|
|
243
|
+
k = int(probs.argmax())
|
|
244
|
+
per_word.append((model.id2tag.get(k, KEEP), float(probs[k])))
|
|
245
|
+
return decode(text, per_word)
|
quire_grammar/py.typed
ADDED
|
File without changes
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
"""The rule registry. Each rule is a pure function ``str -> Iterable[Correction]``
|
|
2
|
+
registered with an id and a category. No state, no I/O, no dependencies.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
from __future__ import annotations
|
|
6
|
+
|
|
7
|
+
from collections.abc import Callable, Iterable
|
|
8
|
+
from dataclasses import dataclass
|
|
9
|
+
|
|
10
|
+
from .types import CATEGORIES, Correction
|
|
11
|
+
|
|
12
|
+
Checker = Callable[[str], Iterable[Correction]]
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
@dataclass(frozen=True)
|
|
16
|
+
class Rule:
|
|
17
|
+
id: str
|
|
18
|
+
category: str
|
|
19
|
+
fn: Checker
|
|
20
|
+
enabled_by_default: bool = True
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
_RULES: dict[str, Rule] = {}
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def rule(id: str, category: str, *, default: bool = True) -> Callable[[Checker], Checker]:
|
|
27
|
+
"""Decorator: register ``fn`` as a rule."""
|
|
28
|
+
if category not in CATEGORIES:
|
|
29
|
+
raise ValueError(f"{id}: unknown category {category!r}")
|
|
30
|
+
|
|
31
|
+
def deco(fn: Checker) -> Checker:
|
|
32
|
+
if id in _RULES:
|
|
33
|
+
raise ValueError(f"duplicate rule id: {id}")
|
|
34
|
+
_RULES[id] = Rule(id, category, fn, default)
|
|
35
|
+
return fn
|
|
36
|
+
|
|
37
|
+
return deco
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def all_rules() -> list[Rule]:
|
|
41
|
+
return list(_RULES.values())
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def rule_ids() -> list[str]:
|
|
45
|
+
return list(_RULES)
|