simthinkd 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.
- simthinkd/__init__.py +15 -0
- simthinkd/bench.py +55 -0
- simthinkd/cli.py +62 -0
- simthinkd/core.py +158 -0
- simthinkd/data/doom_defend_states.json +1 -0
- simthinkd/integrations/__init__.py +0 -0
- simthinkd/integrations/langchain_tool.py +27 -0
- simthinkd/integrations/mcp_server.py +52 -0
- simthinkd/policy.py +181 -0
- simthinkd/server.py +60 -0
- simthinkd/toy.py +53 -0
- simthinkd/train.py +222 -0
- simthinkd/weights/SHA256SUMS +2 -0
- simthinkd/weights/doom-corridor.npz +0 -0
- simthinkd/weights/doom-defend.npz +0 -0
- simthinkd-0.1.0.dist-info/METADATA +170 -0
- simthinkd-0.1.0.dist-info/RECORD +22 -0
- simthinkd-0.1.0.dist-info/WHEEL +5 -0
- simthinkd-0.1.0.dist-info/entry_points.txt +2 -0
- simthinkd-0.1.0.dist-info/licenses/LICENSE +202 -0
- simthinkd-0.1.0.dist-info/licenses/NOTICE +14 -0
- simthinkd-0.1.0.dist-info/top_level.txt +1 -0
|
File without changes
|
|
@@ -0,0 +1,27 @@
|
|
|
1
|
+
"""LangChain / LangGraph tool: route a step to a 2 ms decider instead of a model call.
|
|
2
|
+
|
|
3
|
+
pip install "simthinkd[langchain]"
|
|
4
|
+
from simthinkd.integrations.langchain_tool import simthinkd_tool
|
|
5
|
+
tool = simthinkd_tool("doom-defend")
|
|
6
|
+
tool.invoke({"state": "seen: Demon left a30 d5 | enemies 1 | sway left | gun ready | ammo25"})
|
|
7
|
+
"""
|
|
8
|
+
from ..core import Decider
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def simthinkd_tool(decider='doom-defend', name=None, description=None):
|
|
12
|
+
"""A LangChain StructuredTool around one decider. Input: {"state": str}. Output: {"choice", "confidence", "ms"}."""
|
|
13
|
+
try:
|
|
14
|
+
from langchain_core.tools import StructuredTool
|
|
15
|
+
except ImportError as error:
|
|
16
|
+
raise ImportError('the LangChain tool needs: pip install "simthinkd[langchain]"') from error
|
|
17
|
+
d = decider if isinstance(decider, Decider) else Decider(decider)
|
|
18
|
+
actions = ', '.join(d.actions)
|
|
19
|
+
|
|
20
|
+
def run(state: str) -> dict:
|
|
21
|
+
r = d.decide(state)
|
|
22
|
+
return {'choice': r.choice, 'confidence': round(r.confidence, 4), 'ms': round(r.ms, 3)}
|
|
23
|
+
|
|
24
|
+
return StructuredTool.from_function(
|
|
25
|
+
run, name=name or f'simthinkd_{d.preset["name"].replace("-", "_")}',
|
|
26
|
+
description=description or (f'Pick one action ({actions}) for a short situation sentence in about 2 ms. '
|
|
27
|
+
f'Goal: {d.preset.get("goal", "")}'))
|
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
"""MCP server: lets an agent (Claude Desktop, Cursor, any MCP client) ask a SimThink D decider for a fast decision.
|
|
2
|
+
|
|
3
|
+
pip install "simthinkd[mcp]"
|
|
4
|
+
python -m simthinkd.integrations.mcp_server # stdio transport
|
|
5
|
+
|
|
6
|
+
Client config example:
|
|
7
|
+
{"mcpServers": {"simthinkd": {"command": "python", "args": ["-m", "simthinkd.integrations.mcp_server"]}}}
|
|
8
|
+
"""
|
|
9
|
+
from functools import lru_cache
|
|
10
|
+
|
|
11
|
+
from ..core import Decider, available
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
@lru_cache(maxsize=8)
|
|
15
|
+
def _decider(name):
|
|
16
|
+
return Decider(name)
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def list_deciders() -> dict:
|
|
20
|
+
"""Bundled deciders and what each was trained for."""
|
|
21
|
+
return available()
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def decide(state: str, decider: str = 'doom-defend', actions: dict | None = None, goal: str | None = None) -> dict:
|
|
25
|
+
"""Pick one action for a short situation sentence. Returns choice, confidence, probabilities and time in ms.
|
|
26
|
+
|
|
27
|
+
`actions` ({name: description}) and `goal` default to the ones the decider was trained with.
|
|
28
|
+
"""
|
|
29
|
+
r = _decider(decider).decide(state, actions=actions, goal=goal)
|
|
30
|
+
return {'choice': r.choice, 'confidence': r.confidence, 'probabilities': r.probabilities, 'ms': round(r.ms, 3)}
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def build_server():
|
|
34
|
+
try:
|
|
35
|
+
from mcp.server.mcpserver import MCPServer as Server # mcp 2.x
|
|
36
|
+
except ImportError:
|
|
37
|
+
try:
|
|
38
|
+
from mcp.server.fastmcp import FastMCP as Server # mcp 1.x
|
|
39
|
+
except ImportError as error:
|
|
40
|
+
raise ImportError('the MCP server needs: pip install "simthinkd[mcp]"') from error
|
|
41
|
+
server = Server('simthinkd')
|
|
42
|
+
server.tool()(list_deciders)
|
|
43
|
+
server.tool()(decide)
|
|
44
|
+
return server
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def main():
|
|
48
|
+
build_server().run()
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
if __name__ == '__main__':
|
|
52
|
+
main()
|
simthinkd/policy.py
ADDED
|
@@ -0,0 +1,181 @@
|
|
|
1
|
+
"""Goal/state/candidate encoder and portable learned policy; no task parser/oracle.
|
|
2
|
+
|
|
3
|
+
All text is consumed by deterministic hashed word/character n-gram features.
|
|
4
|
+
Neural projections, cross-candidate context and scores are learned from scratch.
|
|
5
|
+
Hash compression is lossy; this is not a general pretrained language encoder.
|
|
6
|
+
"""
|
|
7
|
+
import hashlib
|
|
8
|
+
import json
|
|
9
|
+
import re
|
|
10
|
+
import unicodedata
|
|
11
|
+
from functools import lru_cache
|
|
12
|
+
from pathlib import Path
|
|
13
|
+
|
|
14
|
+
import numpy as np
|
|
15
|
+
|
|
16
|
+
D = 192
|
|
17
|
+
OPS = ['CLICK', 'TYPE_TEXT', 'SELECT', 'WAIT', 'SCROLL_DOWN', 'SCROLL_UP', 'DONE', 'BLOCKED']
|
|
18
|
+
ROLES = ['link', 'button', 'textbox', 'searchbox', 'combobox', 'checkbox', 'option', 'radio']
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def tokens(text):
|
|
22
|
+
return re.findall(r'\w+|[^\w\s]', unicodedata.normalize('NFKC', str(text)).lower())
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
@lru_cache(maxsize=60000)
|
|
26
|
+
def hashed(text):
|
|
27
|
+
words = tokens(text)
|
|
28
|
+
features = ['w:' + t for t in words]
|
|
29
|
+
features += ['b:' + a + ' ' + b for a, b in zip(words, words[1:])]
|
|
30
|
+
features += ['c:' + t[i:i+3] for t in words for i in range(max(0, len(t)-2))]
|
|
31
|
+
result = np.zeros(D, np.float32)
|
|
32
|
+
for value in features:
|
|
33
|
+
number = int.from_bytes(hashlib.blake2s(value.encode(), digest_size=4).digest(), 'little')
|
|
34
|
+
result[number % D] += 1 if number & 256 else -1
|
|
35
|
+
result /= max(float(np.linalg.norm(result)), 1.)
|
|
36
|
+
return result
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def text(value):
|
|
40
|
+
return value if isinstance(value, str) else json.dumps(value, ensure_ascii=False, sort_keys=True)
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def without_indices(value):
|
|
44
|
+
if isinstance(value, dict):
|
|
45
|
+
return {k: without_indices(v) for k, v in value.items() if k != 'index'}
|
|
46
|
+
if isinstance(value, list):
|
|
47
|
+
return sorted([without_indices(v) for v in value], key=text)
|
|
48
|
+
return value
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def rows_for(body):
|
|
52
|
+
questions = body.get('questions', {})
|
|
53
|
+
if 'operation' not in questions:
|
|
54
|
+
raise ValueError('This decider answers operation/target questions only')
|
|
55
|
+
rows, operations = [], list(questions['operation']['criteria'])
|
|
56
|
+
for group, op in enumerate(operations):
|
|
57
|
+
head = op.lower() + '_target'
|
|
58
|
+
if head in questions:
|
|
59
|
+
for key, value in questions[head]['criteria'].items():
|
|
60
|
+
rows.append({'operation': op, 'target': key, 'head': head, 'group': group, 'criterion': value})
|
|
61
|
+
else:
|
|
62
|
+
rows.append({'operation': op, 'target': None, 'head': None, 'group': group,
|
|
63
|
+
'criterion': {'element': op, 'description': questions['operation']['criteria'][op]}})
|
|
64
|
+
if not rows or len(rows) > 1500:
|
|
65
|
+
raise ValueError('Invalid offered action space')
|
|
66
|
+
return rows, operations
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def overlap(a, b):
|
|
70
|
+
aa, bb = set(tokens(a)), set(tokens(b))
|
|
71
|
+
return [len(aa & bb)/max(len(aa), 1), len(aa & bb)/max(len(bb), 1),
|
|
72
|
+
float(bool(str(b)) and str(b).casefold() in str(a).casefold())]
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def encode(body):
|
|
76
|
+
rows, operations = rows_for(body)
|
|
77
|
+
instructions = body['questions']['operation'].get('instructions', {})
|
|
78
|
+
goal = text(instructions.get('goal', ''))
|
|
79
|
+
state = body['state']
|
|
80
|
+
page = text(state.get('page', {}))
|
|
81
|
+
history = text(state.get('recent_actions', []))
|
|
82
|
+
page_state = state.get('page', {}) if isinstance(state.get('page', {}), dict) else {}
|
|
83
|
+
page_text = text(page_state.get('title', '')) + '\n' + text(page_state.get('text', ''))
|
|
84
|
+
recent = state.get('recent_actions') if isinstance(state.get('recent_actions'), list) else []
|
|
85
|
+
recent = [a for a in recent if isinstance(a, dict)]
|
|
86
|
+
# Order is preserved by the caller; the last entry is the most recent action.
|
|
87
|
+
recent_labels = [str(a.get('action', '')).strip().casefold() for a in recent]
|
|
88
|
+
last_changed = float(bool(recent[-1].get('page_changed'))) if recent else 0.
|
|
89
|
+
# Preserve all observed fields, removing only arbitrary indexing identities.
|
|
90
|
+
elements = without_indices(state.get('elements', []))
|
|
91
|
+
context = text(elements) + '\n' + text(instructions.get('rules', ''))
|
|
92
|
+
shared = [hashed(goal), hashed(page), hashed(history), hashed(context)]
|
|
93
|
+
goal_words = tokens(goal)
|
|
94
|
+
result = []
|
|
95
|
+
for row in rows:
|
|
96
|
+
c = row['criterion']
|
|
97
|
+
c = c if isinstance(c, dict) else {'element': text(c)}
|
|
98
|
+
label = re.sub(r'^\[[^\]]+\]\s*', '', c.get('element', ''))
|
|
99
|
+
value = text(c.get('current_value', ''))
|
|
100
|
+
local = text({**c, 'element': label, 'operation': row['operation']})
|
|
101
|
+
# Generic alignment windows retain order/negation context without parsing goals.
|
|
102
|
+
matching = set(tokens(label + ' ' + value))
|
|
103
|
+
positions = [i for i, word in enumerate(goal_words) if word in matching]
|
|
104
|
+
focus = ' | '.join(' '.join(goal_words[max(0, i-3):i+4]) for i in positions)
|
|
105
|
+
numeric = [float(row['operation'] == op) for op in OPS]
|
|
106
|
+
numeric += [float(c.get('role') == role) for role in ROLES]
|
|
107
|
+
numeric += [float(str(c.get(key, '')).lower() == flag)
|
|
108
|
+
for key in ['checked', 'selected', 'expanded'] for flag in ['true', 'false']]
|
|
109
|
+
numeric += [float(bool(value)), min(len(rows), 100)/100., min(len(positions), 20)/20.]
|
|
110
|
+
numeric += overlap(goal, label) + overlap(goal, value) + overlap(page, label) + overlap(history, label)
|
|
111
|
+
# Generic pair interactions preserve comparison evidence before compression.
|
|
112
|
+
local_hash = hashed(local)
|
|
113
|
+
products = [hashed(goal)*local_hash, hashed(goal)*hashed(page), hashed(page)*local_hash]
|
|
114
|
+
def char_overlap(a, b):
|
|
115
|
+
def grams(v):
|
|
116
|
+
v = unicodedata.normalize('NFKC', str(v)).lower()
|
|
117
|
+
return {v[i:i+2] for i in range(max(0, len(v)-1)) if not v[i:i+2].isspace()}
|
|
118
|
+
aa, bb = grams(a), grams(b)
|
|
119
|
+
return [len(aa & bb)/max(len(aa), 1), len(aa & bb)/max(len(bb), 1)]
|
|
120
|
+
numeric += char_overlap(goal, label) + char_overlap(goal, page)
|
|
121
|
+
# Recency of this candidate itself: a bag overlap cannot tell order from count.
|
|
122
|
+
key = label.strip().casefold()
|
|
123
|
+
seen = [i for i, name in enumerate(recent_labels) if name and name == key]
|
|
124
|
+
distance = len(recent_labels) - 1 - seen[-1] if seen else None
|
|
125
|
+
numeric += [0. if distance is None else 1./(1.+distance),
|
|
126
|
+
min(len(seen), 5)/5.,
|
|
127
|
+
float(distance == 0),
|
|
128
|
+
float(distance == 0)*last_changed,
|
|
129
|
+
float(len(seen) >= 2),
|
|
130
|
+
min(len(recent_labels), 10)/10.]
|
|
131
|
+
# DONE carries no target, so page agreement has to reach its own row explicitly.
|
|
132
|
+
is_done = float(row['operation'] == 'DONE')
|
|
133
|
+
page_match = overlap(goal, page_text) + char_overlap(goal, page_text)
|
|
134
|
+
numeric += page_match + [is_done*v for v in page_match]
|
|
135
|
+
result.append(np.concatenate([*shared, hashed(local), hashed(focus), *products, np.array(numeric, np.float32)]))
|
|
136
|
+
return np.stack(result), np.array([r['group'] for r in rows], np.int64), rows, operations
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
def softmax(v):
|
|
140
|
+
ex = np.exp(v - np.max(v))
|
|
141
|
+
return ex / ex.sum()
|
|
142
|
+
|
|
143
|
+
|
|
144
|
+
def distributions(scores, groups, temperature=1.):
|
|
145
|
+
scores = scores / temperature
|
|
146
|
+
op_scores, conditional = [], {}
|
|
147
|
+
for g in sorted(set(groups.tolist())):
|
|
148
|
+
where = np.flatnonzero(groups == g)
|
|
149
|
+
s = scores[where]
|
|
150
|
+
op_scores.append(float(np.max(s) + np.log(np.exp(s-np.max(s)).mean())))
|
|
151
|
+
conditional[g] = softmax(s)
|
|
152
|
+
op = softmax(np.asarray(op_scores))
|
|
153
|
+
joint = np.zeros(len(scores))
|
|
154
|
+
for g, p in conditional.items():
|
|
155
|
+
joint[groups == g] = op[g] * p
|
|
156
|
+
return op, conditional, joint
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
class Policy:
|
|
160
|
+
def __init__(self, path):
|
|
161
|
+
self.path = Path(path)
|
|
162
|
+
self.digest = hashlib.sha256(self.path.read_bytes()).hexdigest()
|
|
163
|
+
with np.load(path, allow_pickle=False) as saved:
|
|
164
|
+
self.w = {key: saved[key].copy() for key in saved.files}
|
|
165
|
+
self.temperature = float(self.w.get('temperature', 1.))
|
|
166
|
+
|
|
167
|
+
def scores(self, x):
|
|
168
|
+
w = self.w
|
|
169
|
+
h = np.maximum(0, x @ w['local.weight'].T + w['local.bias'])
|
|
170
|
+
context = np.concatenate([h.mean(0), h.max(0)])
|
|
171
|
+
z = np.concatenate([h, np.broadcast_to(context, (len(h), len(context)))], axis=1)
|
|
172
|
+
z = np.maximum(0, z @ w['context.weight'].T + w['context.bias'])
|
|
173
|
+
return (z @ w['score.weight'].T + w['score.bias'])[:, 0]
|
|
174
|
+
|
|
175
|
+
def predict(self, body):
|
|
176
|
+
x, groups, rows, operations = encode(body)
|
|
177
|
+
op, conditional, joint = distributions(self.scores(x), groups, self.temperature)
|
|
178
|
+
chosen_group = int(op.argmax())
|
|
179
|
+
indices = np.flatnonzero(groups == chosen_group)
|
|
180
|
+
chosen = rows[int(indices[conditional[chosen_group].argmax()])]
|
|
181
|
+
return chosen, op, conditional, joint, rows, operations
|
simthinkd/server.py
ADDED
|
@@ -0,0 +1,60 @@
|
|
|
1
|
+
"""HTTP server for a decider (POST /v1/systemone, GET /health). Standard library + NumPy only.
|
|
2
|
+
|
|
3
|
+
simthinkd serve doom-defend --port 11890 [--delay-ms 0]
|
|
4
|
+
|
|
5
|
+
Binds to 127.0.0.1 by default. `--delay-ms` adds a fixed wait after inference (latency-injection experiments).
|
|
6
|
+
"""
|
|
7
|
+
import json
|
|
8
|
+
import time
|
|
9
|
+
from datetime import datetime, timezone
|
|
10
|
+
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
11
|
+
|
|
12
|
+
from .core import Decider
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def make_handler(decider, name, delay_ms):
|
|
16
|
+
class Handler(BaseHTTPRequestHandler):
|
|
17
|
+
protocol_version = 'HTTP/1.1'
|
|
18
|
+
|
|
19
|
+
def log_message(self, *unused):
|
|
20
|
+
pass
|
|
21
|
+
|
|
22
|
+
def _send(self, code, payload):
|
|
23
|
+
data = json.dumps(payload, ensure_ascii=False).encode()
|
|
24
|
+
self.send_response(code)
|
|
25
|
+
self.send_header('Content-Type', 'application/json')
|
|
26
|
+
self.send_header('Content-Length', str(len(data)))
|
|
27
|
+
self.end_headers()
|
|
28
|
+
self.wfile.write(data)
|
|
29
|
+
|
|
30
|
+
def do_GET(self):
|
|
31
|
+
if self.path != '/health':
|
|
32
|
+
return self._send(404, {'error': 'not found'})
|
|
33
|
+
self._send(200, {'model': name, 'weights_sha256': decider.sha256, 'delay_ms': delay_ms})
|
|
34
|
+
|
|
35
|
+
def do_POST(self):
|
|
36
|
+
if self.path != '/v1/systemone':
|
|
37
|
+
return self._send(404, {'error': 'not found'})
|
|
38
|
+
try:
|
|
39
|
+
body = json.loads(self.rfile.read(int(self.headers.get('Content-Length', 0))) or b'{}')
|
|
40
|
+
start = time.perf_counter()
|
|
41
|
+
answers = decider.predict(body)
|
|
42
|
+
infer_ms = (time.perf_counter() - start) * 1000
|
|
43
|
+
except (ValueError, KeyError, TypeError) as error:
|
|
44
|
+
return self._send(400, {'error': f'invalid request: {error}'})
|
|
45
|
+
if delay_ms:
|
|
46
|
+
time.sleep(delay_ms / 1000)
|
|
47
|
+
self._send(200, {'model': name, 'answers': answers, 'usage': {'external_model_calls': 0},
|
|
48
|
+
'meta': {'total_ms': round(infer_ms + delay_ms, 3), 'infer_ms': round(infer_ms, 3),
|
|
49
|
+
'delay_ms': delay_ms, 'weights_sha256': decider.sha256,
|
|
50
|
+
'answered_at': datetime.now(timezone.utc).isoformat()}})
|
|
51
|
+
|
|
52
|
+
return Handler
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def serve(decider='doom-defend', host='127.0.0.1', port=11890, delay_ms=0.0, name=None):
|
|
56
|
+
decider = decider if isinstance(decider, Decider) else Decider(decider)
|
|
57
|
+
name = name or f'simthink-d:{decider.preset["name"]}'
|
|
58
|
+
server = ThreadingHTTPServer((host, port), make_handler(decider, name, delay_ms))
|
|
59
|
+
print(json.dumps({'listening': f'{host}:{port}', 'weights_sha256': decider.sha256, 'model': name}), flush=True)
|
|
60
|
+
server.serve_forever()
|
simthinkd/toy.py
ADDED
|
@@ -0,0 +1,53 @@
|
|
|
1
|
+
"""A toy domain for trying `simthinkd.fit` in seconds: a part passes an inspection station on a conveyor.
|
|
2
|
+
|
|
3
|
+
import simthinkd
|
|
4
|
+
from simthinkd import toy
|
|
5
|
+
d = simthinkd.fit(toy.examples(3000), toy.ACTIONS, goal=toy.GOAL, out="inspection")
|
|
6
|
+
d.decide("part: defect dent | severity severe | image clear | belt normal | rework queue long")
|
|
7
|
+
|
|
8
|
+
Replace `observe()` and `teacher()` with your simulator's perception (exact state -> short binned sentence)
|
|
9
|
+
and rule teacher (exact state -> action). That pair is all a new domain needs.
|
|
10
|
+
"""
|
|
11
|
+
import random
|
|
12
|
+
|
|
13
|
+
GOAL = 'Ship only good parts without stopping the line.'
|
|
14
|
+
ACTIONS = {
|
|
15
|
+
'PASS': 'Let the part continue to packing.',
|
|
16
|
+
'REJECT': 'Divert the part to the reject bin.',
|
|
17
|
+
'REWORK': 'Send the part to the rework loop.',
|
|
18
|
+
'SLOW_BELT': 'Slow the conveyor for a closer look.',
|
|
19
|
+
'ESCALATE': 'Hold the part and call a human inspector.',
|
|
20
|
+
}
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def observe(rng):
|
|
24
|
+
"""A random part: exact state (for the teacher) and its binned sentence (for the decider)."""
|
|
25
|
+
s = {'defect': rng.choice(['none'] * 5 + ['scratch', 'dent', 'missing', 'stain']),
|
|
26
|
+
'severity': rng.choice(['minor', 'moderate', 'severe']),
|
|
27
|
+
'image': rng.choice(['clear', 'blurry']),
|
|
28
|
+
'belt': rng.choice(['normal', 'fast']),
|
|
29
|
+
'queue': rng.choice(['short', 'long'])}
|
|
30
|
+
text = (f"part: defect {s['defect']} | severity {s['severity']} | image {s['image']} | "
|
|
31
|
+
f"belt {s['belt']} | rework queue {s['queue']}")
|
|
32
|
+
return s, text
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def teacher(s):
|
|
36
|
+
"""The rule teacher the decider learns to imitate."""
|
|
37
|
+
if s['defect'] == 'none':
|
|
38
|
+
return 'SLOW_BELT' if s['image'] == 'blurry' and s['belt'] == 'fast' else 'PASS'
|
|
39
|
+
if s['image'] == 'blurry':
|
|
40
|
+
return 'ESCALATE'
|
|
41
|
+
if s['defect'] == 'stain' or (s['severity'] == 'minor' and s['queue'] == 'short'):
|
|
42
|
+
return 'REWORK'
|
|
43
|
+
return 'REJECT'
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def examples(n=3000, seed=0):
|
|
47
|
+
"""n (situation sentence, teacher action) pairs."""
|
|
48
|
+
rng = random.Random(seed)
|
|
49
|
+
out = []
|
|
50
|
+
for _ in range(n):
|
|
51
|
+
s, text = observe(rng)
|
|
52
|
+
out.append((text, teacher(s)))
|
|
53
|
+
return out
|
simthinkd/train.py
ADDED
|
@@ -0,0 +1,222 @@
|
|
|
1
|
+
"""Train a decider from your own examples. Needs PyTorch (`pip install "simthinkd[train]"`); inference needs only NumPy.
|
|
2
|
+
|
|
3
|
+
import simthinkd
|
|
4
|
+
d = simthinkd.fit(examples, actions, goal="Ship only good parts.", out="my_decider")
|
|
5
|
+
d.decide("part: defect dent | severity severe | image clear")
|
|
6
|
+
|
|
7
|
+
`examples` is a list of (situation sentence, action name) pairs, usually written by a rule teacher running in your
|
|
8
|
+
simulator. The model is a randomly initialised small ranker (no pretrained weights), trained with sampled residual
|
|
9
|
+
rewards against a proper scoring rule (Brier); the checkpoint with the lowest validation Brier is kept and its
|
|
10
|
+
temperature is fitted on a separate calibration split. Splits are made by a hash of each example's id.
|
|
11
|
+
|
|
12
|
+
Command line (folder of train/validation/calibration .jsonl.gz rows, see docs/TRAINING.md):
|
|
13
|
+
simthinkd train --data data/toy --out models/toy --steps 600
|
|
14
|
+
"""
|
|
15
|
+
import copy
|
|
16
|
+
import gzip
|
|
17
|
+
import hashlib
|
|
18
|
+
import json
|
|
19
|
+
import os
|
|
20
|
+
import time
|
|
21
|
+
from datetime import datetime, timezone
|
|
22
|
+
from pathlib import Path
|
|
23
|
+
|
|
24
|
+
import numpy as np
|
|
25
|
+
|
|
26
|
+
from .core import Decider, build_request
|
|
27
|
+
from .policy import encode
|
|
28
|
+
|
|
29
|
+
os.environ.setdefault('CUBLAS_WORKSPACE_CONFIG', ':4096:8')
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def _torch():
|
|
33
|
+
try:
|
|
34
|
+
import torch
|
|
35
|
+
except ImportError as error:
|
|
36
|
+
raise ImportError('training needs PyTorch: pip install "simthinkd[train]"') from error
|
|
37
|
+
return torch
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def sha(path):
|
|
41
|
+
return hashlib.sha256(Path(path).read_bytes()).hexdigest()
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def split_for(name):
|
|
45
|
+
digit = int(hashlib.sha256(name.encode()).hexdigest(), 16) % 10
|
|
46
|
+
return 'validation' if digit == 0 else 'calibration' if digit == 1 else 'train'
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _ranker(torch, width):
|
|
50
|
+
class Ranker(torch.nn.Module):
|
|
51
|
+
def __init__(self):
|
|
52
|
+
super().__init__()
|
|
53
|
+
self.local = torch.nn.Linear(width, 128)
|
|
54
|
+
self.context = torch.nn.Linear(128 * 3, 96)
|
|
55
|
+
self.score = torch.nn.Linear(96, 1)
|
|
56
|
+
|
|
57
|
+
def forward(self, x, mask):
|
|
58
|
+
h = torch.relu(self.local(x))
|
|
59
|
+
mean = (h * mask[..., None]).sum(1) / mask.sum(1, keepdim=True)
|
|
60
|
+
maximum = h.masked_fill(~mask[..., None], -1e9).max(1).values
|
|
61
|
+
c = torch.cat([mean, maximum], -1)[:, None].expand(-1, h.shape[1], -1)
|
|
62
|
+
return self.score(torch.relu(self.context(torch.cat([h, c], -1)))).squeeze(-1)
|
|
63
|
+
return Ranker()
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def _probabilities(torch, scores, groups, mask, temperature=1.):
|
|
67
|
+
scores = scores / temperature
|
|
68
|
+
membership = (groups[..., None] == torch.arange(8, device=scores.device)) & mask[..., None]
|
|
69
|
+
count = membership.sum(1)
|
|
70
|
+
expanded = scores[..., None].masked_fill(~membership, -1e9)
|
|
71
|
+
denom = torch.logsumexp(expanded, 1)
|
|
72
|
+
op_logits = (denom - count.clamp_min(1).log()).masked_fill(count == 0, -1e9)
|
|
73
|
+
op = torch.softmax(op_logits, -1)
|
|
74
|
+
conditional = torch.exp(expanded - denom[:, None]) * membership
|
|
75
|
+
return (conditional * op[:, None]).sum(-1), op
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def _encode_rows(torch, items, device):
|
|
79
|
+
features, groups, labels = [], [], []
|
|
80
|
+
for row in items:
|
|
81
|
+
x, g, candidates, _ = encode(row['request'])
|
|
82
|
+
gold = row['expected']
|
|
83
|
+
op = gold['operation']
|
|
84
|
+
target = gold.get(op.lower() + '_target')
|
|
85
|
+
matches = [i for i, c in enumerate(candidates) if (c['operation'], c['target']) == (op, target)]
|
|
86
|
+
if not matches:
|
|
87
|
+
raise ValueError(f'example {row.get("id")}: expected action {op!r} is not among the offered actions')
|
|
88
|
+
features.append(x)
|
|
89
|
+
groups.append(g)
|
|
90
|
+
labels.append(matches[0])
|
|
91
|
+
maximum = max(len(x) for x in features)
|
|
92
|
+
width = features[0].shape[1]
|
|
93
|
+
xx = np.zeros((len(items), maximum, width), np.float32)
|
|
94
|
+
gg = np.zeros((len(items), maximum), np.int64)
|
|
95
|
+
mm = np.zeros((len(items), maximum), bool)
|
|
96
|
+
for i, (x, g) in enumerate(zip(features, groups)):
|
|
97
|
+
xx[i, :len(x)] = x
|
|
98
|
+
gg[i, :len(x)] = g
|
|
99
|
+
mm[i, :len(x)] = True
|
|
100
|
+
return tuple(torch.as_tensor(v, device=device) for v in [xx, gg, mm, np.array(labels)]), width
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
def _brier(torch, p, y):
|
|
104
|
+
q = torch.nn.functional.one_hot(y, p.shape[1]).float()
|
|
105
|
+
return ((p - q) ** 2).sum(-1).mean()
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def _assess(torch, model, part):
|
|
109
|
+
x, g, m, y = part
|
|
110
|
+
loss = 0.
|
|
111
|
+
with torch.inference_mode():
|
|
112
|
+
for start in range(0, len(x), 128):
|
|
113
|
+
sl = slice(start, start + 128)
|
|
114
|
+
p, _ = _probabilities(torch, model(x[sl], m[sl]), g[sl], m[sl])
|
|
115
|
+
loss += float(_brier(torch, p, y[sl])) * len(x[sl])
|
|
116
|
+
return loss / len(x)
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
def _fit_one(torch, initial, train, validation, steps, seed, log):
|
|
120
|
+
model = copy.deepcopy(initial)
|
|
121
|
+
model.train()
|
|
122
|
+
optimizer = torch.optim.AdamW(model.parameters(), lr=.0015, weight_decay=.0001)
|
|
123
|
+
rng = torch.Generator(device=train[0].device).manual_seed(seed + 2121)
|
|
124
|
+
rewards = torch.Generator(device=train[0].device).manual_seed(seed + 7171)
|
|
125
|
+
x, g, mask, y = train
|
|
126
|
+
best, weights, events, best_step = float('inf'), copy.deepcopy(model.state_dict()), [], 0
|
|
127
|
+
started = time.perf_counter()
|
|
128
|
+
for step in range(steps + 1):
|
|
129
|
+
if step % 100 == 0 or step == steps:
|
|
130
|
+
val = _assess(torch, model, validation)
|
|
131
|
+
events.append({'step': step, 'validation_brier': val})
|
|
132
|
+
if step > 0 and val < best:
|
|
133
|
+
best, weights, best_step = val, copy.deepcopy(model.state_dict()), step
|
|
134
|
+
log({'seed': seed, 'step': step, 'validation_brier': round(val, 6)})
|
|
135
|
+
if step == steps:
|
|
136
|
+
break
|
|
137
|
+
batch = torch.randint(len(x), (96,), generator=rng, device=x.device)
|
|
138
|
+
p, _ = _probabilities(torch, model(x[batch], mask[batch]), g[batch], mask[batch])
|
|
139
|
+
q = torch.nn.functional.one_hot(y[batch], p.shape[1]).float()
|
|
140
|
+
reward = 2 * (q - p).detach()
|
|
141
|
+
baseline = (p.detach() * reward).sum(-1, keepdim=True)
|
|
142
|
+
sampled = torch.multinomial(p.detach(), 32, replacement=True, generator=rewards)
|
|
143
|
+
loss = -((reward - baseline).gather(1, sampled) * p.clamp_min(1e-30).log().gather(1, sampled)).mean()
|
|
144
|
+
optimizer.zero_grad(set_to_none=True)
|
|
145
|
+
loss.backward()
|
|
146
|
+
optimizer.step()
|
|
147
|
+
model.load_state_dict(weights)
|
|
148
|
+
model.eval()
|
|
149
|
+
return model, {'seed': seed, 'steps_run': steps, 'selected_step': best_step,
|
|
150
|
+
'seconds': time.perf_counter() - started, 'validation_brier': best, 'events': events}
|
|
151
|
+
|
|
152
|
+
|
|
153
|
+
def fit_rows(rows, out, steps=600, seed=31, card=None, quiet=False):
|
|
154
|
+
"""Train from protocol rows {"id", "request", "expected": {"operation", "<op>_target"?}}; returns a Decider."""
|
|
155
|
+
torch = _torch()
|
|
156
|
+
out = Path(out)
|
|
157
|
+
out.mkdir(parents=True, exist_ok=True)
|
|
158
|
+
log = (lambda record: None) if quiet else (lambda record: print(json.dumps(record), flush=True))
|
|
159
|
+
torch.set_num_threads(min(4, os.cpu_count() or 1))
|
|
160
|
+
torch.use_deterministic_algorithms(True)
|
|
161
|
+
device = 'cuda' if torch.cuda.is_available() else 'cpu'
|
|
162
|
+
splits = {'train': [], 'validation': [], 'calibration': []}
|
|
163
|
+
for row in rows:
|
|
164
|
+
splits[row.get('split') or split_for(row['id'])].append(row)
|
|
165
|
+
if min(len(v) for v in splits.values()) == 0:
|
|
166
|
+
raise ValueError('need enough examples for train/validation/calibration splits (about 100 or more)')
|
|
167
|
+
parts = {}
|
|
168
|
+
for name, items in splits.items():
|
|
169
|
+
parts[name], width = _encode_rows(torch, items, device)
|
|
170
|
+
torch.manual_seed(seed)
|
|
171
|
+
initial = _ranker(torch, width).to(device)
|
|
172
|
+
model, record = _fit_one(torch, initial, parts['train'], parts['validation'], steps, seed, log)
|
|
173
|
+
cx, cg, cm, cy = parts['calibration']
|
|
174
|
+
with torch.inference_mode():
|
|
175
|
+
scores = torch.cat([model(cx[i:i + 128], cm[i:i + 128]) for i in range(0, len(cx), 128)])
|
|
176
|
+
temperatures = np.geomspace(.5, 3., 25)
|
|
177
|
+
losses = [float(_brier(torch, _probabilities(torch, scores, cg, cm, float(t))[0], cy)) for t in temperatures]
|
|
178
|
+
temperature = float(temperatures[int(np.argmin(losses))])
|
|
179
|
+
arrays = {k: v.detach().cpu().numpy() for k, v in model.state_dict().items()}
|
|
180
|
+
path = out / 'weights.npz'
|
|
181
|
+
np.savez_compressed(path, **arrays, temperature=np.array(temperature))
|
|
182
|
+
card = dict(card or {})
|
|
183
|
+
card.update({'weights': path.name, 'sha256': sha(path), 'temperature': temperature,
|
|
184
|
+
'parameters': sum(p.numel() for p in model.parameters()),
|
|
185
|
+
'examples': {k: len(v) for k, v in splits.items()}, 'training': record,
|
|
186
|
+
'calibration_brier': min(losses), 'device': device, 'torch': torch.__version__,
|
|
187
|
+
'created_at': datetime.now(timezone.utc).isoformat(),
|
|
188
|
+
'initialization': 'random (no pretrained weights)'})
|
|
189
|
+
(out / 'decider.json').write_text(json.dumps(card, ensure_ascii=False, indent=1), encoding='utf8')
|
|
190
|
+
return Decider(str(out))
|
|
191
|
+
|
|
192
|
+
|
|
193
|
+
def fit(examples, actions, goal='', out='my_decider', steps=600, seed=31, name=None, quiet=False):
|
|
194
|
+
"""Train a decider from (situation sentence, action name) pairs. Returns a ready Decider."""
|
|
195
|
+
if not actions or len(actions) > 8:
|
|
196
|
+
raise ValueError('offer between 1 and 8 actions: {name: short description}')
|
|
197
|
+
rows = []
|
|
198
|
+
for i, item in enumerate(examples):
|
|
199
|
+
state, action = (item['state'], item['action']) if isinstance(item, dict) else item
|
|
200
|
+
if action not in actions:
|
|
201
|
+
raise ValueError(f'example {i}: action {action!r} is not in actions')
|
|
202
|
+
row_id = f'ex{i:07d}:' + hashlib.sha256(state.encode()).hexdigest()[:8]
|
|
203
|
+
rows.append({'id': row_id, 'request': build_request(state, actions, goal), 'expected': {'operation': action}})
|
|
204
|
+
card = {'name': name or Path(out).name, 'actions': dict(actions), 'goal': goal,
|
|
205
|
+
'url': 'sim://custom', 'title': '', 'model': 'simthinkd'}
|
|
206
|
+
return fit_rows(rows, out, steps=steps, seed=seed, card=card, quiet=quiet)
|
|
207
|
+
|
|
208
|
+
|
|
209
|
+
def fit_dir(data, out, steps=600, seed=31, quiet=False):
|
|
210
|
+
"""Train from a folder holding train/validation/calibration .jsonl.gz protocol rows."""
|
|
211
|
+
rows = []
|
|
212
|
+
for split in ['train', 'validation', 'calibration']:
|
|
213
|
+
with gzip.open(Path(data) / f'{split}.jsonl.gz', 'rt', encoding='utf8') as stream:
|
|
214
|
+
rows += [dict(json.loads(line), split=split) for line in stream]
|
|
215
|
+
first = rows[0]['request']
|
|
216
|
+
card = {'name': Path(out).name, 'goal': first['questions']['operation'].get('instructions', {}).get('goal', ''),
|
|
217
|
+
'actions': {k: (v.get('description') if isinstance(v, dict) else v)
|
|
218
|
+
for k, v in first['questions']['operation']['criteria'].items()},
|
|
219
|
+
'url': first['state']['page'].get('url', 'sim://custom'), 'title': first['state']['page'].get('title', ''),
|
|
220
|
+
'model': first.get('model', 'simthinkd'),
|
|
221
|
+
'data_sha256': {s: sha(Path(data) / f'{s}.jsonl.gz') for s in ['train', 'validation', 'calibration']}}
|
|
222
|
+
return fit_rows(rows, out, steps=steps, seed=seed, card=card, quiet=quiet)
|
|
Binary file
|
|
Binary file
|