steadybatch 0.2.1__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.
- steadybatch/__init__.py +7 -0
- steadybatch/bench.py +317 -0
- steadybatch/chunking.py +39 -0
- steadybatch/ids.py +32 -0
- steadybatch/models.py +85 -0
- steadybatch/providers/__init__.py +35 -0
- steadybatch/providers/anthropic_batch.py +95 -0
- steadybatch/providers/base.py +60 -0
- steadybatch/providers/fake.py +104 -0
- steadybatch/providers/gemini_batch.py +126 -0
- steadybatch/providers/openai_batch.py +127 -0
- steadybatch/providers/vllm_offline.py +86 -0
- steadybatch/runner.py +216 -0
- steadybatch/store.py +193 -0
- steadybatch/validate.py +39 -0
- steadybatch-0.2.1.dist-info/METADATA +198 -0
- steadybatch-0.2.1.dist-info/RECORD +20 -0
- steadybatch-0.2.1.dist-info/WHEEL +4 -0
- steadybatch-0.2.1.dist-info/entry_points.txt +2 -0
- steadybatch-0.2.1.dist-info/licenses/LICENSE +202 -0
steadybatch/__init__.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
1
|
+
"""steadybatch: run millions of LLM requests through batch APIs without silent failures."""
|
|
2
|
+
|
|
3
|
+
from .models import Outcome, Request, Result
|
|
4
|
+
from .runner import RunReport, Runner
|
|
5
|
+
|
|
6
|
+
__all__ = ["Outcome", "Request", "Result", "RunReport", "Runner"]
|
|
7
|
+
__version__ = "0.2.1"
|
steadybatch/bench.py
ADDED
|
@@ -0,0 +1,317 @@
|
|
|
1
|
+
"""Benchmark harness: run one workload on one provider and write down what happened.
|
|
2
|
+
|
|
3
|
+
steadybatch-bench run --provider openai --model gpt-4o-mini \\
|
|
4
|
+
--data examples/support_tickets.jsonl --schema examples/support_ticket.schema.json \\
|
|
5
|
+
--out runs/openai-1 --prices examples/prices.example.json
|
|
6
|
+
|
|
7
|
+
steadybatch-bench compare runs/openai-1 runs/openai-2
|
|
8
|
+
|
|
9
|
+
``run`` writes results.jsonl (one line per record, in input order) and
|
|
10
|
+
summary.json. ``compare`` checks two runs of the same workload against each
|
|
11
|
+
other, which is how we measure reproducibility.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from __future__ import annotations
|
|
15
|
+
|
|
16
|
+
import argparse
|
|
17
|
+
import json
|
|
18
|
+
import logging
|
|
19
|
+
import statistics
|
|
20
|
+
import sys
|
|
21
|
+
from pathlib import Path
|
|
22
|
+
from typing import Any
|
|
23
|
+
|
|
24
|
+
from .models import Outcome, Request
|
|
25
|
+
from .providers import get_provider
|
|
26
|
+
from .runner import Runner
|
|
27
|
+
|
|
28
|
+
SYSTEM_PROMPT = (
|
|
29
|
+
"You read customer support conversations and extract structured facts. "
|
|
30
|
+
"Answer with a single JSON object and nothing else."
|
|
31
|
+
)
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def load_records(path: Path, limit: int | None) -> list[dict[str, Any]]:
|
|
35
|
+
records = []
|
|
36
|
+
with path.open() as f:
|
|
37
|
+
for line in f:
|
|
38
|
+
if line.strip():
|
|
39
|
+
records.append(json.loads(line))
|
|
40
|
+
if limit and len(records) >= limit:
|
|
41
|
+
break
|
|
42
|
+
return records
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def to_requests(records: list[dict[str, Any]], max_tokens: int, instructions: str = "") -> list[Request]:
|
|
46
|
+
system = SYSTEM_PROMPT + ("\n\n" + instructions.strip() if instructions.strip() else "")
|
|
47
|
+
return [
|
|
48
|
+
Request(key=str(r["key"]), system=system, max_tokens=max_tokens, temperature=0.0,
|
|
49
|
+
messages=[{"role": "user", "content": r["text"]}])
|
|
50
|
+
for r in records
|
|
51
|
+
]
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def accuracy(records: list[dict[str, Any]], results_by_key: dict[str, Any]) -> dict[str, float]:
|
|
55
|
+
"""Field-level accuracy against the 'labels' in the dataset, when present."""
|
|
56
|
+
scores: dict[str, list[int]] = {}
|
|
57
|
+
for r in records:
|
|
58
|
+
labels = r.get("labels")
|
|
59
|
+
data = results_by_key.get(str(r["key"]))
|
|
60
|
+
if not labels or not isinstance(data, dict):
|
|
61
|
+
continue
|
|
62
|
+
for field, expected in labels.items():
|
|
63
|
+
scores.setdefault(field, []).append(int(data.get(field) == expected))
|
|
64
|
+
return {f: round(sum(v) / len(v), 4) for f, v in scores.items() if v}
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def cmd_run(args: argparse.Namespace) -> int:
|
|
68
|
+
records = load_records(Path(args.data), args.limit)
|
|
69
|
+
schema = json.loads(Path(args.schema).read_text()) if args.schema else None
|
|
70
|
+
out = Path(args.out)
|
|
71
|
+
out.mkdir(parents=True, exist_ok=True)
|
|
72
|
+
|
|
73
|
+
runner = Runner(get_provider(args.provider, args.model, thinking=args.thinking), args.model, response_schema=schema,
|
|
74
|
+
checkpoint=out / "checkpoint.sqlite", max_attempts=args.max_attempts,
|
|
75
|
+
poll_every=args.poll_every, batch_size=args.batch_size,
|
|
76
|
+
max_open_batches=args.max_open_batches)
|
|
77
|
+
instructions = Path(args.instructions).read_text() if args.instructions else ""
|
|
78
|
+
report = runner.run(to_requests(records, args.max_tokens, instructions))
|
|
79
|
+
|
|
80
|
+
with (out / "results.jsonl").open("w") as f:
|
|
81
|
+
for r in report.results:
|
|
82
|
+
f.write(json.dumps({"key": r.key, "outcome": r.outcome.value, "data": r.data,
|
|
83
|
+
"error": r.error, "attempts": r.attempts, "history": r.history}) + "\n")
|
|
84
|
+
|
|
85
|
+
in_tok = sum(r.input_tokens for r in report.results)
|
|
86
|
+
out_tok = sum(r.output_tokens for r in report.results)
|
|
87
|
+
summary: dict[str, Any] = {
|
|
88
|
+
"provider": args.provider,
|
|
89
|
+
"model": args.model,
|
|
90
|
+
"instructions_file": args.instructions,
|
|
91
|
+
"thinking": args.thinking or "model default",
|
|
92
|
+
"max_tokens": args.max_tokens,
|
|
93
|
+
"records": len(records),
|
|
94
|
+
"ok": report.ok,
|
|
95
|
+
"failed_after_retries": report.failed,
|
|
96
|
+
"outcomes": {o.value: sum(r.outcome is o for r in report.results) for o in Outcome},
|
|
97
|
+
"lines_missing_seen": report.lines_missing,
|
|
98
|
+
"lines_errored_seen": report.lines_errored,
|
|
99
|
+
"lines_invalid_seen": report.lines_invalid,
|
|
100
|
+
"batches_submitted": report.batches_submitted,
|
|
101
|
+
"capacity_refusals": report.capacity_refusals,
|
|
102
|
+
"submit_retries": report.submit_retries,
|
|
103
|
+
"retried_records": sum(r.attempts > 1 for r in report.results),
|
|
104
|
+
"input_tokens": in_tok,
|
|
105
|
+
"output_tokens": out_tok,
|
|
106
|
+
"wall_seconds": round(report.seconds, 1),
|
|
107
|
+
}
|
|
108
|
+
if report.batch_seconds:
|
|
109
|
+
summary["batch_seconds_median"] = round(statistics.median(report.batch_seconds), 1)
|
|
110
|
+
summary["batch_seconds_max"] = round(max(report.batch_seconds), 1)
|
|
111
|
+
if args.prices:
|
|
112
|
+
p = json.loads(Path(args.prices).read_text())[args.provider]
|
|
113
|
+
cost = (in_tok * p["input_per_mtok"] + out_tok * p["output_per_mtok"]) / 1e6 * p.get("batch_multiplier", 1.0)
|
|
114
|
+
summary["cost_usd"] = round(cost, 4)
|
|
115
|
+
summary["cost_per_1000_records_usd"] = round(cost / max(len(records), 1) * 1000, 4)
|
|
116
|
+
acc = accuracy(records, {r.key: r.data for r in report.results})
|
|
117
|
+
if acc:
|
|
118
|
+
summary["field_accuracy"] = acc
|
|
119
|
+
(out / "summary.json").write_text(json.dumps(summary, indent=2))
|
|
120
|
+
print(json.dumps(summary, indent=2))
|
|
121
|
+
return 0
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
def cmd_compare(args: argparse.Namespace) -> int:
|
|
125
|
+
def load(d: str) -> dict[str, Any]:
|
|
126
|
+
path = Path(d) / "results.jsonl"
|
|
127
|
+
if not path.exists() and path.with_suffix(".jsonl.gz").exists():
|
|
128
|
+
import gzip # large published runs are stored compressed
|
|
129
|
+
text = gzip.decompress(path.with_suffix(".jsonl.gz").read_bytes()).decode("utf-8")
|
|
130
|
+
else:
|
|
131
|
+
text = path.read_text()
|
|
132
|
+
rows = [json.loads(l) for l in text.splitlines() if l]
|
|
133
|
+
return {r["key"]: r["data"] for r in rows if r["outcome"] == "ok"}
|
|
134
|
+
|
|
135
|
+
a, b = load(args.run_a), load(args.run_b)
|
|
136
|
+
both = sorted(set(a) & set(b))
|
|
137
|
+
exact = sum(a[k] == b[k] for k in both)
|
|
138
|
+
fields: dict[str, list[int]] = {}
|
|
139
|
+
for k in both:
|
|
140
|
+
if isinstance(a[k], dict) and isinstance(b[k], dict):
|
|
141
|
+
for f in set(a[k]) | set(b[k]):
|
|
142
|
+
fields.setdefault(f, []).append(int(a[k].get(f) == b[k].get(f)))
|
|
143
|
+
result = {
|
|
144
|
+
"records_in_both": len(both),
|
|
145
|
+
"exact_match_rate": round(exact / len(both), 4) if both else None,
|
|
146
|
+
"field_match_rate": {f: round(sum(v) / len(v), 4) for f, v in sorted(fields.items())},
|
|
147
|
+
}
|
|
148
|
+
print(json.dumps(result, indent=2))
|
|
149
|
+
return 0
|
|
150
|
+
|
|
151
|
+
|
|
152
|
+
REPORT_COLUMNS = [
|
|
153
|
+
("run", "Run"),
|
|
154
|
+
("provider", "Provider"),
|
|
155
|
+
("model", "Model"),
|
|
156
|
+
("records", "Records"),
|
|
157
|
+
("ok", "OK"),
|
|
158
|
+
("failed_after_retries", "Failed"),
|
|
159
|
+
("lines_missing_seen", "Missing"),
|
|
160
|
+
("lines_errored_seen", "Errored"),
|
|
161
|
+
("lines_invalid_seen", "Invalid"),
|
|
162
|
+
("retried_records", "Retried"),
|
|
163
|
+
("batch_seconds_median", "Median batch (s)"),
|
|
164
|
+
("batch_seconds_max", "Slowest batch (s)"),
|
|
165
|
+
("cost_per_1000_records_usd", "Cost per 1k ($)"),
|
|
166
|
+
("mean_field_accuracy", "Accuracy"),
|
|
167
|
+
]
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
def load_summaries(run_dirs: list[str]) -> list[dict[str, Any]]:
|
|
171
|
+
rows = []
|
|
172
|
+
for d in run_dirs:
|
|
173
|
+
path = Path(d) / "summary.json"
|
|
174
|
+
if not path.exists():
|
|
175
|
+
print(f"skipping {d}: no summary.json", file=sys.stderr)
|
|
176
|
+
continue
|
|
177
|
+
s = json.loads(path.read_text())
|
|
178
|
+
s["run"] = Path(d).name
|
|
179
|
+
acc = s.get("field_accuracy") or {}
|
|
180
|
+
if acc:
|
|
181
|
+
s["mean_field_accuracy"] = round(sum(acc.values()) / len(acc), 4)
|
|
182
|
+
rows.append(s)
|
|
183
|
+
return rows
|
|
184
|
+
|
|
185
|
+
|
|
186
|
+
def write_charts(rows: list[dict[str, Any]], out: Path) -> list[str]:
|
|
187
|
+
try:
|
|
188
|
+
import matplotlib
|
|
189
|
+
matplotlib.use("Agg")
|
|
190
|
+
import matplotlib.pyplot as plt
|
|
191
|
+
except ImportError:
|
|
192
|
+
print("matplotlib not installed; skipping charts (pip install matplotlib)", file=sys.stderr)
|
|
193
|
+
return []
|
|
194
|
+
labels = [r["run"] for r in rows]
|
|
195
|
+
made = []
|
|
196
|
+
|
|
197
|
+
fig, ax = plt.subplots(figsize=(max(6, len(rows) * 1.2), 4))
|
|
198
|
+
bottom = [0] * len(rows)
|
|
199
|
+
for key, name in [("lines_missing_seen", "Missing"), ("lines_errored_seen", "Errored"),
|
|
200
|
+
("lines_invalid_seen", "Invalid JSON")]:
|
|
201
|
+
vals = [100 * r.get(key, 0) / max(r.get("records", 1), 1) for r in rows]
|
|
202
|
+
ax.bar(labels, vals, bottom=bottom, label=name)
|
|
203
|
+
bottom = [b + v for b, v in zip(bottom, vals)]
|
|
204
|
+
ax.set_ylabel("% of records (before retries)")
|
|
205
|
+
ax.set_title("Problems caught per run")
|
|
206
|
+
ax.legend(loc="upper left", bbox_to_anchor=(1.01, 1), frameon=False)
|
|
207
|
+
fig.tight_layout()
|
|
208
|
+
fig.savefig(out / "problems.png", dpi=150)
|
|
209
|
+
plt.close(fig)
|
|
210
|
+
made.append("problems.png")
|
|
211
|
+
|
|
212
|
+
costs = [(r["run"], r["cost_per_1000_records_usd"]) for r in rows if "cost_per_1000_records_usd" in r]
|
|
213
|
+
if costs:
|
|
214
|
+
fig, ax = plt.subplots(figsize=(max(6, len(costs) * 1.2), 4))
|
|
215
|
+
ax.bar([c[0] for c in costs], [c[1] for c in costs])
|
|
216
|
+
ax.set_ylabel("USD per 1,000 records (incl. retries)")
|
|
217
|
+
ax.set_title("Cost per 1,000 records")
|
|
218
|
+
fig.tight_layout()
|
|
219
|
+
fig.savefig(out / "cost.png", dpi=150)
|
|
220
|
+
plt.close(fig)
|
|
221
|
+
made.append("cost.png")
|
|
222
|
+
return made
|
|
223
|
+
|
|
224
|
+
|
|
225
|
+
def add_costs(rows: list[dict[str, Any]], prices_path: str) -> None:
|
|
226
|
+
"""Fill in cost for runs whose summary has token counts but no cost yet."""
|
|
227
|
+
prices = json.loads(Path(prices_path).read_text())
|
|
228
|
+
for r in rows:
|
|
229
|
+
p = prices.get(r.get("provider", ""))
|
|
230
|
+
if not p or "cost_usd" in r:
|
|
231
|
+
continue
|
|
232
|
+
if p.get("model") and p["model"] != r.get("model"):
|
|
233
|
+
continue # prices were recorded for a different model
|
|
234
|
+
cost = (r.get("input_tokens", 0) * p["input_per_mtok"] + r.get("output_tokens", 0) * p["output_per_mtok"]) / 1e6 * p.get("batch_multiplier", 1.0)
|
|
235
|
+
r["cost_usd"] = round(cost, 6)
|
|
236
|
+
r["cost_per_1000_records_usd"] = round(cost / max(r.get("records", 1), 1) * 1000, 4)
|
|
237
|
+
|
|
238
|
+
|
|
239
|
+
def cmd_report(args: argparse.Namespace) -> int:
|
|
240
|
+
rows = load_summaries(args.runs)
|
|
241
|
+
if not rows:
|
|
242
|
+
print("no runs with a summary.json found", file=sys.stderr)
|
|
243
|
+
return 1
|
|
244
|
+
if args.prices:
|
|
245
|
+
add_costs(rows, args.prices)
|
|
246
|
+
out = Path(args.out)
|
|
247
|
+
out.mkdir(parents=True, exist_ok=True)
|
|
248
|
+
|
|
249
|
+
def cell(row: dict[str, Any], key: str) -> str:
|
|
250
|
+
v = row.get(key)
|
|
251
|
+
return "" if v is None else str(v)
|
|
252
|
+
|
|
253
|
+
header = [name for _, name in REPORT_COLUMNS]
|
|
254
|
+
lines = ["| " + " | ".join(header) + " |", "|" + "---|" * len(header)]
|
|
255
|
+
for r in rows:
|
|
256
|
+
lines.append("| " + " | ".join(cell(r, k) for k, _ in REPORT_COLUMNS) + " |")
|
|
257
|
+
|
|
258
|
+
import csv
|
|
259
|
+
with (out / "summary.csv").open("w", newline="") as f:
|
|
260
|
+
w = csv.writer(f)
|
|
261
|
+
w.writerow(header)
|
|
262
|
+
for r in rows:
|
|
263
|
+
w.writerow([cell(r, k) for k, _ in REPORT_COLUMNS])
|
|
264
|
+
|
|
265
|
+
charts = [] if args.no_charts else write_charts(rows, out)
|
|
266
|
+
md = ["# Benchmark results", "",
|
|
267
|
+
"Counts of missing, errored and invalid lines are before retries; "
|
|
268
|
+
"OK and Failed are after retries.", "", *lines, ""]
|
|
269
|
+
md += [f"" for c in charts]
|
|
270
|
+
(out / "summary.md").write_text("\n".join(md) + "\n")
|
|
271
|
+
print("\n".join(lines))
|
|
272
|
+
print(f"\nwrote {out / 'summary.md'} and {out / 'summary.csv'}" + (f" and {len(charts)} charts" if charts else ""))
|
|
273
|
+
return 0
|
|
274
|
+
|
|
275
|
+
|
|
276
|
+
def main(argv: list[str] | None = None) -> int:
|
|
277
|
+
parser = argparse.ArgumentParser(prog="steadybatch-bench")
|
|
278
|
+
sub = parser.add_subparsers(dest="cmd", required=True)
|
|
279
|
+
|
|
280
|
+
run = sub.add_parser("run", help="run one workload on one provider")
|
|
281
|
+
run.add_argument("--provider", required=True, choices=["fake", "openai", "anthropic", "vllm", "gemini", "bedrock"])
|
|
282
|
+
run.add_argument("--model", required=True)
|
|
283
|
+
run.add_argument("--data", required=True)
|
|
284
|
+
run.add_argument("--schema")
|
|
285
|
+
run.add_argument("--instructions", help="text file with field definitions, added to the system prompt")
|
|
286
|
+
run.add_argument("--out", required=True)
|
|
287
|
+
run.add_argument("--limit", type=int)
|
|
288
|
+
run.add_argument("--max-tokens", type=int, default=512)
|
|
289
|
+
run.add_argument("--max-attempts", type=int, default=3)
|
|
290
|
+
run.add_argument("--batch-size", type=int, help="at most this many requests per batch")
|
|
291
|
+
run.add_argument("--max-open-batches", type=int,
|
|
292
|
+
help="wait for open batches to finish before submitting more than this many")
|
|
293
|
+
run.add_argument("--thinking", choices=["disabled", "adaptive"],
|
|
294
|
+
help="Anthropic only: set thinking explicitly instead of using the model default")
|
|
295
|
+
run.add_argument("--poll-every", type=float, default=60.0)
|
|
296
|
+
run.add_argument("--prices", help="JSON file with per-provider token prices")
|
|
297
|
+
run.set_defaults(func=cmd_run)
|
|
298
|
+
|
|
299
|
+
cmp_ = sub.add_parser("compare", help="compare two runs of the same workload")
|
|
300
|
+
cmp_.add_argument("run_a")
|
|
301
|
+
cmp_.add_argument("run_b")
|
|
302
|
+
cmp_.set_defaults(func=cmd_compare)
|
|
303
|
+
|
|
304
|
+
rep = sub.add_parser("report", help="turn run summaries into a table, CSV and charts")
|
|
305
|
+
rep.add_argument("runs", nargs="+", help="run folders that contain summary.json")
|
|
306
|
+
rep.add_argument("--out", default="results/latest")
|
|
307
|
+
rep.add_argument("--no-charts", action="store_true")
|
|
308
|
+
rep.add_argument("--prices", help="JSON price file; fills in cost for runs that lack it")
|
|
309
|
+
rep.set_defaults(func=cmd_report)
|
|
310
|
+
|
|
311
|
+
args = parser.parse_args(argv)
|
|
312
|
+
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s")
|
|
313
|
+
return args.func(args)
|
|
314
|
+
|
|
315
|
+
|
|
316
|
+
if __name__ == "__main__":
|
|
317
|
+
sys.exit(main())
|
steadybatch/chunking.py
ADDED
|
@@ -0,0 +1,39 @@
|
|
|
1
|
+
"""Split work into batches that stay under a provider's limits.
|
|
2
|
+
|
|
3
|
+
Providers cap a batch by request count and by file size. Going over the
|
|
4
|
+
size cap does not always produce a clear error, so we stay under both,
|
|
5
|
+
with a safety margin on size.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
from typing import Callable, Iterable, Iterator, TypeVar
|
|
11
|
+
|
|
12
|
+
T = TypeVar("T")
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def chunk(
|
|
16
|
+
items: Iterable[T],
|
|
17
|
+
max_count: int,
|
|
18
|
+
max_bytes: int,
|
|
19
|
+
size_of: Callable[[T], int],
|
|
20
|
+
margin: float = 0.9,
|
|
21
|
+
) -> Iterator[list[T]]:
|
|
22
|
+
if max_count < 1 or max_bytes < 1:
|
|
23
|
+
raise ValueError("limits must be positive")
|
|
24
|
+
byte_budget = int(max_bytes * margin)
|
|
25
|
+
batch: list[T] = []
|
|
26
|
+
used = 0
|
|
27
|
+
for item in items:
|
|
28
|
+
size = size_of(item)
|
|
29
|
+
if size > byte_budget:
|
|
30
|
+
raise ValueError(
|
|
31
|
+
f"a single request is {size} bytes, over the per-batch budget of {byte_budget}"
|
|
32
|
+
)
|
|
33
|
+
if batch and (len(batch) >= max_count or used + size > byte_budget):
|
|
34
|
+
yield batch
|
|
35
|
+
batch, used = [], 0
|
|
36
|
+
batch.append(item)
|
|
37
|
+
used += size
|
|
38
|
+
if batch:
|
|
39
|
+
yield batch
|
steadybatch/ids.py
ADDED
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
"""Stable request IDs.
|
|
2
|
+
|
|
3
|
+
The same record with the same settings always gets the same ID. That is
|
|
4
|
+
what makes retries safe: if a line comes back twice, or you resubmit work
|
|
5
|
+
after a crash, you can tell it is the same request and not count it twice.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import hashlib
|
|
11
|
+
import json
|
|
12
|
+
from typing import Any
|
|
13
|
+
|
|
14
|
+
from .models import Request
|
|
15
|
+
|
|
16
|
+
# Anthropic allows [a-zA-Z0-9_-]{1,64}; OpenAI is looser. Stay within both.
|
|
17
|
+
_PREFIX = "sb-"
|
|
18
|
+
_HASH_CHARS = 32
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def custom_id_for(request: Request, model: str, response_schema: dict[str, Any] | None = None) -> str:
|
|
22
|
+
payload = {
|
|
23
|
+
"key": request.key,
|
|
24
|
+
"model": model,
|
|
25
|
+
"messages": request.messages,
|
|
26
|
+
"system": request.system,
|
|
27
|
+
"max_tokens": request.max_tokens,
|
|
28
|
+
"temperature": request.temperature,
|
|
29
|
+
"schema": response_schema,
|
|
30
|
+
}
|
|
31
|
+
blob = json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=False)
|
|
32
|
+
return _PREFIX + hashlib.sha256(blob.encode("utf-8")).hexdigest()[:_HASH_CHARS]
|
steadybatch/models.py
ADDED
|
@@ -0,0 +1,85 @@
|
|
|
1
|
+
"""Plain data types shared by the runner and the providers."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from dataclasses import dataclass, field
|
|
6
|
+
from enum import Enum
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
@dataclass(frozen=True)
|
|
11
|
+
class Request:
|
|
12
|
+
"""One unit of work: a record you want the model to process.
|
|
13
|
+
|
|
14
|
+
``key`` is your own identifier for the record (a row id, a file name).
|
|
15
|
+
It must be unique within a job; steadybatch uses it to put results
|
|
16
|
+
back in order and to make retries safe.
|
|
17
|
+
"""
|
|
18
|
+
|
|
19
|
+
key: str
|
|
20
|
+
messages: list[dict[str, str]]
|
|
21
|
+
system: str | None = None
|
|
22
|
+
max_tokens: int = 1024
|
|
23
|
+
temperature: float = 0.0
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
@dataclass(frozen=True)
|
|
27
|
+
class PreparedRequest:
|
|
28
|
+
"""A Request plus everything a provider needs to send it."""
|
|
29
|
+
|
|
30
|
+
custom_id: str
|
|
31
|
+
request: Request
|
|
32
|
+
model: str
|
|
33
|
+
response_schema: dict[str, Any] | None = None
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
class BatchState(str, Enum):
|
|
37
|
+
PENDING = "pending" # accepted, not started
|
|
38
|
+
RUNNING = "running"
|
|
39
|
+
DONE = "done" # provider says it finished (results may still be incomplete)
|
|
40
|
+
FAILED = "failed" # the whole batch failed or expired
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
@dataclass
|
|
44
|
+
class BatchStatus:
|
|
45
|
+
state: BatchState
|
|
46
|
+
total: int = 0
|
|
47
|
+
succeeded: int = 0
|
|
48
|
+
failed: int = 0
|
|
49
|
+
detail: str = ""
|
|
50
|
+
# True when the provider refused the work because the account's queue is full,
|
|
51
|
+
# not because anything is wrong with the requests.
|
|
52
|
+
over_capacity: bool = False
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
@dataclass
|
|
56
|
+
class RawResult:
|
|
57
|
+
"""What a provider returned for one line, before any checking."""
|
|
58
|
+
|
|
59
|
+
custom_id: str
|
|
60
|
+
ok: bool
|
|
61
|
+
text: str | None = None
|
|
62
|
+
error: str | None = None
|
|
63
|
+
input_tokens: int = 0
|
|
64
|
+
output_tokens: int = 0
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
class Outcome(str, Enum):
|
|
68
|
+
OK = "ok"
|
|
69
|
+
ERROR = "error" # the provider returned an error for this line
|
|
70
|
+
MISSING = "missing" # the provider never returned this line at all
|
|
71
|
+
INVALID = "invalid" # a response came back but failed the schema check
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
@dataclass
|
|
75
|
+
class Result:
|
|
76
|
+
key: str
|
|
77
|
+
custom_id: str
|
|
78
|
+
outcome: Outcome
|
|
79
|
+
data: Any = None
|
|
80
|
+
text: str | None = None
|
|
81
|
+
error: str | None = None
|
|
82
|
+
attempts: int = 0
|
|
83
|
+
input_tokens: int = 0
|
|
84
|
+
output_tokens: int = 0
|
|
85
|
+
history: list[str] = field(default_factory=list)
|
|
@@ -0,0 +1,35 @@
|
|
|
1
|
+
"""Batch providers. Each one is imported lazily so you only need the SDKs you use."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from .base import BatchProvider
|
|
6
|
+
from .fake import FakeProvider, demo_answer
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def get_provider(name: str, model: str | None = None, *, thinking: str | None = None) -> BatchProvider:
|
|
10
|
+
if name == "fake":
|
|
11
|
+
# Misbehaves like a real API: a few lines dropped, errored or malformed.
|
|
12
|
+
return FakeProvider(answer=demo_answer, drop_rate=0.01, error_rate=0.01, invalid_rate=0.01,
|
|
13
|
+
fail_first_attempt_only=True)
|
|
14
|
+
if name == "openai":
|
|
15
|
+
from .openai_batch import OpenAIBatch
|
|
16
|
+
return OpenAIBatch()
|
|
17
|
+
if name == "anthropic":
|
|
18
|
+
from .anthropic_batch import AnthropicBatch
|
|
19
|
+
return AnthropicBatch(thinking=thinking)
|
|
20
|
+
if name == "vllm":
|
|
21
|
+
from .vllm_offline import VLLMOffline
|
|
22
|
+
if not model:
|
|
23
|
+
raise ValueError("vllm needs a model name")
|
|
24
|
+
return VLLMOffline(model)
|
|
25
|
+
if name == "gemini":
|
|
26
|
+
from .gemini_batch import GeminiBatch
|
|
27
|
+
return GeminiBatch()
|
|
28
|
+
if name == "bedrock":
|
|
29
|
+
raise NotImplementedError(
|
|
30
|
+
"the Bedrock adapter is next on the list; see docs/ROADMAP.md. Contributions welcome."
|
|
31
|
+
)
|
|
32
|
+
raise ValueError(f"unknown provider: {name}")
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
__all__ = ["BatchProvider", "FakeProvider", "get_provider"]
|
|
@@ -0,0 +1,95 @@
|
|
|
1
|
+
"""Anthropic Message Batches API.
|
|
2
|
+
|
|
3
|
+
Flow: create a batch with up to 100,000 requests in one call, poll until
|
|
4
|
+
processing_status is "ended", then stream results. Each result is
|
|
5
|
+
succeeded, errored, canceled or expired.
|
|
6
|
+
|
|
7
|
+
When a JSON Schema is given, it is sent as a structured output format
|
|
8
|
+
(`output_config.format`), so the reply is constrained to the schema, the
|
|
9
|
+
same as the OpenAI and Gemini adapters. Pass `native_schema=False` to
|
|
10
|
+
describe the schema in the system prompt instead. `thinking` ("disabled" or
|
|
11
|
+
"adaptive") is sent when given; left out, the model uses its default, and on
|
|
12
|
+
Claude Haiku 5.5 that default is adaptive thinking, which counts against
|
|
13
|
+
max_tokens. `temperature` is not sent:
|
|
14
|
+
current Claude models reject non-default sampling settings, and anthropic
|
|
15
|
+
SDK 1.x no longer takes the parameter.
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
from __future__ import annotations
|
|
19
|
+
|
|
20
|
+
import json
|
|
21
|
+
from typing import Iterator
|
|
22
|
+
|
|
23
|
+
from ..models import BatchState, BatchStatus, PreparedRequest, RawResult
|
|
24
|
+
from .base import BatchProvider
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
class AnthropicBatch(BatchProvider):
|
|
28
|
+
name = "anthropic"
|
|
29
|
+
max_requests_per_batch = 100_000
|
|
30
|
+
max_bytes_per_batch = 256 * 1024 * 1024
|
|
31
|
+
|
|
32
|
+
def __init__(self, client=None, *, native_schema: bool = True, thinking: str | None = None):
|
|
33
|
+
if client is None:
|
|
34
|
+
import anthropic # imported here so the package works without it
|
|
35
|
+
client = anthropic.Anthropic()
|
|
36
|
+
self.client = client
|
|
37
|
+
self.native_schema = native_schema
|
|
38
|
+
self.thinking = thinking
|
|
39
|
+
|
|
40
|
+
def to_line(self, req: PreparedRequest) -> dict:
|
|
41
|
+
params = {
|
|
42
|
+
"model": req.model,
|
|
43
|
+
"max_tokens": req.request.max_tokens,
|
|
44
|
+
"messages": req.request.messages,
|
|
45
|
+
}
|
|
46
|
+
if self.thinking:
|
|
47
|
+
params["thinking"] = {"type": self.thinking}
|
|
48
|
+
system = req.request.system or ""
|
|
49
|
+
if req.response_schema is not None and self.native_schema:
|
|
50
|
+
params["output_config"] = {"format": {"type": "json_schema", "schema": req.response_schema}}
|
|
51
|
+
elif req.response_schema is not None:
|
|
52
|
+
# Ask for JSON in the prompt; the runner checks every answer against the schema.
|
|
53
|
+
system = (system + "\n\nReply with only a JSON object that matches this JSON Schema:\n"
|
|
54
|
+
+ json.dumps(req.response_schema)).strip()
|
|
55
|
+
if system:
|
|
56
|
+
params["system"] = system
|
|
57
|
+
return {"custom_id": req.custom_id, "params": params}
|
|
58
|
+
|
|
59
|
+
def submit(self, batch: list[PreparedRequest]) -> str:
|
|
60
|
+
created = self.client.messages.batches.create(requests=[self.to_line(r) for r in batch])
|
|
61
|
+
return created.id
|
|
62
|
+
|
|
63
|
+
def is_over_capacity(self, exc: Exception) -> bool:
|
|
64
|
+
return getattr(exc, "status_code", None) == 429
|
|
65
|
+
|
|
66
|
+
def status(self, batch_id: str) -> BatchStatus:
|
|
67
|
+
b = self.client.messages.batches.retrieve(batch_id)
|
|
68
|
+
c = b.request_counts
|
|
69
|
+
if b.processing_status == "ended":
|
|
70
|
+
state = BatchState.DONE
|
|
71
|
+
elif b.processing_status == "canceling":
|
|
72
|
+
state = BatchState.FAILED
|
|
73
|
+
else:
|
|
74
|
+
state = BatchState.RUNNING
|
|
75
|
+
total = c.processing + c.succeeded + c.errored + c.canceled + c.expired
|
|
76
|
+
return BatchStatus(state, total=total, succeeded=c.succeeded,
|
|
77
|
+
failed=c.errored + c.canceled + c.expired, detail=b.processing_status)
|
|
78
|
+
|
|
79
|
+
def results(self, batch_id: str) -> Iterator[RawResult]:
|
|
80
|
+
for entry in self.client.messages.batches.results(batch_id):
|
|
81
|
+
result = entry.result
|
|
82
|
+
if result.type == "succeeded":
|
|
83
|
+
msg = result.message
|
|
84
|
+
if getattr(msg, "stop_reason", None) == "refusal":
|
|
85
|
+
yield RawResult(entry.custom_id, ok=False, error="refusal")
|
|
86
|
+
continue
|
|
87
|
+
text = "".join(block.text for block in msg.content if block.type == "text")
|
|
88
|
+
truncated = getattr(msg, "stop_reason", None) == "max_tokens"
|
|
89
|
+
yield RawResult(entry.custom_id, ok=True, text=text,
|
|
90
|
+
error="truncated (max_tokens)" if truncated else None,
|
|
91
|
+
input_tokens=msg.usage.input_tokens, output_tokens=msg.usage.output_tokens)
|
|
92
|
+
elif result.type == "errored":
|
|
93
|
+
yield RawResult(entry.custom_id, ok=False, error=str(getattr(result, "error", "errored")))
|
|
94
|
+
else:
|
|
95
|
+
yield RawResult(entry.custom_id, ok=False, error=result.type)
|
|
@@ -0,0 +1,60 @@
|
|
|
1
|
+
"""The interface every batch provider implements.
|
|
2
|
+
|
|
3
|
+
A provider only has to do four things: say what its limits are, submit a
|
|
4
|
+
list of requests, report a batch's status, and hand back whatever lines it
|
|
5
|
+
has. Everything else (retries, checks, ordering, checkpoints) lives in the
|
|
6
|
+
runner, so every provider gets the same treatment and the comparison is fair.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import json
|
|
12
|
+
from abc import ABC, abstractmethod
|
|
13
|
+
from typing import Iterator
|
|
14
|
+
|
|
15
|
+
from ..models import BatchStatus, PreparedRequest, RawResult
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class BatchProvider(ABC):
|
|
19
|
+
name: str = "base"
|
|
20
|
+
max_requests_per_batch: int = 10_000
|
|
21
|
+
max_bytes_per_batch: int = 100 * 1024 * 1024
|
|
22
|
+
|
|
23
|
+
def size_of(self, req: PreparedRequest) -> int:
|
|
24
|
+
"""Rough size of one request line in the upload file."""
|
|
25
|
+
return len(json.dumps(self.to_line(req)).encode("utf-8")) + 1
|
|
26
|
+
|
|
27
|
+
def to_line(self, req: PreparedRequest) -> dict:
|
|
28
|
+
"""The provider-specific JSON for one request. Used for sizing and upload."""
|
|
29
|
+
return {"custom_id": req.custom_id, "messages": req.request.messages}
|
|
30
|
+
|
|
31
|
+
def is_over_capacity(self, exc: Exception) -> bool:
|
|
32
|
+
"""True if `submit` failed because the account's queue or quota is full.
|
|
33
|
+
|
|
34
|
+
The runner then sends smaller batches instead of stopping. Providers that
|
|
35
|
+
can tell this apart from other errors override it.
|
|
36
|
+
"""
|
|
37
|
+
return False
|
|
38
|
+
|
|
39
|
+
def is_transient(self, exc: Exception) -> bool:
|
|
40
|
+
"""True if `submit` failed for a reason worth retrying as-is: a dropped
|
|
41
|
+
connection, a timeout or a server error. The SDKs retry a couple of times
|
|
42
|
+
on their own; a long job needs more patience than that."""
|
|
43
|
+
names = {c.__name__ for c in type(exc).__mro__}
|
|
44
|
+
if names & {"APIConnectionError", "APITimeoutError", "ConnectionError", "TimeoutError",
|
|
45
|
+
"ServerError", "InternalServerError", "ServiceUnavailableError"}:
|
|
46
|
+
return True
|
|
47
|
+
status = getattr(exc, "status_code", None) or getattr(exc, "code", None)
|
|
48
|
+
return isinstance(status, int) and status >= 500
|
|
49
|
+
|
|
50
|
+
@abstractmethod
|
|
51
|
+
def submit(self, batch: list[PreparedRequest]) -> str:
|
|
52
|
+
"""Send a batch. Return the provider's batch ID."""
|
|
53
|
+
|
|
54
|
+
@abstractmethod
|
|
55
|
+
def status(self, batch_id: str) -> BatchStatus:
|
|
56
|
+
"""Report where a batch is."""
|
|
57
|
+
|
|
58
|
+
@abstractmethod
|
|
59
|
+
def results(self, batch_id: str) -> Iterator[RawResult]:
|
|
60
|
+
"""Yield every line the provider returned, successes and errors alike."""
|