akasha-model 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.
- akasha_model/__init__.py +183 -0
- akasha_model/audit.py +185 -0
- akasha_model/authority.py +163 -0
- akasha_model/calibration.py +87 -0
- akasha_model/catalog.py +72 -0
- akasha_model/data.py +374 -0
- akasha_model/decision.py +312 -0
- akasha_model/ensemble.py +133 -0
- akasha_model/eval.py +84 -0
- akasha_model/gate.py +212 -0
- akasha_model/gate_multitask.py +278 -0
- akasha_model/host.py +123 -0
- akasha_model/model.py +147 -0
- akasha_model/multitask.py +422 -0
- akasha_model/multitask_calibrate.py +145 -0
- akasha_model/multitask_eval.py +285 -0
- akasha_model/multitask_train.py +120 -0
- akasha_model/outcomes.py +248 -0
- akasha_model/predict.py +62 -0
- akasha_model/primitives.py +212 -0
- akasha_model/rewards.py +53 -0
- akasha_model/rlcd.py +244 -0
- akasha_model/schemas/gate-audit-envelope.schema.json +44 -0
- akasha_model/schemas/host-outcome.schema.json +43 -0
- akasha_model/schemas/tool-call-plan.schema.json +44 -0
- akasha_model/schemas/tool-gate-request.schema.json +101 -0
- akasha_model/sequence.py +199 -0
- akasha_model/tool_calling.py +257 -0
- akasha_model/train.py +117 -0
- akasha_model/typed_decisions.py +301 -0
- akasha_model/vision.py +228 -0
- akasha_model/wire.py +281 -0
- akasha_model-0.2.0.dist-info/METADATA +503 -0
- akasha_model-0.2.0.dist-info/RECORD +38 -0
- akasha_model-0.2.0.dist-info/WHEEL +5 -0
- akasha_model-0.2.0.dist-info/entry_points.txt +13 -0
- akasha_model-0.2.0.dist-info/licenses/LICENSE +21 -0
- akasha_model-0.2.0.dist-info/top_level.txt +1 -0
akasha_model/__init__.py
ADDED
|
@@ -0,0 +1,183 @@
|
|
|
1
|
+
"""Small one-pass text and visual option scorers.
|
|
2
|
+
|
|
3
|
+
Path A (planner + host + outcomes + wire) imports without torch. Path B scorers,
|
|
4
|
+
training, and vision load lazily and require the ``torch`` extra::
|
|
5
|
+
|
|
6
|
+
pip install "akasha-model[torch]"
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
from typing import Any
|
|
12
|
+
|
|
13
|
+
from .primitives import (
|
|
14
|
+
ChoiceQuestion,
|
|
15
|
+
ChoiceResult,
|
|
16
|
+
DecisionRequest,
|
|
17
|
+
NoulQuestion,
|
|
18
|
+
NoulResult,
|
|
19
|
+
OptionSpec,
|
|
20
|
+
ScoreLevel,
|
|
21
|
+
ScoreQuestion,
|
|
22
|
+
ScoreResult,
|
|
23
|
+
)
|
|
24
|
+
from .authority import AuthorityProfile
|
|
25
|
+
from .catalog import CatalogPolicy, check_catalog_policy
|
|
26
|
+
from .gate import (
|
|
27
|
+
DEFAULT_MAX_RISK_SCORE,
|
|
28
|
+
DEFAULT_MIN_CHOICE_CONFIDENCE,
|
|
29
|
+
DEFAULT_MIN_CHOICE_PROBABILITY,
|
|
30
|
+
DEFAULT_NOUL_THRESHOLD,
|
|
31
|
+
GateSignals,
|
|
32
|
+
ToolProposal,
|
|
33
|
+
default_gate_planner,
|
|
34
|
+
describe_plan,
|
|
35
|
+
evaluate_gate,
|
|
36
|
+
)
|
|
37
|
+
from .host import (
|
|
38
|
+
HostOutcome,
|
|
39
|
+
ToolHost,
|
|
40
|
+
describe_outcome,
|
|
41
|
+
dispatch_plan,
|
|
42
|
+
run_gated_call,
|
|
43
|
+
)
|
|
44
|
+
from .outcomes import (
|
|
45
|
+
GateOutcomeRecord,
|
|
46
|
+
append_outcome,
|
|
47
|
+
load_outcomes,
|
|
48
|
+
record_from_host_outcome,
|
|
49
|
+
suggest_threshold_updates,
|
|
50
|
+
summarize_outcomes,
|
|
51
|
+
)
|
|
52
|
+
from .tool_calling import ToolCallPlan, ToolCallPlanner, ToolSpec, validate_tool_arguments
|
|
53
|
+
from .audit import (
|
|
54
|
+
GateAuditEnvelope,
|
|
55
|
+
build_audit_envelope,
|
|
56
|
+
content_hash,
|
|
57
|
+
envelope_from_dict,
|
|
58
|
+
envelope_to_dict,
|
|
59
|
+
verify_envelope,
|
|
60
|
+
)
|
|
61
|
+
from .wire import (
|
|
62
|
+
CONTRACT_VERSION,
|
|
63
|
+
outcome_from_dict,
|
|
64
|
+
outcome_to_dict,
|
|
65
|
+
plan_from_dict,
|
|
66
|
+
plan_to_dict,
|
|
67
|
+
request_from_dict,
|
|
68
|
+
request_to_dict,
|
|
69
|
+
schema_path,
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
__version__ = "0.2.0"
|
|
73
|
+
|
|
74
|
+
# Symbols that require the optional torch stack. Resolved via __getattr__.
|
|
75
|
+
_TORCH_EXPORTS = {
|
|
76
|
+
"CHESS_OPTION_IDS",
|
|
77
|
+
"DOOM_OPTION_IDS",
|
|
78
|
+
"TOTAL_OPTIONS",
|
|
79
|
+
"DecisionModel",
|
|
80
|
+
"DoomScorerV2",
|
|
81
|
+
"MultiQuestionCollator",
|
|
82
|
+
"MultiQuestionDataset",
|
|
83
|
+
"MultiQuestionTinyScorer",
|
|
84
|
+
"convert_typed_row",
|
|
85
|
+
"grpo_loss",
|
|
86
|
+
"load_decision_checkpoint",
|
|
87
|
+
"plan_scored_proposal",
|
|
88
|
+
"run_scored_gated_call",
|
|
89
|
+
"score_proposal",
|
|
90
|
+
}
|
|
91
|
+
|
|
92
|
+
__all__ = [
|
|
93
|
+
"CHESS_OPTION_IDS", "DOOM_OPTION_IDS", "TOTAL_OPTIONS", "DoomScorerV2",
|
|
94
|
+
"CONTRACT_VERSION", "GateAuditEnvelope",
|
|
95
|
+
"ChoiceQuestion", "ChoiceResult", "DecisionRequest", "NoulQuestion",
|
|
96
|
+
"NoulResult", "OptionSpec", "ScoreLevel", "ScoreQuestion", "ScoreResult",
|
|
97
|
+
"DEFAULT_MAX_RISK_SCORE", "DEFAULT_MIN_CHOICE_CONFIDENCE",
|
|
98
|
+
"DEFAULT_MIN_CHOICE_PROBABILITY", "DEFAULT_NOUL_THRESHOLD",
|
|
99
|
+
"AuthorityProfile", "CatalogPolicy", "DecisionModel", "MultiQuestionCollator",
|
|
100
|
+
"MultiQuestionDataset", "MultiQuestionTinyScorer", "GateOutcomeRecord",
|
|
101
|
+
"GateSignals", "HostOutcome", "ToolHost", "ToolProposal", "ToolCallPlan",
|
|
102
|
+
"ToolCallPlanner", "ToolSpec", "append_outcome", "build_audit_envelope",
|
|
103
|
+
"check_catalog_policy", "content_hash", "convert_typed_row",
|
|
104
|
+
"default_gate_planner", "describe_outcome", "describe_plan",
|
|
105
|
+
"dispatch_plan", "envelope_from_dict", "envelope_to_dict", "evaluate_gate",
|
|
106
|
+
"grpo_loss",
|
|
107
|
+
"load_decision_checkpoint", "load_outcomes", "outcome_from_dict",
|
|
108
|
+
"outcome_to_dict", "plan_from_dict", "plan_scored_proposal", "plan_to_dict",
|
|
109
|
+
"record_from_host_outcome", "request_from_dict", "request_to_dict",
|
|
110
|
+
"run_gated_call", "run_scored_gated_call", "schema_path", "score_proposal",
|
|
111
|
+
"suggest_threshold_updates", "summarize_outcomes", "validate_tool_arguments",
|
|
112
|
+
"verify_envelope",
|
|
113
|
+
]
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
def __getattr__(name: str) -> Any:
|
|
117
|
+
"""Lazy-load torch-backed exports so Path A stays installable without torch."""
|
|
118
|
+
if name not in _TORCH_EXPORTS:
|
|
119
|
+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
|
120
|
+
try:
|
|
121
|
+
if name in {"DecisionModel", "load_decision_checkpoint"}:
|
|
122
|
+
from .decision import DecisionModel, load_decision_checkpoint
|
|
123
|
+
|
|
124
|
+
values = {
|
|
125
|
+
"DecisionModel": DecisionModel,
|
|
126
|
+
"load_decision_checkpoint": load_decision_checkpoint,
|
|
127
|
+
}
|
|
128
|
+
elif name in {
|
|
129
|
+
"MultiQuestionCollator",
|
|
130
|
+
"MultiQuestionDataset",
|
|
131
|
+
"MultiQuestionTinyScorer",
|
|
132
|
+
}:
|
|
133
|
+
from .multitask import (
|
|
134
|
+
MultiQuestionCollator,
|
|
135
|
+
MultiQuestionDataset,
|
|
136
|
+
MultiQuestionTinyScorer,
|
|
137
|
+
)
|
|
138
|
+
|
|
139
|
+
values = {
|
|
140
|
+
"MultiQuestionCollator": MultiQuestionCollator,
|
|
141
|
+
"MultiQuestionDataset": MultiQuestionDataset,
|
|
142
|
+
"MultiQuestionTinyScorer": MultiQuestionTinyScorer,
|
|
143
|
+
}
|
|
144
|
+
elif name == "grpo_loss":
|
|
145
|
+
from .rlcd import grpo_loss
|
|
146
|
+
|
|
147
|
+
values = {"grpo_loss": grpo_loss}
|
|
148
|
+
elif name in {
|
|
149
|
+
"plan_scored_proposal",
|
|
150
|
+
"run_scored_gated_call",
|
|
151
|
+
"score_proposal",
|
|
152
|
+
}:
|
|
153
|
+
from .gate_multitask import (
|
|
154
|
+
plan_scored_proposal,
|
|
155
|
+
run_scored_gated_call,
|
|
156
|
+
score_proposal,
|
|
157
|
+
)
|
|
158
|
+
|
|
159
|
+
values = {
|
|
160
|
+
"plan_scored_proposal": plan_scored_proposal,
|
|
161
|
+
"run_scored_gated_call": run_scored_gated_call,
|
|
162
|
+
"score_proposal": score_proposal,
|
|
163
|
+
}
|
|
164
|
+
elif name == "convert_typed_row":
|
|
165
|
+
from .typed_decisions import convert_typed_row
|
|
166
|
+
|
|
167
|
+
values = {"convert_typed_row": convert_typed_row}
|
|
168
|
+
else:
|
|
169
|
+
from .vision import CHESS_OPTION_IDS, DOOM_OPTION_IDS, TOTAL_OPTIONS, DoomScorerV2
|
|
170
|
+
|
|
171
|
+
values = {
|
|
172
|
+
"CHESS_OPTION_IDS": CHESS_OPTION_IDS,
|
|
173
|
+
"DOOM_OPTION_IDS": DOOM_OPTION_IDS,
|
|
174
|
+
"TOTAL_OPTIONS": TOTAL_OPTIONS,
|
|
175
|
+
"DoomScorerV2": DoomScorerV2,
|
|
176
|
+
}
|
|
177
|
+
except ImportError as exc: # pragma: no cover - exercised in gate-only CI
|
|
178
|
+
raise ImportError(
|
|
179
|
+
f"{name} requires the torch extra. "
|
|
180
|
+
'Install with: pip install "akasha-model[torch]"'
|
|
181
|
+
) from exc
|
|
182
|
+
globals().update(values)
|
|
183
|
+
return values[name]
|
akasha_model/audit.py
ADDED
|
@@ -0,0 +1,185 @@
|
|
|
1
|
+
"""Tamper-evident (hash-only) audit envelopes for Path A plans / outcomes.
|
|
2
|
+
|
|
3
|
+
v1 is content-hash integrity without cryptography. Hosts may attach signatures
|
|
4
|
+
later. Construction has no side effects and does not execute tools.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import hashlib
|
|
10
|
+
import json
|
|
11
|
+
from dataclasses import dataclass
|
|
12
|
+
from datetime import datetime, timezone
|
|
13
|
+
from typing import Any, Mapping
|
|
14
|
+
|
|
15
|
+
from .gate import GateSignals, ToolProposal
|
|
16
|
+
from .host import HostOutcome
|
|
17
|
+
from .tool_calling import ToolCallPlan, ToolSpec
|
|
18
|
+
from .wire import (
|
|
19
|
+
CONTRACT_VERSION,
|
|
20
|
+
outcome_to_dict,
|
|
21
|
+
plan_to_dict,
|
|
22
|
+
proposal_to_dict,
|
|
23
|
+
signals_to_dict,
|
|
24
|
+
tool_spec_to_dict,
|
|
25
|
+
)
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def _canonical_json(payload: Mapping[str, Any]) -> bytes:
|
|
29
|
+
return json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=True).encode(
|
|
30
|
+
"utf-8",
|
|
31
|
+
)
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def content_hash(payload: Mapping[str, Any]) -> str:
|
|
35
|
+
"""Stable SHA-256 hex digest over canonical JSON (Python/Rust parity target)."""
|
|
36
|
+
return hashlib.sha256(_canonical_json(payload)).hexdigest()
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
@dataclass(frozen=True)
|
|
40
|
+
class GateAuditEnvelope:
|
|
41
|
+
"""Package-owned audit record for one evaluate / dispatch decision."""
|
|
42
|
+
|
|
43
|
+
contract_version: int
|
|
44
|
+
request_id: str | None
|
|
45
|
+
recorded_at: str
|
|
46
|
+
proposal: Mapping[str, Any]
|
|
47
|
+
signals: Mapping[str, Any] | None
|
|
48
|
+
plan: Mapping[str, Any]
|
|
49
|
+
host_action: str | None
|
|
50
|
+
host_reason: str | None
|
|
51
|
+
tools: Mapping[str, Any] | None
|
|
52
|
+
content_sha256: str
|
|
53
|
+
|
|
54
|
+
def to_dict(self) -> dict[str, Any]:
|
|
55
|
+
return envelope_to_dict(self)
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def envelope_to_dict(envelope: GateAuditEnvelope) -> dict[str, Any]:
|
|
59
|
+
return {
|
|
60
|
+
"contract_version": envelope.contract_version,
|
|
61
|
+
"request_id": envelope.request_id,
|
|
62
|
+
"recorded_at": envelope.recorded_at,
|
|
63
|
+
"proposal": dict(envelope.proposal),
|
|
64
|
+
"signals": None if envelope.signals is None else dict(envelope.signals),
|
|
65
|
+
"plan": dict(envelope.plan),
|
|
66
|
+
"host_action": envelope.host_action,
|
|
67
|
+
"host_reason": envelope.host_reason,
|
|
68
|
+
"tools": None if envelope.tools is None else dict(envelope.tools),
|
|
69
|
+
"content_sha256": envelope.content_sha256,
|
|
70
|
+
}
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def envelope_from_dict(data: Mapping[str, Any]) -> GateAuditEnvelope:
|
|
74
|
+
version = int(data.get("contract_version", CONTRACT_VERSION))
|
|
75
|
+
if version != CONTRACT_VERSION:
|
|
76
|
+
raise ValueError(
|
|
77
|
+
f"unsupported envelope contract_version {version}; "
|
|
78
|
+
f"this package speaks {CONTRACT_VERSION}"
|
|
79
|
+
)
|
|
80
|
+
return GateAuditEnvelope(
|
|
81
|
+
contract_version=version,
|
|
82
|
+
request_id=data.get("request_id"),
|
|
83
|
+
recorded_at=str(data["recorded_at"]),
|
|
84
|
+
proposal=dict(data["proposal"]),
|
|
85
|
+
signals=None if data.get("signals") is None else dict(data["signals"]),
|
|
86
|
+
plan=dict(data["plan"]),
|
|
87
|
+
host_action=data.get("host_action"),
|
|
88
|
+
host_reason=data.get("host_reason"),
|
|
89
|
+
tools=None if data.get("tools") is None else dict(data["tools"]),
|
|
90
|
+
content_sha256=str(data["content_sha256"]),
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def _body_for_hash(
|
|
95
|
+
*,
|
|
96
|
+
contract_version: int,
|
|
97
|
+
request_id: str | None,
|
|
98
|
+
recorded_at: str,
|
|
99
|
+
proposal: Mapping[str, Any],
|
|
100
|
+
signals: Mapping[str, Any] | None,
|
|
101
|
+
plan: Mapping[str, Any],
|
|
102
|
+
host_action: str | None,
|
|
103
|
+
host_reason: str | None,
|
|
104
|
+
tools: Mapping[str, Any] | None,
|
|
105
|
+
) -> dict[str, Any]:
|
|
106
|
+
# Hash excludes content_sha256 itself.
|
|
107
|
+
return {
|
|
108
|
+
"contract_version": contract_version,
|
|
109
|
+
"request_id": request_id,
|
|
110
|
+
"recorded_at": recorded_at,
|
|
111
|
+
"proposal": dict(proposal),
|
|
112
|
+
"signals": None if signals is None else dict(signals),
|
|
113
|
+
"plan": dict(plan),
|
|
114
|
+
"host_action": host_action,
|
|
115
|
+
"host_reason": host_reason,
|
|
116
|
+
"tools": None if tools is None else dict(tools),
|
|
117
|
+
}
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
def build_audit_envelope(
|
|
121
|
+
*,
|
|
122
|
+
plan: ToolCallPlan,
|
|
123
|
+
proposal: ToolProposal,
|
|
124
|
+
signals: GateSignals | None = None,
|
|
125
|
+
tools: Mapping[str, ToolSpec] | None = None,
|
|
126
|
+
outcome: HostOutcome | None = None,
|
|
127
|
+
request_id: str | None = None,
|
|
128
|
+
recorded_at: str | None = None,
|
|
129
|
+
) -> GateAuditEnvelope:
|
|
130
|
+
"""Build an envelope without I/O. Idempotency / persistence stay host-owned."""
|
|
131
|
+
stamp = recorded_at or datetime.now(timezone.utc).isoformat()
|
|
132
|
+
plan_wire = plan_to_dict(plan, request_id=request_id)
|
|
133
|
+
proposal_wire = proposal_to_dict(proposal)
|
|
134
|
+
signals_wire = None if signals is None else signals_to_dict(signals)
|
|
135
|
+
tools_wire = (
|
|
136
|
+
None
|
|
137
|
+
if tools is None
|
|
138
|
+
else {name: tool_spec_to_dict(spec) for name, spec in tools.items()}
|
|
139
|
+
)
|
|
140
|
+
host_action = outcome.action if outcome is not None else None
|
|
141
|
+
host_reason = outcome.host_reason if outcome is not None else None
|
|
142
|
+
if outcome is not None:
|
|
143
|
+
# Prefer outcome plan wire (same plan, ensures consistency).
|
|
144
|
+
plan_wire = outcome_to_dict(outcome, request_id=request_id)["plan"]
|
|
145
|
+
body = _body_for_hash(
|
|
146
|
+
contract_version=CONTRACT_VERSION,
|
|
147
|
+
request_id=request_id,
|
|
148
|
+
recorded_at=stamp,
|
|
149
|
+
proposal=proposal_wire,
|
|
150
|
+
signals=signals_wire,
|
|
151
|
+
plan=plan_wire,
|
|
152
|
+
host_action=host_action,
|
|
153
|
+
host_reason=host_reason,
|
|
154
|
+
tools=tools_wire,
|
|
155
|
+
)
|
|
156
|
+
digest = content_hash(body)
|
|
157
|
+
return GateAuditEnvelope(
|
|
158
|
+
contract_version=CONTRACT_VERSION,
|
|
159
|
+
request_id=request_id,
|
|
160
|
+
recorded_at=stamp,
|
|
161
|
+
proposal=proposal_wire,
|
|
162
|
+
signals=signals_wire,
|
|
163
|
+
plan=plan_wire,
|
|
164
|
+
host_action=host_action,
|
|
165
|
+
host_reason=host_reason,
|
|
166
|
+
tools=tools_wire,
|
|
167
|
+
content_sha256=digest,
|
|
168
|
+
)
|
|
169
|
+
|
|
170
|
+
|
|
171
|
+
def verify_envelope(envelope: GateAuditEnvelope | Mapping[str, Any]) -> bool:
|
|
172
|
+
data = envelope_to_dict(envelope) if isinstance(envelope, GateAuditEnvelope) else dict(envelope)
|
|
173
|
+
expected = data.get("content_sha256")
|
|
174
|
+
body = {k: v for k, v in data.items() if k != "content_sha256"}
|
|
175
|
+
return expected == content_hash(body)
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
__all__ = [
|
|
179
|
+
"GateAuditEnvelope",
|
|
180
|
+
"build_audit_envelope",
|
|
181
|
+
"content_hash",
|
|
182
|
+
"envelope_from_dict",
|
|
183
|
+
"envelope_to_dict",
|
|
184
|
+
"verify_envelope",
|
|
185
|
+
]
|
|
@@ -0,0 +1,163 @@
|
|
|
1
|
+
"""Optional Path A authority profiles (MIDAS-inspired, executor-free).
|
|
2
|
+
|
|
3
|
+
Profiles tune *how far* authorization goes once a tool is proposed. Catalog
|
|
4
|
+
policy (#15) decides *what* may be proposed; OS ``check_permissions`` remains
|
|
5
|
+
the final grant check.
|
|
6
|
+
|
|
7
|
+
Statuses stay ``ready`` / ``abstain`` / ``blocked`` for backward compatibility.
|
|
8
|
+
When a profile is supplied, ``plan.reason`` is prefixed with a stable code:
|
|
9
|
+
|
|
10
|
+
- ``escalate:`` — within authority chain but over profile threshold (human review)
|
|
11
|
+
- ``reject:`` — hard policy / risk deny
|
|
12
|
+
- ``clarify:`` — insufficient context / missing required context keys
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
from __future__ import annotations
|
|
16
|
+
|
|
17
|
+
from dataclasses import dataclass
|
|
18
|
+
from typing import Mapping
|
|
19
|
+
|
|
20
|
+
from .tool_calling import ToolCallPlan
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
@dataclass(frozen=True)
|
|
24
|
+
class AuthorityProfile:
|
|
25
|
+
"""Optional thresholds bound to a decision surface / tool class.
|
|
26
|
+
|
|
27
|
+
All fields are optional; ``None`` means “use the planner defaults / ignore”.
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
name: str = "default"
|
|
31
|
+
min_authorized: float | None = None
|
|
32
|
+
min_sufficient_context: float | None = None
|
|
33
|
+
max_risk_score: float | None = None
|
|
34
|
+
max_consequence: float | None = None
|
|
35
|
+
required_context_keys: tuple[str, ...] = ()
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def annotate_reason(code: str, detail: str) -> str:
|
|
39
|
+
return f"{code}: {detail}"
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def apply_authority_profile(
|
|
43
|
+
plan: ToolCallPlan,
|
|
44
|
+
*,
|
|
45
|
+
profile: AuthorityProfile | None,
|
|
46
|
+
authorized: float,
|
|
47
|
+
sufficient_context: float,
|
|
48
|
+
risk_score: float | None,
|
|
49
|
+
consequence: float | None,
|
|
50
|
+
context: Mapping[str, object] | None,
|
|
51
|
+
) -> ToolCallPlan:
|
|
52
|
+
"""Re-map or annotate a plan when an authority profile is active.
|
|
53
|
+
|
|
54
|
+
Callers that omit ``profile`` keep today's behaviour unchanged.
|
|
55
|
+
"""
|
|
56
|
+
if profile is None:
|
|
57
|
+
return plan
|
|
58
|
+
|
|
59
|
+
if profile.required_context_keys:
|
|
60
|
+
ctx = context or {}
|
|
61
|
+
missing = [key for key in profile.required_context_keys if key not in ctx]
|
|
62
|
+
if missing:
|
|
63
|
+
return ToolCallPlan(
|
|
64
|
+
"abstain",
|
|
65
|
+
plan.tool_name,
|
|
66
|
+
None,
|
|
67
|
+
annotate_reason(
|
|
68
|
+
"clarify",
|
|
69
|
+
"missing required context key(s): " + ", ".join(missing),
|
|
70
|
+
),
|
|
71
|
+
plan.choice_probability,
|
|
72
|
+
plan.choice_confidence,
|
|
73
|
+
plan.risk_score,
|
|
74
|
+
)
|
|
75
|
+
|
|
76
|
+
if (
|
|
77
|
+
profile.min_sufficient_context is not None
|
|
78
|
+
and sufficient_context < profile.min_sufficient_context
|
|
79
|
+
):
|
|
80
|
+
return ToolCallPlan(
|
|
81
|
+
"abstain",
|
|
82
|
+
plan.tool_name,
|
|
83
|
+
None,
|
|
84
|
+
annotate_reason(
|
|
85
|
+
"clarify",
|
|
86
|
+
"sufficient_context below authority profile minimum",
|
|
87
|
+
),
|
|
88
|
+
plan.choice_probability,
|
|
89
|
+
plan.choice_confidence,
|
|
90
|
+
plan.risk_score,
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
if profile.min_authorized is not None and authorized < profile.min_authorized:
|
|
94
|
+
return ToolCallPlan(
|
|
95
|
+
"blocked",
|
|
96
|
+
plan.tool_name,
|
|
97
|
+
None,
|
|
98
|
+
annotate_reason("reject", "authorized below authority profile minimum"),
|
|
99
|
+
plan.choice_probability,
|
|
100
|
+
plan.choice_confidence,
|
|
101
|
+
plan.risk_score,
|
|
102
|
+
)
|
|
103
|
+
|
|
104
|
+
if (
|
|
105
|
+
profile.max_risk_score is not None
|
|
106
|
+
and risk_score is not None
|
|
107
|
+
and risk_score > profile.max_risk_score
|
|
108
|
+
):
|
|
109
|
+
return ToolCallPlan(
|
|
110
|
+
"blocked",
|
|
111
|
+
plan.tool_name,
|
|
112
|
+
None,
|
|
113
|
+
annotate_reason("reject", "risk score exceeds authority profile maximum"),
|
|
114
|
+
plan.choice_probability,
|
|
115
|
+
plan.choice_confidence,
|
|
116
|
+
risk_score,
|
|
117
|
+
)
|
|
118
|
+
|
|
119
|
+
if (
|
|
120
|
+
profile.max_consequence is not None
|
|
121
|
+
and consequence is not None
|
|
122
|
+
and consequence > profile.max_consequence
|
|
123
|
+
):
|
|
124
|
+
return ToolCallPlan(
|
|
125
|
+
"abstain",
|
|
126
|
+
plan.tool_name,
|
|
127
|
+
None,
|
|
128
|
+
annotate_reason(
|
|
129
|
+
"escalate",
|
|
130
|
+
"consequence exceeds authority profile; human review required",
|
|
131
|
+
),
|
|
132
|
+
plan.choice_probability,
|
|
133
|
+
plan.choice_confidence,
|
|
134
|
+
plan.risk_score,
|
|
135
|
+
)
|
|
136
|
+
|
|
137
|
+
# Already non-ready from planner: prefix stable codes where obvious.
|
|
138
|
+
if plan.status == "blocked" and plan.reason.startswith(
|
|
139
|
+
("authorization", "Risk score", "required capability", "explicit human"),
|
|
140
|
+
):
|
|
141
|
+
return ToolCallPlan(
|
|
142
|
+
plan.status,
|
|
143
|
+
plan.tool_name,
|
|
144
|
+
plan.arguments,
|
|
145
|
+
annotate_reason("reject", plan.reason),
|
|
146
|
+
plan.choice_probability,
|
|
147
|
+
plan.choice_confidence,
|
|
148
|
+
plan.risk_score,
|
|
149
|
+
)
|
|
150
|
+
if plan.status == "blocked" and "context sufficiency" in plan.reason:
|
|
151
|
+
return ToolCallPlan(
|
|
152
|
+
"abstain",
|
|
153
|
+
plan.tool_name,
|
|
154
|
+
None,
|
|
155
|
+
annotate_reason("clarify", plan.reason),
|
|
156
|
+
plan.choice_probability,
|
|
157
|
+
plan.choice_confidence,
|
|
158
|
+
plan.risk_score,
|
|
159
|
+
)
|
|
160
|
+
return plan
|
|
161
|
+
|
|
162
|
+
|
|
163
|
+
__all__ = ["AuthorityProfile", "annotate_reason", "apply_authority_profile"]
|
|
@@ -0,0 +1,87 @@
|
|
|
1
|
+
"""Post-hoc temperature calibration for one-pass choice models."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import argparse
|
|
6
|
+
import json
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
|
|
9
|
+
import torch
|
|
10
|
+
from torch.nn import functional as F
|
|
11
|
+
from torch.utils.data import DataLoader
|
|
12
|
+
|
|
13
|
+
from .data import JsonlDataset
|
|
14
|
+
from .model import load_checkpoint, select_device
|
|
15
|
+
from .train import move
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def _pad_logits(items: list[torch.Tensor]) -> torch.Tensor:
|
|
19
|
+
width = max(item.shape[1] for item in items)
|
|
20
|
+
fill = -1e4
|
|
21
|
+
return torch.cat([
|
|
22
|
+
F.pad(item.clamp_min(fill), (0, width - item.shape[1]), value=fill)
|
|
23
|
+
for item in items
|
|
24
|
+
])
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@torch.no_grad()
|
|
28
|
+
def collect_logits(model, loader, device):
|
|
29
|
+
model.eval()
|
|
30
|
+
logits, labels = [], []
|
|
31
|
+
for host_batch in loader:
|
|
32
|
+
batch = move(host_batch, device)
|
|
33
|
+
logits.append(model(batch).cpu())
|
|
34
|
+
labels.append(batch["labels"].cpu())
|
|
35
|
+
return _pad_logits(logits), torch.cat(labels)
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def fit_temperature(logits: torch.Tensor, labels: torch.Tensor) -> float:
|
|
39
|
+
"""Fit one positive temperature by minimizing held-out NLL."""
|
|
40
|
+
log_temperature = torch.zeros((), requires_grad=True)
|
|
41
|
+
optimiser = torch.optim.Adam([log_temperature], lr=0.05)
|
|
42
|
+
for _ in range(200):
|
|
43
|
+
optimiser.zero_grad()
|
|
44
|
+
temperature = log_temperature.clamp(-5.0, 5.0).exp()
|
|
45
|
+
loss = F.cross_entropy(logits / temperature, labels)
|
|
46
|
+
loss.backward()
|
|
47
|
+
optimiser.step()
|
|
48
|
+
with torch.no_grad():
|
|
49
|
+
log_temperature.clamp_(-5.0, 5.0)
|
|
50
|
+
return float(log_temperature.detach().clamp(-5.0, 5.0).exp())
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def main() -> None:
|
|
54
|
+
parser = argparse.ArgumentParser()
|
|
55
|
+
parser.add_argument("checkpoint")
|
|
56
|
+
parser.add_argument("calibration_data")
|
|
57
|
+
parser.add_argument("--output", default="runs/calibrated.pt")
|
|
58
|
+
parser.add_argument("--batch-size", type=int, default=64)
|
|
59
|
+
parser.add_argument("--device", choices=("auto", "cpu", "mps", "cuda"), default="auto")
|
|
60
|
+
args = parser.parse_args()
|
|
61
|
+
|
|
62
|
+
device = select_device(args.device)
|
|
63
|
+
model, collator, config = load_checkpoint(args.checkpoint, device)
|
|
64
|
+
loader = DataLoader(
|
|
65
|
+
JsonlDataset(args.calibration_data), batch_size=args.batch_size,
|
|
66
|
+
collate_fn=collator,
|
|
67
|
+
)
|
|
68
|
+
logits, labels = collect_logits(model, loader, device)
|
|
69
|
+
old_temperature = float(config.get("temperature", 1.0))
|
|
70
|
+
temperature = fit_temperature(logits, labels) * old_temperature
|
|
71
|
+
|
|
72
|
+
payload = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
|
|
73
|
+
payload["config"]["temperature"] = temperature
|
|
74
|
+
output = Path(args.output)
|
|
75
|
+
output.parent.mkdir(parents=True, exist_ok=True)
|
|
76
|
+
torch.save(payload, output)
|
|
77
|
+
print(json.dumps({
|
|
78
|
+
"checkpoint": str(output),
|
|
79
|
+
"calibration_examples": int(labels.numel()),
|
|
80
|
+
"old_temperature": old_temperature,
|
|
81
|
+
"temperature": temperature,
|
|
82
|
+
"device": str(device),
|
|
83
|
+
}))
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
if __name__ == "__main__":
|
|
87
|
+
main()
|
akasha_model/catalog.py
ADDED
|
@@ -0,0 +1,72 @@
|
|
|
1
|
+
"""Optional phase-3 catalog constraints for Path A.
|
|
2
|
+
|
|
3
|
+
Layering (outer → inner):
|
|
4
|
+
|
|
5
|
+
1. **CatalogPolicy** (this module) — which tools may be proposed on a surface
|
|
6
|
+
(allowlist / denylist / required placement tags).
|
|
7
|
+
2. **ToolSpec** — per-tool schema, capabilities, confirmation, irreversible.
|
|
8
|
+
3. **OS ``check_permissions``** — final ACL / capability grant on the host.
|
|
9
|
+
|
|
10
|
+
No execution here. Outcomes suggestions (#8) never auto-write these rules.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
from dataclasses import dataclass, field
|
|
16
|
+
from typing import Mapping
|
|
17
|
+
|
|
18
|
+
from .tool_calling import ToolCallPlan, ToolSpec
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
@dataclass(frozen=True)
|
|
22
|
+
class CatalogPolicy:
|
|
23
|
+
"""Optional constraints applied before planner thresholds.
|
|
24
|
+
|
|
25
|
+
- ``allowlist``: if set, proposal tool must be in the set.
|
|
26
|
+
- ``denylist``: proposal tool must not be in the set.
|
|
27
|
+
- ``required_placement``: every listed tag must appear on the tool's
|
|
28
|
+
``placement_tags`` (e.g. ``offline``, ``preview``).
|
|
29
|
+
"""
|
|
30
|
+
|
|
31
|
+
allowlist: frozenset[str] | None = None
|
|
32
|
+
denylist: frozenset[str] = field(default_factory=frozenset)
|
|
33
|
+
required_placement: frozenset[str] = field(default_factory=frozenset)
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def check_catalog_policy(
|
|
37
|
+
tools: Mapping[str, ToolSpec],
|
|
38
|
+
tool_name: str,
|
|
39
|
+
policy: CatalogPolicy | None,
|
|
40
|
+
) -> ToolCallPlan | None:
|
|
41
|
+
"""Return a blocked plan when policy rejects the proposal, else ``None``."""
|
|
42
|
+
if policy is None:
|
|
43
|
+
return None
|
|
44
|
+
if policy.allowlist is not None and tool_name not in policy.allowlist:
|
|
45
|
+
return ToolCallPlan(
|
|
46
|
+
"blocked",
|
|
47
|
+
tool_name,
|
|
48
|
+
None,
|
|
49
|
+
f"catalog_allowlist: tool not permitted: {tool_name}",
|
|
50
|
+
)
|
|
51
|
+
if tool_name in policy.denylist:
|
|
52
|
+
return ToolCallPlan(
|
|
53
|
+
"blocked",
|
|
54
|
+
tool_name,
|
|
55
|
+
None,
|
|
56
|
+
f"catalog_denylist: tool denied: {tool_name}",
|
|
57
|
+
)
|
|
58
|
+
if policy.required_placement:
|
|
59
|
+
spec = tools.get(tool_name)
|
|
60
|
+
tags = frozenset(getattr(spec, "placement_tags", ()) or ()) if spec else frozenset()
|
|
61
|
+
missing = sorted(policy.required_placement - tags)
|
|
62
|
+
if missing:
|
|
63
|
+
return ToolCallPlan(
|
|
64
|
+
"blocked",
|
|
65
|
+
tool_name,
|
|
66
|
+
None,
|
|
67
|
+
"catalog_placement: missing tag(s): " + ", ".join(missing),
|
|
68
|
+
)
|
|
69
|
+
return None
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
__all__ = ["CatalogPolicy", "check_catalog_policy"]
|