sdm-learn 0.1.0.dev0__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.
- sdm/__init__.py +71 -0
- sdm/__main__.py +51 -0
- sdm/data.py +238 -0
- sdm/format.md +84 -0
- sdm/formats.py +1182 -0
- sdm/gateway.py +169 -0
- sdm/learn.py +699 -0
- sdm/signals.py +386 -0
- sdm_learn-0.1.0.dev0.dist-info/METADATA +35 -0
- sdm_learn-0.1.0.dev0.dist-info/RECORD +14 -0
- sdm_learn-0.1.0.dev0.dist-info/WHEEL +5 -0
- sdm_learn-0.1.0.dev0.dist-info/entry_points.txt +2 -0
- sdm_learn-0.1.0.dev0.dist-info/licenses/LICENSE +201 -0
- sdm_learn-0.1.0.dev0.dist-info/top_level.txt +1 -0
sdm/gateway.py
ADDED
|
@@ -0,0 +1,169 @@
|
|
|
1
|
+
"""The gateway to an LLM, and the environment checks that come with it.
|
|
2
|
+
|
|
3
|
+
Deliberately not called ``model``: in this project the skill document *is* the
|
|
4
|
+
model. This file is only the wire to a language model that reads it.
|
|
5
|
+
|
|
6
|
+
The bottom of the stack, so nothing here knows about skills, datasets, or the
|
|
7
|
+
learning loop. It turns a list of content blocks into one completion, retries
|
|
8
|
+
what is worth retrying, and refuses to paper over a truncated response. It is
|
|
9
|
+
also the one seam the offline checks replace, which is what keeps them honest:
|
|
10
|
+
there is no second path to the network.
|
|
11
|
+
|
|
12
|
+
``SetupNeeded`` lives here because this is the module every other one can
|
|
13
|
+
reach. It marks the failures a traceback cannot fix: no API key, no SVG
|
|
14
|
+
rasterizer, no numpy.
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
from __future__ import annotations
|
|
18
|
+
|
|
19
|
+
import base64
|
|
20
|
+
import contextlib
|
|
21
|
+
import io
|
|
22
|
+
import json
|
|
23
|
+
import os
|
|
24
|
+
import time
|
|
25
|
+
from pathlib import Path
|
|
26
|
+
|
|
27
|
+
import httpx
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
class SetupNeeded(Exception):
|
|
31
|
+
"""Something about the environment needs a human, not a traceback."""
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
GATEWAY_URL = os.environ.get(
|
|
35
|
+
"SDM_GATEWAY_URL", "https://inference-api.nvidia.com/v1/chat/completions"
|
|
36
|
+
)
|
|
37
|
+
MODEL = os.environ.get("SDM_MODEL", "azure/anthropic/claude-opus-5")
|
|
38
|
+
EFFORT = os.environ.get("SDM_EFFORT", "low")
|
|
39
|
+
KEY_FILE = Path.home() / ".sdm_key"
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def _api_key() -> str:
|
|
43
|
+
key = os.environ.get("SDM_API_KEY")
|
|
44
|
+
if key:
|
|
45
|
+
return key.strip()
|
|
46
|
+
path = Path(os.environ.get("SDM_API_KEY_FILE", KEY_FILE))
|
|
47
|
+
if path.is_file():
|
|
48
|
+
return path.read_text().strip()
|
|
49
|
+
raise SetupNeeded(
|
|
50
|
+
"no model to talk to yet. Set SDM_API_KEY, or write the key to "
|
|
51
|
+
f"{KEY_FILE}."
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def client(read_timeout: float = 600.0, connect_timeout: float = 30.0):
|
|
56
|
+
"""An authenticated session for the gateway."""
|
|
57
|
+
return httpx.Client(
|
|
58
|
+
headers={"Authorization": f"Bearer {_api_key()}",
|
|
59
|
+
"Content-Type": "application/json"},
|
|
60
|
+
timeout=httpx.Timeout(read_timeout, connect=connect_timeout),
|
|
61
|
+
)
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
@contextlib.contextmanager
|
|
65
|
+
def _borrowed(session):
|
|
66
|
+
"""Yield a session, closing it only if we were the one who opened it."""
|
|
67
|
+
owned = session is None
|
|
68
|
+
session = session or client()
|
|
69
|
+
try:
|
|
70
|
+
yield session
|
|
71
|
+
finally:
|
|
72
|
+
if owned:
|
|
73
|
+
session.close()
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def _stream(session, payload: dict, deadline: float = 1800.0):
|
|
77
|
+
chunks, finish = [], None
|
|
78
|
+
started = time.monotonic()
|
|
79
|
+
with session.stream("POST", GATEWAY_URL,
|
|
80
|
+
json={**payload, "stream": True}) as response:
|
|
81
|
+
if response.status_code != 200:
|
|
82
|
+
response.read()
|
|
83
|
+
raise httpx.HTTPStatusError(
|
|
84
|
+
f"{response.status_code}: {response.text[:200]}",
|
|
85
|
+
request=response.request, response=response,
|
|
86
|
+
)
|
|
87
|
+
for line in response.iter_lines():
|
|
88
|
+
if time.monotonic() - started > deadline:
|
|
89
|
+
raise httpx.TransportError(f"deadline {deadline}s hit")
|
|
90
|
+
if not line.startswith("data: "):
|
|
91
|
+
continue
|
|
92
|
+
body = line[6:]
|
|
93
|
+
if body.strip() == "[DONE]":
|
|
94
|
+
break
|
|
95
|
+
event = json.loads(body)
|
|
96
|
+
if not event.get("choices"):
|
|
97
|
+
continue
|
|
98
|
+
choice = event["choices"][0]
|
|
99
|
+
if (choice.get("delta") or {}).get("content"):
|
|
100
|
+
chunks.append(choice["delta"]["content"])
|
|
101
|
+
if choice.get("finish_reason"):
|
|
102
|
+
finish = choice["finish_reason"]
|
|
103
|
+
return "".join(chunks), finish
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
def complete(session, messages: list, max_tokens: int,
|
|
107
|
+
effort: str | None = None) -> str:
|
|
108
|
+
"""One chat completion, retrying through the failures a long run hits."""
|
|
109
|
+
payload = {
|
|
110
|
+
"model": MODEL,
|
|
111
|
+
"max_tokens": max_tokens,
|
|
112
|
+
"output_config": {"effort": effort or EFFORT},
|
|
113
|
+
"messages": messages,
|
|
114
|
+
}
|
|
115
|
+
for attempt in range(30):
|
|
116
|
+
try:
|
|
117
|
+
content, finish = _stream(session, payload)
|
|
118
|
+
if finish == "content_filter":
|
|
119
|
+
raise RuntimeError("model refused to answer")
|
|
120
|
+
if not content or finish == "length":
|
|
121
|
+
payload["max_tokens"] = min(payload["max_tokens"] * 2, 32000)
|
|
122
|
+
print(f" [truncated ({finish}), retrying with "
|
|
123
|
+
f"max_tokens={payload['max_tokens']}]", flush=True)
|
|
124
|
+
continue
|
|
125
|
+
if finish != "stop":
|
|
126
|
+
print(f" [incomplete stream (finish={finish!r}), retrying]",
|
|
127
|
+
flush=True)
|
|
128
|
+
continue
|
|
129
|
+
return content
|
|
130
|
+
except (httpx.HTTPStatusError, httpx.TransportError,
|
|
131
|
+
json.JSONDecodeError) as error:
|
|
132
|
+
wait = min(15 * (attempt + 1), 600)
|
|
133
|
+
print(f" [retry in {wait}s: {type(error).__name__}: "
|
|
134
|
+
f"{str(error)[:120]}]", flush=True)
|
|
135
|
+
time.sleep(wait)
|
|
136
|
+
raise RuntimeError("too many retries")
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
def _png_block(png: bytes) -> dict:
|
|
140
|
+
"""One content block holding already-encoded PNG bytes."""
|
|
141
|
+
encoded = base64.b64encode(png).decode("ascii")
|
|
142
|
+
return {"type": "image_url",
|
|
143
|
+
"image_url": {"url": f"data:image/png;base64,{encoded}"}}
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def image_block(image, size: int = 224) -> dict:
|
|
147
|
+
"""Encode an RGB array, PIL image, or image path as one content block."""
|
|
148
|
+
from PIL import Image
|
|
149
|
+
|
|
150
|
+
if isinstance(image, Image.Image):
|
|
151
|
+
pil = image.convert("RGB")
|
|
152
|
+
elif isinstance(image, (str, Path)):
|
|
153
|
+
with Image.open(image) as source:
|
|
154
|
+
pil = source.convert("RGB")
|
|
155
|
+
elif hasattr(image, "shape"):
|
|
156
|
+
import numpy as np
|
|
157
|
+
|
|
158
|
+
if image.ndim != 3 or image.shape[2] != 3:
|
|
159
|
+
raise ValueError("expected one HxWx3 RGB image")
|
|
160
|
+
pil = Image.fromarray(np.asarray(image).astype("uint8"), mode="RGB")
|
|
161
|
+
else:
|
|
162
|
+
raise TypeError("expected an RGB array, PIL image, or image path")
|
|
163
|
+
pil = pil.resize((size, size), Image.Resampling.NEAREST)
|
|
164
|
+
buffer = io.BytesIO()
|
|
165
|
+
pil.save(buffer, format="PNG", optimize=False)
|
|
166
|
+
return _png_block(buffer.getvalue())
|
|
167
|
+
|
|
168
|
+
|
|
169
|
+
# --------------------------------------------------------------------------
|