rag-wright 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.
- rag_wright/__init__.py +13 -0
- rag_wright/api/__init__.py +33 -0
- rag_wright/api/config.py +59 -0
- rag_wright/api/discover.py +70 -0
- rag_wright/api/documents.py +39 -0
- rag_wright/api/ids.py +31 -0
- rag_wright/api/invoke.py +99 -0
- rag_wright/api/kg.py +61 -0
- rag_wright/api/mcp.py +94 -0
- rag_wright/api/usage.py +30 -0
- rag_wright/api/workspace.py +85 -0
- rag_wright/capabilities/__init__.py +8 -0
- rag_wright/capabilities/answer_generator.py +427 -0
- rag_wright/capabilities/ard.py +286 -0
- rag_wright/capabilities/assertion_extraction.py +79 -0
- rag_wright/capabilities/chunk_read.py +58 -0
- rag_wright/capabilities/chunk_write.py +163 -0
- rag_wright/capabilities/claim_extraction.py +153 -0
- rag_wright/capabilities/clause_exception_linking.py +117 -0
- rag_wright/capabilities/compliance_judgment.py +322 -0
- rag_wright/capabilities/compliance_store.py +87 -0
- rag_wright/capabilities/contract_kg_serve.py +156 -0
- rag_wright/capabilities/contract_kg_store.py +251 -0
- rag_wright/capabilities/dg_extraction.py +585 -0
- rag_wright/capabilities/disambiguation.py +163 -0
- rag_wright/capabilities/document_parse.py +87 -0
- rag_wright/capabilities/document_scope.py +49 -0
- rag_wright/capabilities/embedding.py +164 -0
- rag_wright/capabilities/embedding_profiles.py +43 -0
- rag_wright/capabilities/entity_resolution.py +154 -0
- rag_wright/capabilities/fusion.py +64 -0
- rag_wright/capabilities/graph_extraction.py +243 -0
- rag_wright/capabilities/graph_query.py +73 -0
- rag_wright/capabilities/graph_storage.py +111 -0
- rag_wright/capabilities/highlight_serve.py +142 -0
- rag_wright/capabilities/hybrid_search.py +65 -0
- rag_wright/capabilities/invoke.py +31 -0
- rag_wright/capabilities/jev_decision.py +38 -0
- rag_wright/capabilities/manifests.py +872 -0
- rag_wright/capabilities/okf_navigate.py +456 -0
- rag_wright/capabilities/parsing.py +286 -0
- rag_wright/capabilities/property_boosted_retrieval.py +125 -0
- rag_wright/capabilities/query_function_classifier.py +94 -0
- rag_wright/capabilities/query_understanding.py +109 -0
- rag_wright/capabilities/registry.py +262 -0
- rag_wright/capabilities/remote_encoders.py +94 -0
- rag_wright/capabilities/requirement_extraction.py +247 -0
- rag_wright/capabilities/reranking.py +123 -0
- rag_wright/capabilities/retrieval_core.py +126 -0
- rag_wright/capabilities/rlm_chunking.py +808 -0
- rag_wright/capabilities/rlm_synthesis.py +316 -0
- rag_wright/capabilities/scan_quality.py +136 -0
- rag_wright/capabilities/span_relevance_judgment.py +191 -0
- rag_wright/capabilities/vision_to_text.py +85 -0
- rag_wright/capabilities/vlm_ocr.py +85 -0
- rag_wright/contracts/__init__.py +6 -0
- rag_wright/contracts/chunk.py +79 -0
- rag_wright/contracts/compliance.py +303 -0
- rag_wright/contracts/contract_meta.py +27 -0
- rag_wright/contracts/extraction.py +130 -0
- rag_wright/contracts/function.py +167 -0
- rag_wright/contracts/function_routing.py +91 -0
- rag_wright/contracts/highlight.py +74 -0
- rag_wright/contracts/identifiers.py +153 -0
- rag_wright/contracts/jurisdiction.py +96 -0
- rag_wright/contracts/ontology.py +142 -0
- rag_wright/contracts/property.py +201 -0
- rag_wright/contracts/provenance.py +78 -0
- rag_wright/contracts/query_intent.py +53 -0
- rag_wright/contracts/span.py +76 -0
- rag_wright/contracts/value_match.py +84 -0
- rag_wright/corpus/__init__.py +0 -0
- rag_wright/corpus/canonicalize.py +116 -0
- rag_wright/corpus/cuad.py +153 -0
- rag_wright/corpus/cuad_ingestion.py +72 -0
- rag_wright/corpus/document_parser.py +299 -0
- rag_wright/corpus/edgar.py +231 -0
- rag_wright/corpus/gcs_ingestion.py +120 -0
- rag_wright/corpus/http.py +110 -0
- rag_wright/corpus/selection.py +152 -0
- rag_wright/mcp/__init__.py +11 -0
- rag_wright/mcp/compliance_server.py +299 -0
- rag_wright/mcp/intra_document_qa_server.py +170 -0
- rag_wright/mcp/relational_qa_server.py +171 -0
- rag_wright/mcp/session_store.py +64 -0
- rag_wright/mcp/typed_property_retrieval_server.py +191 -0
- rag_wright/models/__init__.py +8 -0
- rag_wright/models/profiles.py +331 -0
- rag_wright/models/seam.py +497 -0
- rag_wright/models/tag_structured.py +285 -0
- rag_wright/models/tracing.py +179 -0
- rag_wright/models/usage.py +102 -0
- rag_wright/okf/__init__.py +11 -0
- rag_wright/okf/compile.py +292 -0
- rag_wright/okf/document.py +47 -0
- rag_wright/okf/enrich.py +176 -0
- rag_wright/okf/links.py +190 -0
- rag_wright/okf/lint.py +105 -0
- rag_wright/ontology/__init__.py +6 -0
- rag_wright/ontology/_generated_template_meta.py +60 -0
- rag_wright/ontology/_generated_vocab.py +52 -0
- rag_wright/ontology/clause_template.py +964 -0
- rag_wright/ontology/codegen.py +84 -0
- rag_wright/ontology/compliance_bridge.ttl +186 -0
- rag_wright/ontology/contract_bridge.ttl +2685 -0
- rag_wright/ontology/contract_taxonomy.py +24 -0
- rag_wright/ontology/derive.py +58 -0
- rag_wright/ontology/loader.py +435 -0
- rag_wright/ontology/packs/ftc_16cfr255.ttl +29 -0
- rag_wright/ontology/registry.py +87 -0
- rag_wright/ontology/template_introspect.py +100 -0
- rag_wright/py.typed +0 -0
- rag_wright/reference/__init__.py +2 -0
- rag_wright/reference/compliance.py +41 -0
- rag_wright/reference/contract_seam.py +123 -0
- rag_wright/skills/__init__.py +7 -0
- rag_wright/skills/claim_extraction/SKILL.md +47 -0
- rag_wright/skills/claim_extraction/__init__.py +1 -0
- rag_wright/skills/claim_extraction/template.py +50 -0
- rag_wright/skills/compliance_judgment/SKILL.md +59 -0
- rag_wright/skills/corpus_ingest/SKILL.md +106 -0
- rag_wright/skills/extraction_semantic_judge/SKILL.md +51 -0
- rag_wright/skills/extraction_semantic_judge/__init__.py +1 -0
- rag_wright/skills/generation/SKILL.md +64 -0
- rag_wright/skills/generation/__init__.py +1 -0
- rag_wright/skills/generic_compliance_judgment/SKILL.md +58 -0
- rag_wright/skills/okf_navigate/SKILL.md +137 -0
- rag_wright/skills/requirement_extraction/SKILL.md +47 -0
- rag_wright/skills/requirement_extraction/__init__.py +1 -0
- rag_wright/skills/requirement_extraction/template.py +50 -0
- rag_wright/skills/rlm/SKILL.md +186 -0
- rag_wright/skills/rlm/__init__.py +31 -0
- rag_wright/skills/rlm/agent.py +292 -0
- rag_wright/skills/span_relevance_judgment/SKILL.md +67 -0
- rag_wright/skills/vision_to_text/SKILL.md +36 -0
- rag_wright/skills/vision_to_text/__init__.py +1 -0
- rag_wright/spans/__init__.py +1 -0
- rag_wright/spans/boundary.py +78 -0
- rag_wright/spans/clause_function_classifier.py +490 -0
- rag_wright/spans/clause_kg_extractor.py +337 -0
- rag_wright/spans/cuad_labels.py +81 -0
- rag_wright/spans/dim_classifier.py +158 -0
- rag_wright/spans/dim_fleet.json +411 -0
- rag_wright/spans/function_classifier.py +77 -0
- rag_wright/spans/function_families.py +62 -0
- rag_wright/spans/hybrid_classifier.py +103 -0
- rag_wright/spans/legalbert_classifier.py +83 -0
- rag_wright/spans/model_capabilities.py +107 -0
- rag_wright/spans/new_function_labels.py +111 -0
- rag_wright/spans/page_map.py +68 -0
- rag_wright/spans/property_extractor.py +365 -0
- rag_wright/spans/property_grounding.py +182 -0
- rag_wright/spans/reclassify.py +77 -0
- rag_wright/spans/scarce_function_labels.py +105 -0
- rag_wright/spans/segment.py +341 -0
- rag_wright/spans/semantic_judge.py +197 -0
- rag_wright/spans/symbolic_validation.py +131 -0
- rag_wright/spans/tag_clause_extractor.py +182 -0
- rag_wright/store/__init__.py +6 -0
- rag_wright/store/arcadedb.py +1135 -0
- rag_wright/store/chunk_text.py +66 -0
- rag_wright/store/seam.py +213 -0
- rag_wright/subgraphs/__init__.py +0 -0
- rag_wright/subgraphs/async_ingestion.py +204 -0
- rag_wright/subgraphs/compliance_check.py +1042 -0
- rag_wright/subgraphs/compliance_ingestion.py +306 -0
- rag_wright/subgraphs/contract_ingestion_pipeline.py +999 -0
- rag_wright/subgraphs/graph_extraction.py +102 -0
- rag_wright/subgraphs/intra_document_qa.py +328 -0
- rag_wright/subgraphs/observability.py +140 -0
- rag_wright/subgraphs/query_constraint_extraction.py +73 -0
- rag_wright/subgraphs/relational_qa.py +165 -0
- rag_wright/subgraphs/requirement_extraction.py +137 -0
- rag_wright/subgraphs/scaffold.py +65 -0
- rag_wright/subgraphs/semantic_chunking.py +183 -0
- rag_wright/subgraphs/typed_clause_extraction.py +172 -0
- rag_wright/subgraphs/typed_property_retrieval.py +278 -0
- rag_wright/util/__init__.py +1 -0
- rag_wright/util/concurrent.py +153 -0
- rag_wright/util/spacy_model.py +45 -0
- rag_wright-0.1.0.dist-info/METADATA +168 -0
- rag_wright-0.1.0.dist-info/RECORD +184 -0
- rag_wright-0.1.0.dist-info/WHEEL +4 -0
- rag_wright-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,490 @@
|
|
|
1
|
+
"""INGEST-LLM-CLASSIFIER (ADR-0048): the clause-level function-classification seam.
|
|
2
|
+
|
|
3
|
+
Replaces the span-level, single-label LegalBERT call at ingestion with a FULL-CLAUSE, MULTI-LABEL, confidence-
|
|
4
|
+
scored classifier by the graph-building LLM. Two implementations behind one `classify(clause_text) -> [FunctionScore]`
|
|
5
|
+
seam: `LlmClauseClassifier` (the default going forward, via the model-profile structured seam) and
|
|
6
|
+
`LegalBertClauseAdapter` (wraps the existing single-label classifier for back-compat). Injected into
|
|
7
|
+
`production_document_ingest` (a later step); the injection point query-time already had.
|
|
8
|
+
|
|
9
|
+
Robustness: the post-processor canonicalizes the LLM's labels, drops off-taxonomy / below-floor labels, and caps
|
|
10
|
+
the count -- so LLM slop never crashes ingestion, and a runnable failure degrades to NO function (the clause
|
|
11
|
+
becomes NONE, exactly as an off-taxonomy span does today)."""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
import asyncio
|
|
16
|
+
import re
|
|
17
|
+
from typing import Any, Protocol, runtime_checkable
|
|
18
|
+
|
|
19
|
+
from pydantic import BaseModel, Field
|
|
20
|
+
|
|
21
|
+
from rag_wright.contracts.function import (
|
|
22
|
+
FUNCTION_LABELS,
|
|
23
|
+
FunctionConfidence,
|
|
24
|
+
FunctionScore,
|
|
25
|
+
canonical_function,
|
|
26
|
+
)
|
|
27
|
+
|
|
28
|
+
_FLOOR: frozenset[FunctionConfidence] = frozenset({FunctionConfidence.HIGH, FunctionConfidence.MEDIUM})
|
|
29
|
+
_MAX_FUNCTIONS = 3
|
|
30
|
+
|
|
31
|
+
# ADR-0048 Phase A: the structured `function` field advertises the closed label set (the 52 taxonomy labels +
|
|
32
|
+
# "OTHER" for a real clause type we lack a label for). Emitted into the JSON schema so guided decoding HARD-
|
|
33
|
+
# constrains the model to a valid label (killing the granite-8B failure mode: inventing free-form names like
|
|
34
|
+
# "Exclusive Source of Supply" that then fall to NONE), and strongly guides it under function-calling. Kept as a
|
|
35
|
+
# `str` field (json_schema_extra is schema-only, not pydantic-enforced), so a stray value never crashes a whole
|
|
36
|
+
# sub-batch and the taxonomy-gap channel (`other_label`) still works. "OTHER" is not a FUNCTION_LABELS entry.
|
|
37
|
+
_FUNCTION_ENUM: list[str] = [*FUNCTION_LABELS, "OTHER"]
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
@runtime_checkable
|
|
41
|
+
class ClauseFunctionClassifier(Protocol):
|
|
42
|
+
"""The ingestion seam (option B): classify a chunk's spans with the chunk as SHARED context, returning one
|
|
43
|
+
ranked `FunctionScore` list per span (aligned to the input order). One LLM call per chunk, not per span."""
|
|
44
|
+
|
|
45
|
+
def classify_spans(self, chunk_text: str, span_texts: list[str]) -> list[list[FunctionScore]]: ...
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
class RawScore(BaseModel):
|
|
49
|
+
"""The LLM's raw (pre-validation) score: a function label (one of the 52 taxonomy labels or "OTHER") + coarse
|
|
50
|
+
confidence. When `function` is "OTHER" (a real clause type not in our taxonomy), `other_label` names it -- the
|
|
51
|
+
taxonomy-gap signal (ADR-0048 option 2). Filtered/categorized downstream. `function` advertises the closed
|
|
52
|
+
label enum in the JSON schema (guided decoding), but stays a `str` so a stray never crashes the sub-batch."""
|
|
53
|
+
|
|
54
|
+
function: str = Field(json_schema_extra={"enum": _FUNCTION_ENUM})
|
|
55
|
+
confidence: str
|
|
56
|
+
other_label: str = ""
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
class ClauseFunctionClassification(BaseModel):
|
|
60
|
+
"""One span/clause's structured output: applicable functions RANKED primary-first, each with a coarse
|
|
61
|
+
`high|medium|low` confidence. Kept raw (str fields) so post-processing can drop LLM slop rather than fail the
|
|
62
|
+
whole call."""
|
|
63
|
+
|
|
64
|
+
functions: list[RawScore]
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
class SpanFunctions(BaseModel):
|
|
68
|
+
"""One span's classification, keyed by its EXPLICIT `span_index` (the `[n]` in the prompt) -- alignment is by
|
|
69
|
+
index, NOT list position, so a dropped/reordered span can't silently shift every later span's label. The
|
|
70
|
+
classifier may OMIT spans it assigns no function."""
|
|
71
|
+
|
|
72
|
+
span_index: int
|
|
73
|
+
functions: list[RawScore]
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
class BatchSpanClassification(BaseModel):
|
|
77
|
+
"""The batched (option B) output: per-span classifications keyed by `span_index` (sparse -- no-function spans
|
|
78
|
+
may be omitted)."""
|
|
79
|
+
|
|
80
|
+
spans: list[SpanFunctions]
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def _to_scores(raw_scores: list[RawScore]) -> list[FunctionScore]:
|
|
84
|
+
"""Canonicalize + confidence-floor (>= medium) + cap (<=3), preserving rank. Drops off-taxonomy / NONE /
|
|
85
|
+
below-floor / unparseable-confidence entries (robust to LLM slop)."""
|
|
86
|
+
out: list[FunctionScore] = []
|
|
87
|
+
for r in raw_scores:
|
|
88
|
+
canon = canonical_function(r.function)
|
|
89
|
+
if canon is None: # off-taxonomy or NONE
|
|
90
|
+
continue
|
|
91
|
+
try:
|
|
92
|
+
conf = FunctionConfidence(r.confidence.strip().lower())
|
|
93
|
+
except ValueError:
|
|
94
|
+
continue
|
|
95
|
+
if conf in _FLOOR:
|
|
96
|
+
out.append(FunctionScore(function=canon, confidence=conf))
|
|
97
|
+
if len(out) >= _MAX_FUNCTIONS:
|
|
98
|
+
break
|
|
99
|
+
return out
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
def categorize_raw(raw_scores: list[RawScore]) -> tuple[list[FunctionScore], list[str]]:
|
|
103
|
+
"""ADR-0048 option 2: split a span's raw scores into (in-taxonomy FunctionScores [canonical+floored],
|
|
104
|
+
out-of-taxonomy labels). An entry whose `function` is off-taxonomy / "OTHER" is captured by its real clause
|
|
105
|
+
type (`other_label`, else the raw `function`) -- distinguishing "no function" from "a function we lack a
|
|
106
|
+
label for" (the taxonomy-gap signal), instead of silently collapsing both to NONE."""
|
|
107
|
+
in_tax = _to_scores(raw_scores)
|
|
108
|
+
others: list[str] = []
|
|
109
|
+
for r in raw_scores:
|
|
110
|
+
if canonical_function(r.function) is None: # off-taxonomy
|
|
111
|
+
# the LLM may follow the convention (function="OTHER", real type in other_label) OR put the real
|
|
112
|
+
# type directly in `function` (with other_label empty/"None"). Prefer other_label ONLY when function
|
|
113
|
+
# is the literal "OTHER"; otherwise the off-taxonomy `function` string IS the real type.
|
|
114
|
+
if r.function.strip().upper() == "OTHER":
|
|
115
|
+
label = r.other_label.strip()
|
|
116
|
+
else:
|
|
117
|
+
label = r.function.strip()
|
|
118
|
+
if label and label.upper() not in ("NONE", "OTHER"):
|
|
119
|
+
others.append(label)
|
|
120
|
+
return in_tax, others
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
_PROMPT = (
|
|
124
|
+
"You are classifying a single contract CLAUSE by its legal function. Read the whole clause (not a fragment) "
|
|
125
|
+
"and list EVERY function it genuinely serves, ranked most-relevant first. A clause may serve more than one "
|
|
126
|
+
"function; most serve exactly one. For each, give a coarse confidence: high, medium, or low. Use ONLY these "
|
|
127
|
+
"function types; if none apply, return an empty list.\n\nFunction types:\n{labels}\n\nClause:\n{text}"
|
|
128
|
+
)
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
class LlmClauseClassifier:
|
|
132
|
+
"""Classify a full clause via the graph-building LLM (structured). `runnable` is the injected structured seam
|
|
133
|
+
(`.invoke(prompt) -> ClauseFunctionClassification`); production wires it through the model-profile seam."""
|
|
134
|
+
|
|
135
|
+
def __init__(self, runnable: Any) -> None:
|
|
136
|
+
self._runnable = runnable
|
|
137
|
+
|
|
138
|
+
def classify(self, clause_text: str) -> list[FunctionScore]:
|
|
139
|
+
from rag_wright.contracts.function import FUNCTION_LABELS
|
|
140
|
+
|
|
141
|
+
prompt = _PROMPT.format(labels="\n".join(FUNCTION_LABELS), text=clause_text)
|
|
142
|
+
try:
|
|
143
|
+
raw = self._runnable.invoke(prompt)
|
|
144
|
+
except Exception: # noqa: BLE001 - a classify failure degrades to NO function (never crash ingestion)
|
|
145
|
+
return []
|
|
146
|
+
if not isinstance(raw, ClauseFunctionClassification):
|
|
147
|
+
return []
|
|
148
|
+
return _to_scores(raw.functions)
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
_BATCH_CAP = 10 # max spans per LLM call -- a big multi-provision chunk is split into sub-batches (bounded prompt)
|
|
152
|
+
|
|
153
|
+
_BATCH_PROMPT = (
|
|
154
|
+
"You are classifying the operative provisions of ONE contract section by their legal function. Read the WHOLE "
|
|
155
|
+
"section for context, then classify each numbered span below by the function(s) it serves -- ranked "
|
|
156
|
+
"most-relevant first, each with a coarse confidence (high, medium, or low). For `function`, use EXACTLY one of "
|
|
157
|
+
"the function types listed below. If a span's real function is NOT in the list, set `function` to \"OTHER\" and "
|
|
158
|
+
"put the actual clause type in `other_label` -- do NOT invent a name in `function`. For EACH span you classify, "
|
|
159
|
+
"return its `span_index` (the [n] number) and its functions; OMIT any span that serves no function. Classify "
|
|
160
|
+
"only indices 0..{max_index}.\n\nFunction types:\n{labels}\n\nSECTION (context):\n{context}\n\nSPANS:\n{spans}"
|
|
161
|
+
)
|
|
162
|
+
|
|
163
|
+
_SUBBATCH_CONCURRENCY = 8 # max concurrent sub-batch LLM calls per chunk (async), bounded against provider limits
|
|
164
|
+
|
|
165
|
+
|
|
166
|
+
def _subs(span_texts: list[str]) -> list[tuple[int, list[str]]]:
|
|
167
|
+
"""Split spans into sub-batches of at most `_BATCH_CAP`, each a `(start_index, spans)` pair."""
|
|
168
|
+
return [(start, span_texts[start:start + _BATCH_CAP]) for start in range(0, len(span_texts), _BATCH_CAP)]
|
|
169
|
+
|
|
170
|
+
|
|
171
|
+
def _sub_prompt(chunk_text: str, sub: list[str]) -> str:
|
|
172
|
+
"""The batch-classify prompt for one sub-batch (shared by the sync and async paths)."""
|
|
173
|
+
from rag_wright.contracts.function import FUNCTION_LABELS
|
|
174
|
+
|
|
175
|
+
numbered = "\n".join(f"[{i}] {t}" for i, t in enumerate(sub))
|
|
176
|
+
return _BATCH_PROMPT.format(
|
|
177
|
+
labels="\n".join(FUNCTION_LABELS), context=chunk_text, spans=numbered, max_index=len(sub) - 1)
|
|
178
|
+
|
|
179
|
+
|
|
180
|
+
def _merge_subbatches(results: list[tuple[int, int, Any]], span_texts: list[str]) -> list[list[RawScore]]:
|
|
181
|
+
"""Merge sub-batch results into index-aligned per-span scores, aligning by the returned `span_index` (relative
|
|
182
|
+
to each sub-batch, offset by its start) so a dropped/reordered span can't shift later labels."""
|
|
183
|
+
out: list[list[RawScore]] = [[] for _ in span_texts]
|
|
184
|
+
for start, sublen, raw in results:
|
|
185
|
+
if raw is None:
|
|
186
|
+
continue
|
|
187
|
+
for sf in raw.spans:
|
|
188
|
+
idx = start + sf.span_index
|
|
189
|
+
if start <= idx < start + sublen:
|
|
190
|
+
out[idx] = list(sf.functions)
|
|
191
|
+
return out
|
|
192
|
+
|
|
193
|
+
|
|
194
|
+
class LlmBatchClauseClassifier:
|
|
195
|
+
"""Option B: classify a chunk's spans with the chunk as shared context. A big chunk is split into sub-batches
|
|
196
|
+
of `_BATCH_CAP` spans (bounded prompt), each ONE LLM call. Output is aligned by the returned `span_index` (not
|
|
197
|
+
list position), so a dropped/reordered span can't shift later labels. `runnable` is the injected structured
|
|
198
|
+
seam (`.invoke(prompt) -> BatchSpanClassification`). Robust: a failed sub-batch leaves its spans empty."""
|
|
199
|
+
|
|
200
|
+
def __init__(self, runnable: Any) -> None:
|
|
201
|
+
self._runnable = runnable
|
|
202
|
+
|
|
203
|
+
def _classify_raw(self, chunk_text: str, span_texts: list[str]) -> list[list[RawScore]]:
|
|
204
|
+
"""Sub-batched LLM calls, run CONCURRENTLY within the chunk (a big multi-provision chunk's sub-batches are
|
|
205
|
+
independent, so a 29-span chunk costs ~1 call's latency, not 3x). Index-aligned RAW per-span scores (the
|
|
206
|
+
LLM's function+confidence strings, NO canonicalize/floor/cap). A failed sub-batch leaves its spans empty."""
|
|
207
|
+
if not span_texts:
|
|
208
|
+
return []
|
|
209
|
+
from concurrent.futures import ThreadPoolExecutor
|
|
210
|
+
|
|
211
|
+
subs = _subs(span_texts)
|
|
212
|
+
|
|
213
|
+
def _call(item):
|
|
214
|
+
start, sub = item
|
|
215
|
+
try:
|
|
216
|
+
raw = self._runnable.invoke(_sub_prompt(chunk_text, sub))
|
|
217
|
+
except Exception: # noqa: BLE001 - a failed sub-batch leaves its spans empty (never crash)
|
|
218
|
+
return (start, len(sub), None)
|
|
219
|
+
return (start, len(sub), raw if isinstance(raw, BatchSpanClassification) else None)
|
|
220
|
+
|
|
221
|
+
if len(subs) <= 1: # single sub-batch -> no thread pool
|
|
222
|
+
results = [_call(subs[0])]
|
|
223
|
+
else: # concurrent sub-batches (independent calls, merged by span_index)
|
|
224
|
+
with ThreadPoolExecutor(max_workers=len(subs)) as ex:
|
|
225
|
+
results = list(ex.map(_call, subs))
|
|
226
|
+
return _merge_subbatches(results, span_texts)
|
|
227
|
+
|
|
228
|
+
async def _aclassify_raw(self, chunk_text: str, span_texts: list[str],
|
|
229
|
+
*, sem: asyncio.Semaphore | None = None) -> list[list[RawScore]]:
|
|
230
|
+
"""ASYNC-B1 (ADR-0057): the async twin of `_classify_raw`. Sub-batches run concurrently via
|
|
231
|
+
`asyncio.gather` bounded by a `Semaphore` (native form of the sync thread pool); each `.ainvoke` carries
|
|
232
|
+
the true wall-clock deadline. A failed sub-batch -- including a `ModelCallTimeout` (an Exception, so it is
|
|
233
|
+
caught here) -- leaves its spans empty: the degrade path stays reachable and the node never raises.
|
|
234
|
+
|
|
235
|
+
CLASSIFY-CONCURRENCY-1: a caller classifying MANY chunks concurrently passes ONE shared `sem`, so the total
|
|
236
|
+
in-flight sub-batch calls across all chunks are bounded by a single deliberate knob (else each chunk got
|
|
237
|
+
its own `_SUBBATCH_CONCURRENCY` budget). The sem is acquired only at the leaf `.ainvoke` -- never held
|
|
238
|
+
across another acquire -- so nesting the chunk gather over it cannot deadlock."""
|
|
239
|
+
if not span_texts:
|
|
240
|
+
return []
|
|
241
|
+
subs = _subs(span_texts)
|
|
242
|
+
sem = sem if sem is not None else asyncio.Semaphore(_SUBBATCH_CONCURRENCY)
|
|
243
|
+
|
|
244
|
+
async def _acall(item):
|
|
245
|
+
start, sub = item
|
|
246
|
+
async with sem:
|
|
247
|
+
try:
|
|
248
|
+
raw = await self._runnable.ainvoke(_sub_prompt(chunk_text, sub))
|
|
249
|
+
except Exception: # noqa: BLE001 - failed sub-batch -> empty spans (never crash the document)
|
|
250
|
+
return (start, len(sub), None)
|
|
251
|
+
return (start, len(sub), raw if isinstance(raw, BatchSpanClassification) else None)
|
|
252
|
+
|
|
253
|
+
results = await asyncio.gather(*(_acall(s) for s in subs))
|
|
254
|
+
return _merge_subbatches(results, span_texts)
|
|
255
|
+
|
|
256
|
+
def classify_spans(self, chunk_text: str, span_texts: list[str]) -> list[list[FunctionScore]]:
|
|
257
|
+
return [_to_scores(raws) for raws in self._classify_raw(chunk_text, span_texts)]
|
|
258
|
+
|
|
259
|
+
async def aclassify_spans(self, chunk_text: str, span_texts: list[str],
|
|
260
|
+
*, sem: asyncio.Semaphore | None = None) -> list[list[FunctionScore]]:
|
|
261
|
+
return [_to_scores(raws) for raws in await self._aclassify_raw(chunk_text, span_texts, sem=sem)]
|
|
262
|
+
|
|
263
|
+
def classify_spans_raw(self, chunk_text: str, span_texts: list[str]) -> list[list[RawScore]]:
|
|
264
|
+
"""Debug: the RAW per-span scores (pre-floor), to see what the confidence floor drops."""
|
|
265
|
+
return self._classify_raw(chunk_text, span_texts)
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
class LegalBertClauseAdapter:
|
|
269
|
+
"""Back-compat: wrap the span-level single-label LegalBERT classifier. Its one label becomes the PRIMARY at
|
|
270
|
+
`high` confidence; NONE / off-taxonomy -> [] (no function). It is span-level, so `classify_spans` ignores the
|
|
271
|
+
chunk context (each span classified independently, exactly as today)."""
|
|
272
|
+
|
|
273
|
+
def __init__(self, classifier: Any) -> None:
|
|
274
|
+
self._classifier = classifier
|
|
275
|
+
|
|
276
|
+
def classify(self, clause_text: str) -> list[FunctionScore]:
|
|
277
|
+
return self.classify_spans("", [clause_text])[0]
|
|
278
|
+
|
|
279
|
+
def classify_spans(self, chunk_text: str, span_texts: list[str]) -> list[list[FunctionScore]]: # noqa: ARG002
|
|
280
|
+
out: list[list[FunctionScore]] = []
|
|
281
|
+
for label in self._classifier.classify(span_texts):
|
|
282
|
+
canon = canonical_function(label)
|
|
283
|
+
out.append([FunctionScore(function=canon, confidence=FunctionConfidence.HIGH)] if canon else [])
|
|
284
|
+
return out
|
|
285
|
+
|
|
286
|
+
|
|
287
|
+
# --- issue 0005 / route (b): CLIENT-SIDE free-text tag classification (no server-side guided decoding) ---------
|
|
288
|
+
#
|
|
289
|
+
# The classifier schemas are nested list[BaseModel] (BatchSpanClassification.spans -> SpanFunctions.functions ->
|
|
290
|
+
# RawScore), so `build_tag_structured` (FLAT only) can't serve them. Server-side guided decoding
|
|
291
|
+
# (`build_structured`) runs away on self-hosted Granite (profiles.py: no `client_side_structured` escape) -- the
|
|
292
|
+
# ADR-0058 boundary-call failure at a new call site. So, like `answer_generator.parse_tagged_answer`, the
|
|
293
|
+
# classifier uses its OWN compact free-text tag format that FLATTENS the nesting into the tag BODY, parsed
|
|
294
|
+
# CLIENT-SIDE into the same Pydantic contract. `function` stays raw -- the downstream categorize/floor/cap cleans
|
|
295
|
+
# a stray, so a malformed token never crashes a sub-batch.
|
|
296
|
+
|
|
297
|
+
_BATCH_TAG_INSTRUCTIONS = (
|
|
298
|
+
"\n\nReturn ONLY the classification as tagged lines -- ONE per span you assign a function to (OMIT any span "
|
|
299
|
+
"with no function). Between the tags put the applicable function labels PRIMARY-FIRST as `Label:confidence` "
|
|
300
|
+
"(confidence = high, medium, or low), comma-separated. For a real clause type NOT in the list above, use "
|
|
301
|
+
"`OTHER:confidence:short_name`. Use ONLY the [n] span indices shown. Example:\n"
|
|
302
|
+
'<span index="0">Cap On Liability:high, Indemnification:medium</span>\n'
|
|
303
|
+
'<span index="3">Governing Law:high</span>'
|
|
304
|
+
)
|
|
305
|
+
_CLAUSE_TAG_INSTRUCTIONS = (
|
|
306
|
+
"\n\nReturn ONLY the applicable function labels PRIMARY-FIRST as `Label:confidence` (confidence = high, "
|
|
307
|
+
"medium, or low), comma-separated, between <functions></functions> tags; for a real clause type NOT in the "
|
|
308
|
+
"list above, use `OTHER:confidence:short_name`. If none apply, return `<functions></functions>`. Example:\n"
|
|
309
|
+
"<functions>Cap On Liability:high, Indemnification:low</functions>"
|
|
310
|
+
)
|
|
311
|
+
|
|
312
|
+
_SPAN_TAG_RE = re.compile(r'<span\s+index="?(\d+)"?\s*>(.*?)</span>', re.DOTALL | re.IGNORECASE)
|
|
313
|
+
_FUNCTIONS_TAG_RE = re.compile(r"<functions>(.*?)</functions>", re.DOTALL | re.IGNORECASE)
|
|
314
|
+
|
|
315
|
+
|
|
316
|
+
def _parse_score_item(item: str) -> RawScore | None:
|
|
317
|
+
"""One `Label:confidence` (or `OTHER:confidence:other_label`) token -> a RawScore, kept RAW."""
|
|
318
|
+
parts = [p.strip() for p in item.split(":")]
|
|
319
|
+
if not parts or not parts[0]:
|
|
320
|
+
return None
|
|
321
|
+
if parts[0].upper() == "OTHER":
|
|
322
|
+
return RawScore(function="OTHER", confidence=parts[1] if len(parts) > 1 else "low",
|
|
323
|
+
other_label=parts[2] if len(parts) > 2 else "")
|
|
324
|
+
if len(parts) == 1:
|
|
325
|
+
return RawScore(function=parts[0], confidence="low") # no confidence given -> lowest (dropped by floor)
|
|
326
|
+
return RawScore(function=":".join(parts[:-1]), confidence=parts[-1]) # rejoin a label that itself had a colon
|
|
327
|
+
|
|
328
|
+
|
|
329
|
+
def _parse_scores(body: str) -> list[RawScore]:
|
|
330
|
+
return [rs for item in re.split(r"[,\n]", body) if (rs := _parse_score_item(item.strip())) is not None]
|
|
331
|
+
|
|
332
|
+
|
|
333
|
+
def parse_batch_span_tags(text: str) -> BatchSpanClassification:
|
|
334
|
+
"""CLIENT-SIDE parse of the batched free-text tags -> BatchSpanClassification. Sparse (a no-function span is
|
|
335
|
+
absent) and aligned by the EXPLICIT index in each tag, matching the schema's span_index contract."""
|
|
336
|
+
spans = [SpanFunctions(span_index=int(m.group(1)), functions=scores)
|
|
337
|
+
for m in _SPAN_TAG_RE.finditer(text) if (scores := _parse_scores(m.group(2)))]
|
|
338
|
+
return BatchSpanClassification(spans=spans)
|
|
339
|
+
|
|
340
|
+
|
|
341
|
+
def parse_clause_function_tags(text: str) -> ClauseFunctionClassification:
|
|
342
|
+
"""CLIENT-SIDE parse of the single-clause free-text tags -> ClauseFunctionClassification."""
|
|
343
|
+
m = _FUNCTIONS_TAG_RE.search(text)
|
|
344
|
+
return ClauseFunctionClassification(functions=_parse_scores(m.group(1)) if m else [])
|
|
345
|
+
|
|
346
|
+
|
|
347
|
+
class _TagClassifierRunnable:
|
|
348
|
+
"""A `build_structured`-shaped runnable (`.invoke`/`.ainvoke(prompt) -> the Pydantic contract`) that drives the
|
|
349
|
+
classifier CLIENT-SIDE: a plain free-text call (NO server guided decoding) whose tagged output is parsed by
|
|
350
|
+
`parse`. `.ainvoke` streams via `astream_text` (true wall-clock deadline + the stage `label`, issue 0005)."""
|
|
351
|
+
|
|
352
|
+
def __init__(self, model_id: str, *, instructions: str, parse: Any, label: str, max_tokens: int = 2048) -> None:
|
|
353
|
+
self._model_id = model_id
|
|
354
|
+
self._instructions = instructions
|
|
355
|
+
self._parse = parse
|
|
356
|
+
self._label = label
|
|
357
|
+
self._max_tokens = max_tokens
|
|
358
|
+
|
|
359
|
+
def invoke(self, prompt: Any, config: Any = None) -> Any: # config accepted for runnable-compat, unused
|
|
360
|
+
from rag_wright.models.seam import build_model
|
|
361
|
+
|
|
362
|
+
text = build_model(self._model_id, max_tokens=self._max_tokens).invoke(prompt + self._instructions).content
|
|
363
|
+
return self._parse(str(text))
|
|
364
|
+
|
|
365
|
+
async def ainvoke(self, prompt: Any, config: Any = None) -> Any:
|
|
366
|
+
from rag_wright.models.seam import astream_text
|
|
367
|
+
|
|
368
|
+
text = await astream_text(self._model_id, prompt + self._instructions,
|
|
369
|
+
max_tokens=self._max_tokens, label=self._label)
|
|
370
|
+
return self._parse(text)
|
|
371
|
+
|
|
372
|
+
|
|
373
|
+
def production_llm_clause_classifier(model_id: str) -> LlmClauseClassifier:
|
|
374
|
+
"""Wire the single-clause classifier over CLIENT-SIDE free-text tag parse (issue 0005: no server-side guided
|
|
375
|
+
decoding, which runs away on self-hosted Granite). Same `ClauseFunctionClassification` contract."""
|
|
376
|
+
return LlmClauseClassifier(_TagClassifierRunnable(
|
|
377
|
+
model_id, instructions=_CLAUSE_TAG_INSTRUCTIONS, parse=parse_clause_function_tags,
|
|
378
|
+
label="clause_function_classifier.classify"))
|
|
379
|
+
|
|
380
|
+
|
|
381
|
+
def production_batch_clause_classifier(model_id: str) -> LlmBatchClauseClassifier:
|
|
382
|
+
"""Wire the BATCHED (option B) classifier -- the ingestion default -- over CLIENT-SIDE free-text tag parse
|
|
383
|
+
(issue 0005). Same `BatchSpanClassification` contract; parsed by `parse_batch_span_tags`."""
|
|
384
|
+
return LlmBatchClauseClassifier(_TagClassifierRunnable(
|
|
385
|
+
model_id, instructions=_BATCH_TAG_INSTRUCTIONS, parse=parse_batch_span_tags,
|
|
386
|
+
label="clause_function_classifier.classify_spans"))
|
|
387
|
+
|
|
388
|
+
|
|
389
|
+
# --- T55 / SETFIT-SEG-1: trained SetFit soft-tagger (replaces the LLM classifier for ingestion latency) ----------
|
|
390
|
+
# Behind the SAME ClauseFunctionClassifier seam as LlmBatchClauseClassifier / LegalBertClauseAdapter -- a new
|
|
391
|
+
# IMPLEMENTATION of the already-registered `clause_function_classification` capability, no contract/API change.
|
|
392
|
+
# Function is a SOFT tag (ADR-0047), so emitting multiple tags per span is intended.
|
|
393
|
+
class SetFitClauseAdapter:
|
|
394
|
+
"""In-process ENSEMBLE of trained SetFit models. Each model on disk is a Sentence-Transformer body
|
|
395
|
+
(`model.safetensors` + configs) + a joblib-pickled sklearn head (`model_head.pkl`), so inference needs only
|
|
396
|
+
sentence-transformers + scikit-learn + joblib -- NO `setfit` dependency. `classify_spans` encodes every span
|
|
397
|
+
with each body, AVERAGES the per-class probabilities across the ensemble, and emits the top-k labels above
|
|
398
|
+
`threshold` as `FunctionScore` soft tags (primary-first). Span-level (ignores `chunk_text`, like
|
|
399
|
+
`LegalBertClauseAdapter`). Load-don't-retrain from local checkpoints; milliseconds/span, no network hop."""
|
|
400
|
+
|
|
401
|
+
def __init__(self, model_dirs, *, top_k: int = 3, threshold: float = 0.0,
|
|
402
|
+
hi: float = 0.6, mid: float = 0.3, device: str | None = None, batch_size: int = 32) -> None:
|
|
403
|
+
# DEFAULT = the finalized operating point: pure avg-prob TOP-3 (threshold 0.0) -> ~3.0 tags/span, the
|
|
404
|
+
# validated 50/52 classes >0.65 recall. A higher threshold trims tags but drops recall (e.g. 0.10 -> ~2.0
|
|
405
|
+
# tags, ~47/52); tune via RAG_SETFIT_THRESHOLD / RAG_SETFIT_TOPK only with a re-measured operating point.
|
|
406
|
+
import json
|
|
407
|
+
from pathlib import Path
|
|
408
|
+
|
|
409
|
+
import joblib
|
|
410
|
+
import numpy as np
|
|
411
|
+
from sentence_transformers import SentenceTransformer
|
|
412
|
+
|
|
413
|
+
self._np = np
|
|
414
|
+
self._top_k, self._threshold, self._hi, self._mid, self._batch = top_k, threshold, hi, mid, batch_size
|
|
415
|
+
self._models: list[tuple] = []
|
|
416
|
+
for d in model_dirs:
|
|
417
|
+
d = Path(d)
|
|
418
|
+
cfg = json.loads((d / "config_setfit.json").read_text())
|
|
419
|
+
body = SentenceTransformer(str(d), device=device)
|
|
420
|
+
head = joblib.load(d / "model_head.pkl")
|
|
421
|
+
self._models.append((body, head, bool(cfg.get("normalize_embeddings", False)),
|
|
422
|
+
[str(c) for c in head.classes_]))
|
|
423
|
+
if not self._models:
|
|
424
|
+
raise ValueError("SetFitClauseAdapter needs at least one model dir")
|
|
425
|
+
self._labels = sorted({lab for *_, cols in self._models for lab in cols}) # union label space (robust)
|
|
426
|
+
self._lab_idx = {lab: i for i, lab in enumerate(self._labels)}
|
|
427
|
+
|
|
428
|
+
def classify(self, clause_text: str) -> list[FunctionScore]:
|
|
429
|
+
return self.classify_spans("", [clause_text])[0]
|
|
430
|
+
|
|
431
|
+
async def aclassify_spans(self, chunk_text: str, span_texts: list[str],
|
|
432
|
+
*, sem: asyncio.Semaphore | None = None) -> list[list[FunctionScore]]:
|
|
433
|
+
# the ingestion pipeline calls aclassify_spans under a shared concurrency bound. SetFit is CPU-bound and
|
|
434
|
+
# in-process, so run the sync encode+predict off the event loop in a thread; the sem bounds in-flight work.
|
|
435
|
+
if sem is None:
|
|
436
|
+
return await asyncio.to_thread(self.classify_spans, chunk_text, span_texts)
|
|
437
|
+
async with sem:
|
|
438
|
+
return await asyncio.to_thread(self.classify_spans, chunk_text, span_texts)
|
|
439
|
+
|
|
440
|
+
def classify_spans(self, chunk_text: str, span_texts: list[str]) -> list[list[FunctionScore]]: # noqa: ARG002
|
|
441
|
+
if not span_texts:
|
|
442
|
+
return []
|
|
443
|
+
np = self._np
|
|
444
|
+
agg = np.zeros((len(span_texts), len(self._labels)))
|
|
445
|
+
for body, head, norm, cols in self._models:
|
|
446
|
+
proba = np.asarray(head.predict_proba(
|
|
447
|
+
body.encode(list(span_texts), normalize_embeddings=norm, batch_size=self._batch)), dtype=float)
|
|
448
|
+
for j, c in enumerate(cols):
|
|
449
|
+
idx = self._lab_idx.get(c)
|
|
450
|
+
if idx is not None:
|
|
451
|
+
agg[:, idx] += proba[:, j]
|
|
452
|
+
agg /= len(self._models)
|
|
453
|
+
out: list[list[FunctionScore]] = []
|
|
454
|
+
for row in agg:
|
|
455
|
+
scores: list[FunctionScore] = []
|
|
456
|
+
for j in np.argsort(-row)[: self._top_k]:
|
|
457
|
+
p = float(row[j])
|
|
458
|
+
if p < self._threshold:
|
|
459
|
+
break
|
|
460
|
+
canon = canonical_function(self._labels[int(j)])
|
|
461
|
+
if not canon:
|
|
462
|
+
continue
|
|
463
|
+
conf = (FunctionConfidence.HIGH if p >= self._hi
|
|
464
|
+
else FunctionConfidence.MEDIUM if p >= self._mid else FunctionConfidence.LOW)
|
|
465
|
+
scores.append(FunctionScore(function=canon, confidence=conf))
|
|
466
|
+
out.append(scores)
|
|
467
|
+
return out
|
|
468
|
+
|
|
469
|
+
|
|
470
|
+
def production_setfit_clause_classifier(model_root: str | None = None, **kwargs) -> SetFitClauseAdapter:
|
|
471
|
+
"""Wire the ENSEMBLE SetFit soft-tagger -- the ingestion default (replaces the LLM classifier for latency,
|
|
472
|
+
T55/SETFIT-SEG-1). Loads the 3 finalized checkpoints (LegalBERT + BGE-large + MPNet) from `model_root`
|
|
473
|
+
(env RAG_SETFIT_CLAUSE_DIR; default data/models/setfit_clause). In-process, no network hop, no `setfit` dep.
|
|
474
|
+
Tunables: RAG_SETFIT_TOPK, RAG_SETFIT_THRESHOLD, RAG_SETFIT_DEVICE."""
|
|
475
|
+
import os
|
|
476
|
+
from pathlib import Path
|
|
477
|
+
|
|
478
|
+
root = Path(model_root or os.getenv("RAG_SETFIT_CLAUSE_DIR", "data/models/setfit_clause"))
|
|
479
|
+
subdirs = [root / n for n in ("cap128b_legalbert", "cap128b_bge", "cap128b_mpnet")]
|
|
480
|
+
present = [d for d in subdirs if (d / "model_head.pkl").exists()]
|
|
481
|
+
if not present:
|
|
482
|
+
raise FileNotFoundError(
|
|
483
|
+
f"No SetFit clause checkpoints under {root} (expected cap128b_legalbert/bge/mpnet). "
|
|
484
|
+
"Download them from the model store, or set RAG_FUNCTION_CLASSIFIER=llm to use the LLM classifier.")
|
|
485
|
+
return SetFitClauseAdapter(
|
|
486
|
+
present,
|
|
487
|
+
top_k=int(os.getenv("RAG_SETFIT_TOPK", str(kwargs.pop("top_k", 3)))),
|
|
488
|
+
threshold=float(os.getenv("RAG_SETFIT_THRESHOLD", str(kwargs.pop("threshold", 0.0)))),
|
|
489
|
+
device=os.getenv("RAG_SETFIT_DEVICE") or kwargs.pop("device", None),
|
|
490
|
+
**kwargs)
|