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.
@@ -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)
@@ -0,0 +1,11 @@
1
+ """Importing this package registers every rule as a side effect."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from . import ( # noqa: F401
6
+ agreement,
7
+ confusables,
8
+ mechanics,
9
+ punctuation,
10
+ style,
11
+ )