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,337 @@
|
|
|
1
|
+
"""KG-2 (FR-C.6, ADR-0033/0028): per-clause typed property extraction into the unified contract KG.
|
|
2
|
+
|
|
3
|
+
The extraction MODEL is granite-4.2-8b (the Leg-C winner, GP-1B; no A/B -- DeepSeek is a KG-6 below-par
|
|
4
|
+
contingency only). The MECHANISM is the GP-1B recipe (`kg-extraction-recipe` Skill): docling-graph +
|
|
5
|
+
`ontology.clause_template.Clause` (the KG-1 bridge template) via the `capabilities.dg_extraction` seam.
|
|
6
|
+
|
|
7
|
+
This module holds the two pieces that turn a raw extraction into a gated, contract-shaped record:
|
|
8
|
+
|
|
9
|
+
1. `clause_to_record` -- the PURE, unit-tested adapter from the KG-1 template (`Clause`, typed fields:
|
|
10
|
+
enums + lists + nested odrl:Constraint models) to the existing property contract (`ClausePropertyRecord`
|
|
11
|
+
/ `PropertyAssertion`, T57a). Each filled field becomes one assertion (dimension, value) carrying
|
|
12
|
+
provenance (FR-S.4) + the span citation (FR-Q.6). The OTHER escape means "not asserted" -> dropped;
|
|
13
|
+
`CapBasis.cap_other` maps to the canonical `other` (the one documented KG-1 vocab divergence). Every
|
|
14
|
+
emitted closed value is in `property.py::CLOSED_VOCAB` (the KG-1 test pins the two vocabularies equal).
|
|
15
|
+
|
|
16
|
+
2. `DGClausePropertyExtractor` -- composes an (injectable) Clause-extraction fn with the adapter and the
|
|
17
|
+
deterministic grounding-judge gate (`property_grounding.reground`, ADR-0028): an EXTRACTED value on a
|
|
18
|
+
lexically-anchored dimension whose cue is absent from the text is downgraded to AMBIGUOUS. The extraction
|
|
19
|
+
fn is injected so the mapping + gate are testable with no model call; the live default is granite-4.2-8b.
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
from __future__ import annotations
|
|
23
|
+
|
|
24
|
+
import os
|
|
25
|
+
import re
|
|
26
|
+
from enum import Enum
|
|
27
|
+
from typing import Any, Callable, Optional
|
|
28
|
+
|
|
29
|
+
from rag_wright.contracts.function import canonical_function
|
|
30
|
+
from rag_wright.contracts.identifiers import ChunkId
|
|
31
|
+
from rag_wright.contracts.property import (
|
|
32
|
+
CLOSED_VOCAB,
|
|
33
|
+
FOLIO_CLAUSE_IRI,
|
|
34
|
+
ClausePropertyRecord,
|
|
35
|
+
PropertyAssertion,
|
|
36
|
+
PropertyDimension,
|
|
37
|
+
)
|
|
38
|
+
from rag_wright.contracts.provenance import ConfidenceTag, Provenance
|
|
39
|
+
from rag_wright.ontology._generated_vocab import VALUE_SYNONYMS
|
|
40
|
+
from rag_wright.ontology.clause_template import DamageType, ExceptionModel, _normalize_enum
|
|
41
|
+
from rag_wright.spans.property_grounding import reground
|
|
42
|
+
from rag_wright.spans.semantic_judge import asemantic_judge, semantic_judge
|
|
43
|
+
from rag_wright.spans.symbolic_validation import symbolic_validate
|
|
44
|
+
|
|
45
|
+
_D = PropertyDimension
|
|
46
|
+
_OTHER = "Other" # the compiler's auto-added OTHER escape sentinel -> "not asserted"
|
|
47
|
+
|
|
48
|
+
# scalar enum field on Clause -> the dimension it asserts
|
|
49
|
+
_SCALAR_ENUM_DIMS: dict[str, PropertyDimension] = {
|
|
50
|
+
"has_mutuality": _D.MUTUALITY,
|
|
51
|
+
"has_favorability": _D.FAVORABILITY,
|
|
52
|
+
"has_asymmetry": _D.PARTY_ASYMMETRY,
|
|
53
|
+
"has_warranty_scope": _D.WARRANTY_SCOPE,
|
|
54
|
+
"has_claim_scope": _D.CLAIM_SCOPE,
|
|
55
|
+
"has_ip_ownership": _D.IP_OWNERSHIP,
|
|
56
|
+
"has_renewal": _D.RENEWAL_MECHANISM,
|
|
57
|
+
"covers_party_scope": _D.COVERED_PARTIES,
|
|
58
|
+
"prohibits_solicit": _D.NONSOLICIT_TARGET,
|
|
59
|
+
"requires_duty": _D.PROCEDURAL,
|
|
60
|
+
# tier 3 -- CUAD-family extensions (KG-4)
|
|
61
|
+
"has_exclusivity_type": _D.EXCLUSIVITY_TYPE,
|
|
62
|
+
"has_right_of_first_type": _D.RIGHT_OF_FIRST_TYPE,
|
|
63
|
+
"has_restriction_scope": _D.RESTRICTION_SCOPE,
|
|
64
|
+
"has_coc_consent": _D.COC_CONSENT,
|
|
65
|
+
"has_assignment_consent": _D.ASSIGNMENT_CONSENT,
|
|
66
|
+
"has_escrow_release_trigger": _D.ESCROW_RELEASE_TRIGGER,
|
|
67
|
+
"has_mfn_scope": _D.MFN_SCOPE,
|
|
68
|
+
"has_termination_right": _D.TERMINATION_RIGHT,
|
|
69
|
+
"dispute_method": _D.DISPUTE_METHOD, # ADR-0049 (2): Dispute Resolution method
|
|
70
|
+
"royalty_basis": _D.ROYALTY_BASIS, # ADR-0049 (2): Royalties basis
|
|
71
|
+
"condition_type": _D.CONDITION_TYPE, # ADR-0049 (2): Condition Precedent kind
|
|
72
|
+
}
|
|
73
|
+
# open-valued CUAD dims: direct string fields on Clause -> dimension
|
|
74
|
+
_OPEN_STR_DIMS: dict[str, PropertyDimension] = {
|
|
75
|
+
"audit_frequency": _D.AUDIT_FREQUENCY,
|
|
76
|
+
"commitment_quantum": _D.COMMITMENT_QUANTUM,
|
|
77
|
+
"ld_trigger": _D.LD_TRIGGER,
|
|
78
|
+
}
|
|
79
|
+
# list enum field on Clause -> the (multi-valued) dimension it asserts
|
|
80
|
+
_LIST_ENUM_DIMS: dict[str, PropertyDimension] = {
|
|
81
|
+
"covers": _D.COVERED_SUBJECT, # issue 0040: CLOSED conduct vocab -> out-of-vocab drops (not verbatim-retained)
|
|
82
|
+
"collateral_type": _D.COLLATERAL_TYPE, # ADR-0049 (2): Security Interest collateral (multi-valued)
|
|
83
|
+
"force_majeure_event": _D.FORCE_MAJEURE_EVENT, # ADR-0049 (2): Force Majeure events (multi-valued)
|
|
84
|
+
"confidentiality_exception": _D.CONFIDENTIALITY_EXCEPTION, # ADR-0049 (2): NDA carve-outs (multi-valued)
|
|
85
|
+
}
|
|
86
|
+
# issue 0037: the OPEN descriptive list-dims -- field -> (dimension, closed-vocab enum). Captured VERBATIM on the
|
|
87
|
+
# Clause; canonicalized here (exact/keyword -> canonical value; else skos:broader synonym -> canonical; else the
|
|
88
|
+
# verbatim phrase is KEPT, never dropped to OTHER). These dims are unbounded in symbolic_validation + lexically
|
|
89
|
+
# grounded (ADR-0028), so a spurious value whose cue is absent from the text is still downgraded by the gate.
|
|
90
|
+
_OPEN_LIST_DIMS: dict[str, tuple[PropertyDimension, type[Enum]]] = {
|
|
91
|
+
"excepts": (_D.CARVE_OUT, ExceptionModel),
|
|
92
|
+
"prohibits_damage": (_D.DAMAGE_TYPE, DamageType),
|
|
93
|
+
}
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def _canonical_value(member: Any) -> str | None:
|
|
97
|
+
"""An enum member's canonical value, or None if it is the OTHER escape (not asserted).
|
|
98
|
+
`CapBasis.cap_other` -> canonical `other` (the documented KG-1 divergence)."""
|
|
99
|
+
value = member.value if isinstance(member, Enum) else str(member)
|
|
100
|
+
if value == _OTHER:
|
|
101
|
+
return None
|
|
102
|
+
return "other" if value == "cap_other" else value
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
def _clean(text: Optional[str]) -> str | None:
|
|
106
|
+
"""A non-empty stripped open-valued literal, or None."""
|
|
107
|
+
if text is None:
|
|
108
|
+
return None
|
|
109
|
+
stripped = text.strip()
|
|
110
|
+
return stripped or None
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
def _norm_key(s: str) -> str:
|
|
114
|
+
return re.sub(r"[^A-Za-z0-9]+", "", s).lower()
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
def _open_list_value(dim: PropertyDimension, enum_cls: type[Enum], raw: Any) -> str | None:
|
|
118
|
+
"""issue 0037: canonicalize one OPEN-dim item, keeping it VERBATIM when nothing matches (never OTHER-dropped).
|
|
119
|
+
Order: exact/keyword vocab match -> its canonical value; else an ontology skos:broader synonym -> the canonical
|
|
120
|
+
value (e.g. 'loss of profits' -> 'consequential' for damage_type); else the cleaned verbatim phrase."""
|
|
121
|
+
s = _clean(str(raw.value) if isinstance(raw, Enum) else str(raw))
|
|
122
|
+
if s is None:
|
|
123
|
+
return None
|
|
124
|
+
member = _normalize_enum(enum_cls, s, keyword_fallback=True) # -> a member, or OTHER if no match
|
|
125
|
+
canon = _canonical_value(member) # None iff OTHER
|
|
126
|
+
if canon is not None:
|
|
127
|
+
return canon
|
|
128
|
+
syn = VALUE_SYNONYMS.get(dim.value, {}).get(_norm_key(s)) # skos:broader synonym -> canonical
|
|
129
|
+
return syn if syn is not None else s # else keep verbatim
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
def clause_to_record(
|
|
133
|
+
clause: Any, *, chunk_id: ChunkId, function: str, span_id: str = ""
|
|
134
|
+
) -> ClausePropertyRecord:
|
|
135
|
+
"""Map an extracted `Clause` (KG-1 template) to a validated `ClausePropertyRecord` (pure; no gate, no
|
|
136
|
+
model). `function` is the KNOWN clause type (from the T56 classifier), not the LLM's `clause_type`.
|
|
137
|
+
Every assertion starts EXTRACTED; the grounding gate (applied by the extractor) downgrades the
|
|
138
|
+
ungrounded ones. Closed-vocab values are guaranteed in-vocabulary by the KG-1 template."""
|
|
139
|
+
prov = Provenance.of(chunk_id)
|
|
140
|
+
assertions: list[PropertyAssertion] = []
|
|
141
|
+
|
|
142
|
+
def add(dimension: PropertyDimension, value: str | None,
|
|
143
|
+
confidence: ConfidenceTag = ConfidenceTag.EXTRACTED) -> None:
|
|
144
|
+
if value is None or not str(value).strip():
|
|
145
|
+
return
|
|
146
|
+
assertions.append(
|
|
147
|
+
PropertyAssertion(
|
|
148
|
+
provenance=prov, confidence=confidence,
|
|
149
|
+
dimension=dimension, value=value, span_id=span_id,
|
|
150
|
+
)
|
|
151
|
+
)
|
|
152
|
+
|
|
153
|
+
for field, dim in _SCALAR_ENUM_DIMS.items():
|
|
154
|
+
add(dim, _canonical_value(getattr(clause, field, None)))
|
|
155
|
+
for field, dim in _LIST_ENUM_DIMS.items():
|
|
156
|
+
for member in getattr(clause, field, None) or []:
|
|
157
|
+
add(dim, _canonical_value(member))
|
|
158
|
+
for field, (dim, enum_cls) in _OPEN_LIST_DIMS.items(): # issue 0037: verbatim-retaining open descriptive dims
|
|
159
|
+
vocab = CLOSED_VOCAB.get(dim, frozenset())
|
|
160
|
+
for raw in getattr(clause, field, None) or []:
|
|
161
|
+
val = _open_list_value(dim, enum_cls, raw)
|
|
162
|
+
if val is None:
|
|
163
|
+
continue
|
|
164
|
+
# a canonical (in-vocab) value is EXTRACTED; a retained VERBATIM tail value is admissible only as the
|
|
165
|
+
# AMBIGUOUS "other" escape (PropertyAssertion contract) -- kept + flagged, never dropped (issue 0037).
|
|
166
|
+
add(dim, val, ConfidenceTag.EXTRACTED if val in vocab else ConfidenceTag.AMBIGUOUS)
|
|
167
|
+
for field, dim in _OPEN_STR_DIMS.items(): # open-valued CUAD dims (direct string fields)
|
|
168
|
+
add(dim, _clean(getattr(clause, field, None)))
|
|
169
|
+
|
|
170
|
+
caps = getattr(clause, "caps", None)
|
|
171
|
+
if caps is not None:
|
|
172
|
+
add(_D.CAP_BASIS, _canonical_value(caps.cap_basis))
|
|
173
|
+
add(_D.CAP_QUANTUM, _clean(caps.cap_quantum)) # open-valued
|
|
174
|
+
|
|
175
|
+
bound = getattr(clause, "bounded_by", None)
|
|
176
|
+
if bound is not None:
|
|
177
|
+
duration = _clean(bound.temporal_duration)
|
|
178
|
+
if duration is not None:
|
|
179
|
+
kind = (bound.temporal_kind or "").strip().lower()
|
|
180
|
+
add(_D.NOTICE_PERIOD if kind == "notice_period" else _D.TEMPORAL_BOUND, duration)
|
|
181
|
+
|
|
182
|
+
law = getattr(clause, "governed_by", None)
|
|
183
|
+
if law is not None:
|
|
184
|
+
add(_D.JURISDICTION, _clean(law.jurisdiction_name)) # open-valued
|
|
185
|
+
add(_D.LAW_MULTIPLICITY, _canonical_value(law.law_multiplicity))
|
|
186
|
+
|
|
187
|
+
return ClausePropertyRecord(
|
|
188
|
+
clause_id=str(chunk_id),
|
|
189
|
+
function=canonical_function(function) or function,
|
|
190
|
+
folio_iri=FOLIO_CLAUSE_IRI.get(function, ""),
|
|
191
|
+
span_id=span_id, # the operative span (1:1) -- carried even when the clause has no properties
|
|
192
|
+
assertions=assertions,
|
|
193
|
+
)
|
|
194
|
+
|
|
195
|
+
|
|
196
|
+
ClauseExtractFn = Callable[[str], Any]
|
|
197
|
+
"""text -> an extracted `Clause` instance (or None). Injected so the adapter + gate test with no model."""
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
class DGClausePropertyExtractor:
|
|
201
|
+
"""Extract a clause's typed properties end to end: run the Clause extraction (granite-4.2-8b via
|
|
202
|
+
docling-graph, injected), adapt to `ClausePropertyRecord`, then apply the deterministic
|
|
203
|
+
grounding-judge gate (ADR-0028). Matches the `PropertyExtractor` call shape (T57b) so it drops into
|
|
204
|
+
the ingestion driver. A None extraction yields an empty (but valid) record for that clause."""
|
|
205
|
+
|
|
206
|
+
def __init__(self, extract_fn: ClauseExtractFn, *, semantic_judge_fn: Any = None,
|
|
207
|
+
aextract_fn: Any = None, asemantic_judge_fn: Any = None) -> None:
|
|
208
|
+
self._extract = extract_fn
|
|
209
|
+
# ADR-0040 Layer 3: an optional LLM semantic judge. Injected (default None -> deterministic-only) so
|
|
210
|
+
# existing callers + hermetic tests are unaffected; the production pipeline wires the real granite judge.
|
|
211
|
+
self._semantic_judge_fn = semantic_judge_fn
|
|
212
|
+
# ASYNC-B2b (ADR-0057): the async twins (extraction on the async docling-graph seam + async judge).
|
|
213
|
+
self._aextract = aextract_fn
|
|
214
|
+
self._asemantic_judge_fn = asemantic_judge_fn
|
|
215
|
+
|
|
216
|
+
def _empty(self, chunk_id: ChunkId, function: str, span_id: str = "") -> ClausePropertyRecord:
|
|
217
|
+
return ClausePropertyRecord(
|
|
218
|
+
clause_id=str(chunk_id), function=canonical_function(function) or function,
|
|
219
|
+
folio_iri=FOLIO_CLAUSE_IRI.get(function, ""),
|
|
220
|
+
span_id=span_id, assertions=[]) # carry the operative-span anchor even when property-less (ADR-0025)
|
|
221
|
+
|
|
222
|
+
def _grounded(self, clause: Any, *, chunk_id: ChunkId, function: str, text: str, span_id: str
|
|
223
|
+
) -> ClausePropertyRecord:
|
|
224
|
+
# ADR-0028 lexical grounding gate, then ADR-0040 symbolic (function->dimension applicability) gate.
|
|
225
|
+
record = clause_to_record(clause, chunk_id=chunk_id, function=function, span_id=span_id)
|
|
226
|
+
return symbolic_validate(reground(record, text))
|
|
227
|
+
|
|
228
|
+
def __call__(
|
|
229
|
+
self, *, chunk_id: ChunkId, function: str, text: str, span_id: str = ""
|
|
230
|
+
) -> ClausePropertyRecord:
|
|
231
|
+
clause = self._extract(text)
|
|
232
|
+
if clause is None:
|
|
233
|
+
return self._empty(chunk_id, function, span_id)
|
|
234
|
+
record = self._grounded(clause, chunk_id=chunk_id, function=function, text=text, span_id=span_id)
|
|
235
|
+
if self._semantic_judge_fn is not None: # ADR-0040 Layer 3 LLM semantic gate (production only)
|
|
236
|
+
record = semantic_judge(record, text, self._semantic_judge_fn)
|
|
237
|
+
return record
|
|
238
|
+
|
|
239
|
+
async def aextract(
|
|
240
|
+
self, *, chunk_id: ChunkId, function: str, text: str, span_id: str = ""
|
|
241
|
+
) -> ClausePropertyRecord:
|
|
242
|
+
"""ASYNC-B2b (ADR-0057): the async twin of `__call__`. Runs the docling-graph clause extraction on the
|
|
243
|
+
async seam (true wall-clock deadline), the same deterministic grounding/symbolic gates, then the async
|
|
244
|
+
Layer-3 semantic judge. Same contract as `__call__`."""
|
|
245
|
+
clause = await self._aextract(text)
|
|
246
|
+
if clause is None:
|
|
247
|
+
return self._empty(chunk_id, function, span_id)
|
|
248
|
+
record = self._grounded(clause, chunk_id=chunk_id, function=function, text=text, span_id=span_id)
|
|
249
|
+
if self._asemantic_judge_fn is not None:
|
|
250
|
+
record = await asemantic_judge(record, text, self._asemantic_judge_fn)
|
|
251
|
+
return record
|
|
252
|
+
|
|
253
|
+
|
|
254
|
+
def granite_clause_extractor(model: Any = None, *, semantic_judge_fn: Any = None,
|
|
255
|
+
asemantic_judge_fn: Any = None, list_model: str | None = None,
|
|
256
|
+
samples: int | None = None) -> DGClausePropertyExtractor:
|
|
257
|
+
"""The live default: granite-4.2-8b via the SELECTED serving backend (`default_extraction_model` reads
|
|
258
|
+
`RAG_SERVING` -> vLLM-Granite in product, OpenRouter-Granite in dev; MS1-3, ADR-0079). Pass a different
|
|
259
|
+
`ExtractionModel` to override, or a `semantic_judge_fn`/`asemantic_judge_fn` to enable the ADR-0040 Layer-3
|
|
260
|
+
gate (sync/async). ASYNC-B2b wires the async extraction seam (`aextract_clause`) so `aextract` gets the true
|
|
261
|
+
wall-clock deadline.
|
|
262
|
+
|
|
263
|
+
TAGPARSE-INGEST-1b: `RAG_INGEST_CLAUSE_EXTRACTOR` selects the Clause-producing step -- `tagparse` (DEFAULT
|
|
264
|
+
since ADR-0081: function-independent thematic tag-parse groups) or `docling` (the legacy docling-graph
|
|
265
|
+
server-side-JSON path, kept for rollback). BOTH feed the SAME downstream (adapt to ClausePropertyRecord +
|
|
266
|
+
ADR-0028 grounding + ADR-0040 symbolic gate). tagparse is the default because docling hard-crashes ~89% of
|
|
267
|
+
real CUAD clauses (grounded A/B, 45 clauses: docling success 0.11 vs tagparse 1.00). NOTE: tagparse issues one
|
|
268
|
+
LLM call per thematic GROUP (~8/clause) vs docling's ~1; cost is reduced via `is_extractable_span` (fewer spans),
|
|
269
|
+
NOT by pruning groups or batching clauses -- both were measured in issue 0036 at a ~15-18% property-recall loss.
|
|
270
|
+
|
|
271
|
+
`list_model` (ARGUMENT; else env `RAG_INGEST_LIST_MODEL`; else the profile general model) is the SECOND model
|
|
272
|
+
for the cross-model list union on list-bearing groups -- exposed here (like `model`) so the caller configures
|
|
273
|
+
it explicitly; pass `"off"` to disable. `samples` (else env) sets same-model multi-sample union."""
|
|
274
|
+
from rag_wright.capabilities.dg_extraction import aextract_clause, default_extraction_model, extract_clause
|
|
275
|
+
|
|
276
|
+
chosen = model or default_extraction_model("clause-extract")
|
|
277
|
+
if os.getenv("RAG_INGEST_CLAUSE_EXTRACTOR", "tagparse").strip().lower() == "tagparse":
|
|
278
|
+
from rag_wright.spans.tag_clause_extractor import atag_extract_clause, tag_extract_clause
|
|
279
|
+
model_id = chosen.model
|
|
280
|
+
return DGClausePropertyExtractor(
|
|
281
|
+
lambda text: tag_extract_clause(text, model_id, list_model=list_model, samples=samples),
|
|
282
|
+
aextract_fn=lambda text: atag_extract_clause(text, model_id, list_model=list_model, samples=samples),
|
|
283
|
+
semantic_judge_fn=semantic_judge_fn, asemantic_judge_fn=asemantic_judge_fn,
|
|
284
|
+
)
|
|
285
|
+
return DGClausePropertyExtractor(
|
|
286
|
+
lambda text: extract_clause(text, chosen),
|
|
287
|
+
aextract_fn=lambda text: aextract_clause(text, chosen),
|
|
288
|
+
semantic_judge_fn=semantic_judge_fn, asemantic_judge_fn=asemantic_judge_fn,
|
|
289
|
+
)
|
|
290
|
+
|
|
291
|
+
|
|
292
|
+
class ClassifierPropertyExtractor:
|
|
293
|
+
"""CLS-C (ADR-0115): the classifier-first Step-3a property extractor. Runs the FUNCTION-INDEPENDENT
|
|
294
|
+
`HybridPropertyExtractor` (classifiers for the 21 covered dims + ONE residual LLM call for the 7 numeric/open
|
|
295
|
+
dims) and then the SAME record-level gates as `DGClausePropertyExtractor`: ADR-0028 `reground` -> ADR-0040
|
|
296
|
+
`symbolic_validate` -> optional Layer-3 `semantic_judge` (sync) / `asemantic_judge` (async). This REPLACES the
|
|
297
|
+
full-LLM tag-parse extraction for these dims -- it is the decided path, not a toggle over an LLM fallback.
|
|
298
|
+
Same `PropertyExtractor` Protocol (`__call__` + `aextract`), so nothing downstream changes."""
|
|
299
|
+
|
|
300
|
+
def __init__(self, hybrid: Any, *, semantic_judge_fn: Any = None, asemantic_judge_fn: Any = None) -> None:
|
|
301
|
+
self._hybrid = hybrid
|
|
302
|
+
self._semantic_judge_fn = semantic_judge_fn
|
|
303
|
+
self._asemantic_judge_fn = asemantic_judge_fn
|
|
304
|
+
|
|
305
|
+
def _gate(self, record: ClausePropertyRecord, text: str) -> ClausePropertyRecord:
|
|
306
|
+
return symbolic_validate(reground(record, text)) # ADR-0028 lexical, then ADR-0040 symbolic
|
|
307
|
+
|
|
308
|
+
def __call__(self, *, chunk_id: ChunkId, function: str, text: str, span_id: str = "",
|
|
309
|
+
functions: tuple[str, ...] = ()) -> ClausePropertyRecord:
|
|
310
|
+
record = self._gate(self._hybrid(chunk_id=chunk_id, function=function, text=text, span_id=span_id,
|
|
311
|
+
functions=functions), text)
|
|
312
|
+
if self._semantic_judge_fn is not None:
|
|
313
|
+
record = semantic_judge(record, text, self._semantic_judge_fn)
|
|
314
|
+
return record
|
|
315
|
+
|
|
316
|
+
async def aextract(self, *, chunk_id: ChunkId, function: str, text: str,
|
|
317
|
+
span_id: str = "", functions: tuple[str, ...] = ()) -> ClausePropertyRecord:
|
|
318
|
+
raw = await self._hybrid.aextract(chunk_id=chunk_id, function=function, text=text, span_id=span_id,
|
|
319
|
+
functions=functions)
|
|
320
|
+
record = self._gate(raw, text)
|
|
321
|
+
if self._asemantic_judge_fn is not None:
|
|
322
|
+
record = await asemantic_judge(record, text, self._asemantic_judge_fn)
|
|
323
|
+
return record
|
|
324
|
+
|
|
325
|
+
|
|
326
|
+
def classifier_property_extractor(*, registry: Any = None, runnable: Any = None, model_id: Optional[str] = None,
|
|
327
|
+
semantic_judge_fn: Any = None, asemantic_judge_fn: Any = None,
|
|
328
|
+
classifier_fn: Any = None) -> ClassifierPropertyExtractor:
|
|
329
|
+
"""Build the CLS-C classifier-first Step-3a extractor. The classifier LANE comes from `classifier_fn` when given
|
|
330
|
+
(EP-RT-7: the ingestion pipeline passes the `clause_property_classification` capability dispatch, so the fleet is
|
|
331
|
+
invoked through the capability -- the single production path), else from the local `registry` fleet. `runnable`
|
|
332
|
+
defaults to the structured seam for the residual numeric call; the ADR-0028/0040 gates + the judge apply unchanged."""
|
|
333
|
+
from rag_wright.spans.property_extractor import HybridPropertyExtractor
|
|
334
|
+
|
|
335
|
+
hybrid = HybridPropertyExtractor(registry, runnable=runnable, model_id=model_id, classifier_fn=classifier_fn)
|
|
336
|
+
return ClassifierPropertyExtractor(hybrid, semantic_judge_fn=semantic_judge_fn,
|
|
337
|
+
asemantic_judge_fn=asemantic_judge_fn)
|
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
"""T56 (FR-R, ADR-0025): CUAD gold -> operative-span function labels.
|
|
2
|
+
|
|
3
|
+
CUAD (`CUAD_v1.json`, SQuAD-style) annotates each contract with answer spans per clause type. To train the
|
|
4
|
+
function classifier at the SAME granularity we infer on (operative spans, T55), we segment each CUAD contract's
|
|
5
|
+
text with the operative-span segmenter and label each span by the clause type of the CUAD answer span it
|
|
6
|
+
overlaps most (else NONE). Same segmenter, same unit, both sides -- no train/infer granularity mismatch.
|
|
7
|
+
|
|
8
|
+
The clause type is parsed from the CUAD question (`... related to "<Type>" ...`); types are the canonical CUAD
|
|
9
|
+
label names (== `ClauseCategory` values, ADR-0002). A span overlapping answers of two types takes the one with
|
|
10
|
+
the greater character overlap (single-label). The contract id is kept so callers can split train/test by
|
|
11
|
+
contract (no leakage across the split).
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from __future__ import annotations
|
|
15
|
+
|
|
16
|
+
import json
|
|
17
|
+
import re
|
|
18
|
+
from collections.abc import Iterator
|
|
19
|
+
from pathlib import Path
|
|
20
|
+
|
|
21
|
+
from pydantic import BaseModel
|
|
22
|
+
|
|
23
|
+
from rag_wright.spans.function_classifier import NONE_LABEL
|
|
24
|
+
from rag_wright.spans.segment import segment_clause
|
|
25
|
+
|
|
26
|
+
_TYPE_IN_QUESTION = re.compile(r'related to\s+"([^"]+)"')
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class CuadAnswer(BaseModel):
|
|
30
|
+
clause_type: str
|
|
31
|
+
start: int # char offset into the contract context
|
|
32
|
+
text: str
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class CuadContract(BaseModel):
|
|
36
|
+
contract_id: str
|
|
37
|
+
context: str # the full contract text
|
|
38
|
+
answers: list[CuadAnswer]
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class LabeledSpan(BaseModel):
|
|
42
|
+
contract_id: str
|
|
43
|
+
text: str
|
|
44
|
+
label: str # a CUAD clause type, or NONE
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def parse_cuad(path: Path) -> Iterator[CuadContract]:
|
|
48
|
+
"""Yield one `CuadContract` per CUAD contract (id, full context, its labeled answer spans)."""
|
|
49
|
+
payload = json.loads(Path(path).read_text(encoding="utf-8"))
|
|
50
|
+
for entry in payload["data"]:
|
|
51
|
+
contract_id = str(entry.get("title") or entry.get("id") or "")
|
|
52
|
+
for para in entry["paragraphs"]:
|
|
53
|
+
context = para["context"]
|
|
54
|
+
answers: list[CuadAnswer] = []
|
|
55
|
+
for qa in para["qas"]:
|
|
56
|
+
m = _TYPE_IN_QUESTION.search(qa.get("question", ""))
|
|
57
|
+
if not m:
|
|
58
|
+
continue
|
|
59
|
+
clause_type = m.group(1).strip()
|
|
60
|
+
for ans in qa.get("answers", []):
|
|
61
|
+
text = str(ans["text"])
|
|
62
|
+
if text:
|
|
63
|
+
answers.append(CuadAnswer(clause_type=clause_type, start=int(ans["answer_start"]), text=text))
|
|
64
|
+
yield CuadContract(contract_id=contract_id, context=context, answers=answers)
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def label_operative_spans(contract: CuadContract) -> list[LabeledSpan]:
|
|
68
|
+
"""Segment the contract into operative spans; label each by the max-overlapping CUAD answer's type (else NONE)."""
|
|
69
|
+
spans = segment_clause(contract.contract_id, contract.context)
|
|
70
|
+
out: list[LabeledSpan] = []
|
|
71
|
+
for s in spans:
|
|
72
|
+
best_label, best_overlap = NONE_LABEL, 0
|
|
73
|
+
for ans in contract.answers:
|
|
74
|
+
a_end = ans.start + len(ans.text)
|
|
75
|
+
overlap = min(s.end, a_end) - max(s.start, ans.start)
|
|
76
|
+
if overlap > best_overlap:
|
|
77
|
+
best_overlap, best_label = overlap, ans.clause_type
|
|
78
|
+
text = s.text.strip()
|
|
79
|
+
if text:
|
|
80
|
+
out.append(LabeledSpan(contract_id=contract.contract_id, text=text, label=best_label))
|
|
81
|
+
return out
|
|
@@ -0,0 +1,158 @@
|
|
|
1
|
+
"""CLS-A (FR-I.4): per-dimension property CLASSIFIERS behind one seam, for Step-3a "Extract Clauses".
|
|
2
|
+
|
|
3
|
+
A `DimClassifier` maps a span's text to top-k (value, probability) for ONE closed-vocab PropertyDimension, replacing
|
|
4
|
+
the per-span LLM call for that dimension. Two runtimes implement the seam:
|
|
5
|
+
- `SetFitDimClassifier` — a SentenceTransformer body + joblib sklearn head on disk (NO `setfit` dep at serve time),
|
|
6
|
+
the 14 SetFit keepers (13 LegalBERT + 1 all-mpnet).
|
|
7
|
+
- `LayaDimClassifier` — a fine-tuned Laya (ModernBERT-large RL decision) checkpoint via the `laya` package, the 1
|
|
8
|
+
keeper (termination_right) SetFit couldn't crack.
|
|
9
|
+
|
|
10
|
+
STANDING serving philosophy (same as the model-profile seam for the LLM): device is NOT pinned — use a GPU if one is
|
|
11
|
+
available (CUDA, else Apple MPS), else CPU; the caller may override. Load ONCE (these are heavy, esp. Laya ~820MB).
|
|
12
|
+
Nothing upstream changes: the hybrid extractor (CLS-B) composes these with the LLM behind the `PropertyExtractor`
|
|
13
|
+
Protocol.
|
|
14
|
+
"""
|
|
15
|
+
from __future__ import annotations
|
|
16
|
+
|
|
17
|
+
from pathlib import Path
|
|
18
|
+
from typing import Any, Optional, Protocol, runtime_checkable
|
|
19
|
+
|
|
20
|
+
from rag_wright.contracts.property import PropertyDimension
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def auto_device(pref: Optional[str] = None) -> str:
|
|
24
|
+
"""GPU-if-available-else-CPU: honor an explicit choice, else CUDA -> MPS -> CPU. Never pin a device in code."""
|
|
25
|
+
if pref:
|
|
26
|
+
return pref
|
|
27
|
+
try:
|
|
28
|
+
import torch
|
|
29
|
+
if torch.cuda.is_available():
|
|
30
|
+
return "cuda"
|
|
31
|
+
if getattr(torch.backends, "mps", None) is not None and torch.backends.mps.is_available():
|
|
32
|
+
return "mps"
|
|
33
|
+
except Exception: # noqa: BLE001 - torch import/probe failure -> CPU is the safe floor
|
|
34
|
+
pass
|
|
35
|
+
return "cpu"
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
@runtime_checkable
|
|
39
|
+
class DimClassifier(Protocol):
|
|
40
|
+
"""Span text -> ranked (value, probability) for one dimension. `classify` returns top-k, highest first."""
|
|
41
|
+
|
|
42
|
+
dim: PropertyDimension
|
|
43
|
+
|
|
44
|
+
def classify(self, span_text: str) -> list[tuple[str, float]]: ...
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
class SetFitDimClassifier:
|
|
48
|
+
"""A SetFit keeper served as body+head (no `setfit` dep): `body.encode` -> `head.predict_proba` -> top-k. Honors
|
|
49
|
+
the body's configured normalization; `head.classes_` columns map to the dimension's values."""
|
|
50
|
+
|
|
51
|
+
def __init__(self, dim: PropertyDimension, model_dir, *, top_k: int = 1, device: Optional[str] = None,
|
|
52
|
+
batch_size: int = 32) -> None:
|
|
53
|
+
import json
|
|
54
|
+
from pathlib import Path
|
|
55
|
+
|
|
56
|
+
import joblib
|
|
57
|
+
from sentence_transformers import SentenceTransformer
|
|
58
|
+
|
|
59
|
+
self.dim = dim
|
|
60
|
+
self._top_k = top_k
|
|
61
|
+
self._batch = batch_size
|
|
62
|
+
d = Path(model_dir)
|
|
63
|
+
self._device = auto_device(device)
|
|
64
|
+
self._body = SentenceTransformer(str(d), device=self._device)
|
|
65
|
+
self._head = joblib.load(d / "model_head.pkl")
|
|
66
|
+
self._labels = [str(c) for c in self._head.classes_]
|
|
67
|
+
# SetFit stores whether the body normalizes embeddings; mirror it so serve-time matches train-time.
|
|
68
|
+
cfg = d / "config_setfit.json"
|
|
69
|
+
self._normalize = bool(json.loads(cfg.read_text()).get("normalize_embeddings", True)) if cfg.exists() else True
|
|
70
|
+
|
|
71
|
+
def classify(self, span_text: str) -> list[tuple[str, float]]:
|
|
72
|
+
import numpy as np
|
|
73
|
+
emb = self._body.encode([span_text], normalize_embeddings=self._normalize, batch_size=self._batch)
|
|
74
|
+
proba = np.asarray(self._head.predict_proba(emb), dtype=float)[0]
|
|
75
|
+
order = np.argsort(-proba)[: self._top_k]
|
|
76
|
+
return [(self._labels[i], float(proba[i])) for i in order]
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
class LayaDimClassifier:
|
|
80
|
+
"""A fine-tuned Laya checkpoint served via the `laya` package: one `choice` question over the dimension's values,
|
|
81
|
+
a single non-autoregressive forward pass. Device auto (CUDA/MPS/CPU) inside laya.load; loaded once."""
|
|
82
|
+
|
|
83
|
+
def __init__(self, dim: PropertyDimension, model_dir, *, instructions: str, criteria: dict[str, str],
|
|
84
|
+
top_k: int = 1, device: Optional[str] = None, agent: Any = None) -> None:
|
|
85
|
+
self.dim = dim
|
|
86
|
+
self._top_k = top_k
|
|
87
|
+
self._q = {dim.value: {"type": "choice", "instructions": instructions, "criteria": criteria}}
|
|
88
|
+
# A GROUP checkpoint serves several dims from ONE model: pass a shared pre-loaded `agent` so it is loaded
|
|
89
|
+
# ONCE (a ModernBERT-large is ~820MB) and every dim in the group queries the same forward-pass backbone.
|
|
90
|
+
if agent is not None:
|
|
91
|
+
self._agent = agent
|
|
92
|
+
else:
|
|
93
|
+
import laya # laya.load device=None -> CUDA else MPS else CPU; pass an explicit override if given.
|
|
94
|
+
self._agent = laya.load(str(model_dir), device=device)
|
|
95
|
+
|
|
96
|
+
def classify(self, span_text: str) -> list[tuple[str, float]]:
|
|
97
|
+
ans = self._agent.predict(span_text, self._q)["answers"][self.dim.value]
|
|
98
|
+
probs = ans.get("probabilities", {})
|
|
99
|
+
ranked = sorted(probs.items(), key=lambda kv: -kv[1])[: self._top_k]
|
|
100
|
+
return [(k, float(v)) for k, v in ranked]
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
class DimClassifierRegistry:
|
|
104
|
+
"""Loads the configured per-dimension classifiers ONCE and serves them by dimension. Missing dims (uncovered)
|
|
105
|
+
return None; the hybrid extractor (CLS-B/C) extracts ONLY its 7 numeric/open `RESIDUAL_LLM_DIMS` via the LLM,
|
|
106
|
+
never an uncovered classifier dim (the corpus-starved dims are left unextracted until CLS-F sources data)."""
|
|
107
|
+
|
|
108
|
+
def __init__(self, classifiers: dict[PropertyDimension, DimClassifier]) -> None:
|
|
109
|
+
self._by_dim = dict(classifiers)
|
|
110
|
+
|
|
111
|
+
def get(self, dim: PropertyDimension) -> Optional[DimClassifier]:
|
|
112
|
+
return self._by_dim.get(dim)
|
|
113
|
+
|
|
114
|
+
def covers(self, dim: PropertyDimension) -> bool:
|
|
115
|
+
return dim in self._by_dim
|
|
116
|
+
|
|
117
|
+
@property
|
|
118
|
+
def dims(self) -> list[PropertyDimension]:
|
|
119
|
+
return list(self._by_dim)
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
# CLS-D: the production 21-dim best-of-both fleet. `dim_fleet.json` (committed config) maps each dim to its
|
|
123
|
+
# framework + local model dir + serving params; model weights live under `data/models/` (gitignored). A LAYA
|
|
124
|
+
# group checkpoint serves several dims -> load each unique model ONCE and share the agent.
|
|
125
|
+
_FLEET_CONFIG = Path(__file__).with_name("dim_fleet.json")
|
|
126
|
+
_MODELS_DIR = Path(__file__).resolve().parents[3] / "data" / "models"
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
def load_dim_registry(config_path=None, *, models_dir=None, device: Optional[str] = None) -> "DimClassifierRegistry":
|
|
130
|
+
"""Load the committed fleet: shared Laya group agents (loaded once) + per-dim SetFit models, behind the
|
|
131
|
+
device-agnostic seam. Raises FileNotFoundError with the missing path if a checkpoint has not been fetched."""
|
|
132
|
+
import json
|
|
133
|
+
|
|
134
|
+
cfg = json.loads(Path(config_path or _FLEET_CONFIG).read_text())
|
|
135
|
+
base = Path(models_dir or _MODELS_DIR)
|
|
136
|
+
laya_agents: dict[str, Any] = {} # model dir name -> loaded laya agent (shared across its dims)
|
|
137
|
+
classifiers: dict[PropertyDimension, DimClassifier] = {}
|
|
138
|
+
for dim_str, spec in cfg.items():
|
|
139
|
+
dim = PropertyDimension(dim_str)
|
|
140
|
+
top_k = int(spec.get("top_k", 1))
|
|
141
|
+
if spec["framework"] == "laya":
|
|
142
|
+
mdir = base / "laya" / spec["model"]
|
|
143
|
+
if not mdir.exists():
|
|
144
|
+
raise FileNotFoundError(f"laya checkpoint not fetched: {mdir}")
|
|
145
|
+
agent = laya_agents.get(spec["model"])
|
|
146
|
+
if agent is None:
|
|
147
|
+
import laya
|
|
148
|
+
agent = laya.load(str(mdir), device=device) # device=None -> CUDA/MPS/CPU
|
|
149
|
+
laya_agents[spec["model"]] = agent
|
|
150
|
+
q = spec["question"]
|
|
151
|
+
classifiers[dim] = LayaDimClassifier(dim, mdir, instructions=q["instructions"],
|
|
152
|
+
criteria=q["criteria"], top_k=top_k, agent=agent)
|
|
153
|
+
else:
|
|
154
|
+
mdir = base / "setfit" / spec["model"]
|
|
155
|
+
if not mdir.exists():
|
|
156
|
+
raise FileNotFoundError(f"setfit checkpoint not fetched: {mdir}")
|
|
157
|
+
classifiers[dim] = SetFitDimClassifier(dim, mdir, top_k=top_k, device=device)
|
|
158
|
+
return DimClassifierRegistry(classifiers)
|