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.
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)
@@ -0,0 +1,2 @@
1
+ a164d90514f172ff2f1ac9452aaf83a637ffd8ef3136d0f553179e5c7ab035d8 *doom-corridor.npz
2
+ a814587e03478f451f1ab478558de7e076dbe9a3de9cb2ee4cc84124cdfa2a19 *doom-defend.npz
Binary file
Binary file