Text Classification
Transformers
Safetensors
English
qwen3_5_text
text-generation
decision-model
typed-decisions
one-pass
option-probabilities
Instructions to use thegovind/blink-4b with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use thegovind/blink-4b with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="thegovind/blink-4b")# Load model directly from transformers import AutoTokenizer, AutoModelForCausalLM tokenizer = AutoTokenizer.from_pretrained("thegovind/blink-4b") model = AutoModelForCausalLM.from_pretrained("thegovind/blink-4b", device_map="auto") - Notebooks
- Google Colab
- Kaggle
v1.1: opt-in cross-request batching in serve.py (--batch-window-ms, off by default); weights unchanged
Browse files
README.md
CHANGED
|
@@ -212,6 +212,21 @@ docker build -t blink-4b . && docker run --rm --gpus all -p 127.0.0.1:8000:8000
|
|
| 212 |
- Requests run one at a time. Questions are batched; each batch takes one forward pass (large requests
|
| 213 |
can take more than one). Serving the downloaded folder or Docker image enables Hugging Face offline
|
| 214 |
mode before model loading (`hub_offline: true`). The server doesn't otherwise restrict network access.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 215 |
- blink-4b weights are 8.4 GB in bf16. Long prompts need more memory.
|
| 216 |
|
| 217 |
</details>
|
|
|
|
| 212 |
- Requests run one at a time. Questions are batched; each batch takes one forward pass (large requests
|
| 213 |
can take more than one). Serving the downloaded folder or Docker image enables Hugging Face offline
|
| 214 |
mode before model loading (`hub_offline: true`). The server doesn't otherwise restrict network access.
|
| 215 |
+
- Set `--batch-window-ms 5` to turn on cross-request batching with a 5 ms collection window, up to
|
| 216 |
+
`--max-batch-requests` requests at a time, which defaults to 16. The window defaults to 0, so requests still
|
| 217 |
+
run one at a time. `--max-queued-requests` lets up to 64 requests wait for a batch by default. Excess requests
|
| 218 |
+
get HTTP 503 with `Retry-After`, so clients should retry. On a 1,000-request Decision Index sample over HTTP,
|
| 219 |
+
throughput rose about 20% with 4 concurrent clients and 24% with 16. One client saw no gain. Offline runs on
|
| 220 |
+
long documents showed no meaningful gain. On the Decision Index sample, a set of long workflow documents, and
|
| 221 |
+
the public TypeSafe cases, with each document sent as a JSON object, batched answers passed the same
|
| 222 |
+
numerical-parity checks against an FP32 reference as one-at-a-time answers, covering argmax agreement and
|
| 223 |
+
probability differences. A few near-tied answers can still flip. `v1.1` changes code only, leaving `v1.0`
|
| 224 |
+
weights unchanged. Update the two code files in an existing `v1.0` download, then restart with the flag:
|
| 225 |
+
|
| 226 |
+
```sh
|
| 227 |
+
hf download thegovind/blink-4b serve.py blink.py --revision v1.1 --local-dir blink-4b
|
| 228 |
+
python blink-4b/serve.py --model ./blink-4b --port 8000 --batch-window-ms 5
|
| 229 |
+
```
|
| 230 |
- blink-4b weights are 8.4 GB in bf16. Long prompts need more memory.
|
| 231 |
|
| 232 |
</details>
|
blink.py
CHANGED
|
@@ -56,6 +56,29 @@ MAX_OPTIONS = 255
|
|
| 56 |
MAX_QUESTIONS = 512
|
| 57 |
MAX_INPUT_TOKENS = int(os.environ.get("BLINK_MAX_INPUT_TOKENS", "131072")) # as evaluated
|
| 58 |
TOKEN_BUDGET = int(os.environ.get("BLINK_TOKEN_BUDGET", "32768")) # padded tokens per forward
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 59 |
GPU_DURATION = int(os.environ.get("BLINK_GPU_DURATION", "60"))
|
| 60 |
|
| 61 |
SYSTEM = (
|
|
@@ -427,21 +450,42 @@ class TorchEngine:
|
|
| 427 |
def logits(self, state, questions: dict, bias: dict | None = None) -> tuple[dict[str, list[float]], int]:
|
| 428 |
del bias # demo-only shaping; the trained model reads the evidence instead
|
| 429 |
work = self.render(state, questions)
|
| 430 |
-
|
|
|
|
|
|
|
| 431 |
return (
|
| 432 |
{w["qkey"]: r for w, r in zip(work, rows)},
|
| 433 |
-
sum(len(
|
| 434 |
)
|
| 435 |
|
| 436 |
-
|
| 437 |
-
|
| 438 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 439 |
A sequence longer than the budget runs alone."""
|
| 440 |
order = sorted(range(len(lengths)), key=lambda i: lengths[i])
|
| 441 |
batch, longest = [], 0
|
| 442 |
for i in order:
|
| 443 |
grown = max(longest, lengths[i])
|
| 444 |
-
if batch and grown * (len(batch) + 1) > budget:
|
| 445 |
yield batch
|
| 446 |
batch, grown = [], lengths[i]
|
| 447 |
batch.append(i)
|
|
@@ -450,31 +494,128 @@ def _batches(lengths: list[int], budget: int):
|
|
| 450 |
yield batch
|
| 451 |
|
| 452 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 453 |
@_gpu(GPU_DURATION)
|
| 454 |
def _forward(key: int, seqs: list[list[int]], cands: list[list[int]]):
|
| 455 |
"""Right-padded maskless prefill; every layer is causal, so padding cannot reach the
|
| 456 |
-
last real position. Label rows of lm_head are applied in FP32.
|
|
|
|
| 457 |
import torch
|
| 458 |
|
| 459 |
engine = _LIVE[key]
|
| 460 |
model = engine.model
|
| 461 |
device = next(model.parameters()).device
|
| 462 |
head = model.lm_head.weight
|
|
|
|
| 463 |
out: list = [None] * len(seqs)
|
| 464 |
t0 = time.perf_counter()
|
| 465 |
with torch.no_grad():
|
| 466 |
-
|
| 467 |
-
|
| 468 |
-
|
| 469 |
-
|
| 470 |
-
|
| 471 |
-
|
| 472 |
-
|
| 473 |
-
|
| 474 |
-
h = h[torch.arange(len(b), device=device), last].float()
|
| 475 |
-
for r, i in enumerate(b):
|
| 476 |
-
w = head[torch.tensor(cands[i], device=device)].float()
|
| 477 |
-
out[i] = (w @ h[r]).tolist()
|
| 478 |
return out, round((time.perf_counter() - t0) * 1000, 1)
|
| 479 |
|
| 480 |
|
|
@@ -620,4 +761,144 @@ def decide(state, questions: dict, temperature: float | None = None, bias: dict
|
|
| 620 |
model_ms = getattr(getattr(eng, "live", eng), "last_model_ms", None)
|
| 621 |
if model_ms is not None:
|
| 622 |
meta["model_ms"] = model_ms # forward passes only; excludes any wait for a device
|
|
|
|
|
|
|
|
|
|
| 623 |
return {"answers": answers, "meta": meta}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 56 |
MAX_QUESTIONS = 512
|
| 57 |
MAX_INPUT_TOKENS = int(os.environ.get("BLINK_MAX_INPUT_TOKENS", "131072")) # as evaluated
|
| 58 |
TOKEN_BUDGET = int(os.environ.get("BLINK_TOKEN_BUDGET", "32768")) # padded tokens per forward
|
| 59 |
+
# Shared-prefix prefill (opt-in, BLINK_PREFIX_CACHE=1): the questions of one request share the system prompt and the
|
| 60 |
+
# evidence, so that prefix is encoded once and each question runs only its own tail against a copy of the cached
|
| 61 |
+
# prefix. Same tokens, same positions; only floating-point order differs. Measured 4.8-6.9x less model time on long
|
| 62 |
+
# multi-question documents, but on ~18k-token documents blink-4b's bf16 answers drift slightly further from an FP32
|
| 63 |
+
# reference than the plain path's do, so it is off by default (experiments/t5/PREREG.md, gate A').
|
| 64 |
+
PREFIX_CACHE = os.environ.get("BLINK_PREFIX_CACHE", "0").lower() in ("1", "true", "on", "yes")
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
def _prefix_setting(name: str, default: int) -> int:
|
| 68 |
+
"""A shared-prefix setting, read only when BLINK_PREFIX_CACHE is on: a stray value can't break the default path."""
|
| 69 |
+
if not PREFIX_CACHE:
|
| 70 |
+
return default
|
| 71 |
+
value = int(os.environ.get(name, str(default)))
|
| 72 |
+
if value < 1:
|
| 73 |
+
raise ValueError(f"{name} must be at least 1")
|
| 74 |
+
return value
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
PREFIX_MIN_TOKENS = _prefix_setting("BLINK_PREFIX_MIN_TOKENS", 256) # below this the plain path is as fast
|
| 78 |
+
PREFIX_KV_TOKENS = _prefix_setting("BLINK_PREFIX_KV_TOKENS", 131072) # prefix copies x length held at once
|
| 79 |
+
# Readout implementation: "rows" (release: per-question gather, upcast and GEMV) or "batched" (one gather, one FP32
|
| 80 |
+
# matmul and one host transfer per forward; same arithmetic, different FP32 accumulation order). Opt-in under E1.
|
| 81 |
+
READOUT = os.environ.get("BLINK_READOUT", "rows").strip().lower()
|
| 82 |
GPU_DURATION = int(os.environ.get("BLINK_GPU_DURATION", "60"))
|
| 83 |
|
| 84 |
SYSTEM = (
|
|
|
|
| 450 |
def logits(self, state, questions: dict, bias: dict | None = None) -> tuple[dict[str, list[float]], int]:
|
| 451 |
del bias # demo-only shaping; the trained model reads the evidence instead
|
| 452 |
work = self.render(state, questions)
|
| 453 |
+
seqs = [w["ids"] for w in work]
|
| 454 |
+
rows, self.last_model_ms = _forward(self.key, seqs, [w["cand"] for w in work])
|
| 455 |
+
self.last_prefill_tokens = _prefill_tokens(seqs, getattr(self, "prefix_cache", PREFIX_CACHE))
|
| 456 |
return (
|
| 457 |
{w["qkey"]: r for w, r in zip(work, rows)},
|
| 458 |
+
sum(len(s) for s in seqs),
|
| 459 |
)
|
| 460 |
|
| 461 |
+
def logits_many(self, requests: list) -> list[tuple[dict[str, list[float]], int]]:
|
| 462 |
+
"""Several already-validated requests in one GPU call: every question of every request is packed into the
|
| 463 |
+
same padded forwards, then split back per request. Cross-request batching for a queueing server (E1 arm H1);
|
| 464 |
+
each request's rows come only from its own prompts."""
|
| 465 |
+
return self.logits_rendered([self.render(state, questions) for state, questions in requests])
|
| 466 |
+
|
| 467 |
+
def logits_rendered(self, works: list) -> list[tuple[dict[str, list[float]], int]]:
|
| 468 |
+
"""logits_many for requests already rendered (one render() list per request), so a caller can render each
|
| 469 |
+
request on its own and keep one request's rendering error away from the others."""
|
| 470 |
+
seqs = [w["ids"] for work in works for w in work]
|
| 471 |
+
cands = [w["cand"] for work in works for w in work]
|
| 472 |
+
rows, self.last_model_ms = _forward(self.key, seqs, cands) if seqs else ([], 0.0)
|
| 473 |
+
out, k = [], 0
|
| 474 |
+
for work in works:
|
| 475 |
+
part = rows[k:k + len(work)]
|
| 476 |
+
k += len(work)
|
| 477 |
+
out.append(({w["qkey"]: r for w, r in zip(work, part)}, sum(len(w["ids"]) for w in work)))
|
| 478 |
+
return out
|
| 479 |
+
|
| 480 |
+
|
| 481 |
+
def _batches(lengths: list[int], budget: int, max_rows: int | None = None):
|
| 482 |
+
"""Shortest first; a batch closes when (longest x count) would pass the padded-token budget, or at max_rows.
|
| 483 |
A sequence longer than the budget runs alone."""
|
| 484 |
order = sorted(range(len(lengths)), key=lambda i: lengths[i])
|
| 485 |
batch, longest = [], 0
|
| 486 |
for i in order:
|
| 487 |
grown = max(longest, lengths[i])
|
| 488 |
+
if batch and (grown * (len(batch) + 1) > budget or (max_rows and len(batch) >= max_rows)):
|
| 489 |
yield batch
|
| 490 |
batch, grown = [], lengths[i]
|
| 491 |
batch.append(i)
|
|
|
|
| 494 |
yield batch
|
| 495 |
|
| 496 |
|
| 497 |
+
def _shared_prefix_len(seqs: list[list[int]]) -> int:
|
| 498 |
+
"""Length of the token prefix common to every sequence, leaving each at least one token of its own
|
| 499 |
+
(the answer position must be computed in the per-question pass)."""
|
| 500 |
+
if len(seqs) < 2:
|
| 501 |
+
return 0
|
| 502 |
+
lo, hi = min(seqs), max(seqs) # the common prefix of all = the common prefix of the lexicographic extremes
|
| 503 |
+
n = min(len(lo), len(hi))
|
| 504 |
+
p = 0
|
| 505 |
+
while p < n and lo[p] == hi[p]:
|
| 506 |
+
p += 1
|
| 507 |
+
return min(p, min(len(s) for s in seqs) - 1)
|
| 508 |
+
|
| 509 |
+
|
| 510 |
+
def _prefill_tokens(seqs: list[list[int]], prefix_cache: bool) -> int:
|
| 511 |
+
"""Tokens the model actually computes for these prompts: the shared prefix once when it is reused."""
|
| 512 |
+
total = sum(len(s) for s in seqs)
|
| 513 |
+
shared = _shared_prefix_len(seqs) if prefix_cache else 0
|
| 514 |
+
return total - (len(seqs) - 1) * shared if shared >= PREFIX_MIN_TOKENS else total
|
| 515 |
+
|
| 516 |
+
|
| 517 |
+
def _expand_cache(cache, repeats: int):
|
| 518 |
+
"""A deep copy of a batch-1 cache repeated `repeats` times along the batch axis, for every tensor it holds:
|
| 519 |
+
full-attention keys/values and the linear-attention conv and recurrent states (the library's own
|
| 520 |
+
batch_repeat_interleave only covers keys/values)."""
|
| 521 |
+
import copy
|
| 522 |
+
|
| 523 |
+
import torch
|
| 524 |
+
|
| 525 |
+
cache = copy.deepcopy(cache)
|
| 526 |
+
if repeats == 1:
|
| 527 |
+
return cache
|
| 528 |
+
|
| 529 |
+
def rep(x):
|
| 530 |
+
return x.repeat_interleave(repeats, dim=0) if torch.is_tensor(x) and x.numel() else x
|
| 531 |
+
|
| 532 |
+
for layer in cache.layers:
|
| 533 |
+
for name in ("keys", "values"):
|
| 534 |
+
if torch.is_tensor(getattr(layer, name, None)):
|
| 535 |
+
setattr(layer, name, rep(getattr(layer, name)))
|
| 536 |
+
for name in ("conv_states", "recurrent_states"):
|
| 537 |
+
held = getattr(layer, name, None)
|
| 538 |
+
if isinstance(held, dict):
|
| 539 |
+
setattr(layer, name, {k: rep(v) for k, v in held.items()})
|
| 540 |
+
elif isinstance(held, (list, tuple)):
|
| 541 |
+
setattr(layer, name, type(held)(rep(v) for v in held))
|
| 542 |
+
return cache
|
| 543 |
+
|
| 544 |
+
|
| 545 |
+
def _read(model, head, h, lengths, cands, rows, out, device):
|
| 546 |
+
import torch
|
| 547 |
+
|
| 548 |
+
if READOUT == "batched":
|
| 549 |
+
return _read_batched(head, h, lengths, cands, rows, out, device)
|
| 550 |
+
last = torch.tensor([n - 1 for n in lengths], device=device)
|
| 551 |
+
h = h[torch.arange(len(rows), device=device), last].float()
|
| 552 |
+
for r, i in enumerate(rows):
|
| 553 |
+
w = head[torch.tensor(cands[i], device=device)].float()
|
| 554 |
+
out[i] = (w @ h[r]).tolist()
|
| 555 |
+
|
| 556 |
+
|
| 557 |
+
def _read_batched(head, h, lengths, cands, rows, out, device):
|
| 558 |
+
"""The same FP32 readout as _read, for a whole forward at once: the batch's distinct label rows are gathered
|
| 559 |
+
and upcast once, one FP32 matmul scores every question against them, and one host transfer brings them back.
|
| 560 |
+
Only the order of FP32 accumulation differs from the per-question path (BLINK_READOUT=batched; opt-in)."""
|
| 561 |
+
import torch
|
| 562 |
+
|
| 563 |
+
last = torch.tensor([n - 1 for n in lengths], device=device)
|
| 564 |
+
hs = h[torch.arange(len(rows), device=device), last].float()
|
| 565 |
+
uniq = sorted({t for i in rows for t in cands[i]})
|
| 566 |
+
col = {t: j for j, t in enumerate(uniq)}
|
| 567 |
+
z = (hs @ head[torch.tensor(uniq, device=device)].float().T).tolist()
|
| 568 |
+
for r, i in enumerate(rows):
|
| 569 |
+
out[i] = [z[r][col[t]] for t in cands[i]]
|
| 570 |
+
|
| 571 |
+
|
| 572 |
+
def _padded(seqs, rows, pad_id):
|
| 573 |
+
import torch
|
| 574 |
+
|
| 575 |
+
L = max(len(seqs[i]) for i in rows)
|
| 576 |
+
ids = torch.full((len(rows), L), pad_id, dtype=torch.long)
|
| 577 |
+
for r, i in enumerate(rows):
|
| 578 |
+
ids[r, : len(seqs[i])] = torch.tensor(seqs[i], dtype=torch.long)
|
| 579 |
+
return ids
|
| 580 |
+
|
| 581 |
+
|
| 582 |
+
def _forward_shared(model, head, seqs, cands, prefix_len: int, pad_id: int, budget: int, out, device) -> None:
|
| 583 |
+
"""Encode seqs[0][:prefix_len] once, then run every tail against an expanded copy of that cache."""
|
| 584 |
+
import torch
|
| 585 |
+
|
| 586 |
+
prefix = torch.tensor([seqs[0][:prefix_len]], dtype=torch.long, device=device)
|
| 587 |
+
cache = model.model(input_ids=prefix, use_cache=True).past_key_values
|
| 588 |
+
tails = [s[prefix_len:] for s in seqs]
|
| 589 |
+
rows_cap = max(1, PREFIX_KV_TOKENS // (prefix_len + max(len(t) for t in tails)))
|
| 590 |
+
for b in _batches([len(t) for t in tails], budget, max_rows=rows_cap):
|
| 591 |
+
ids = _padded(tails, b, pad_id).to(device)
|
| 592 |
+
h = model.model(input_ids=ids, past_key_values=_expand_cache(cache, len(b)), use_cache=True).last_hidden_state
|
| 593 |
+
_read(model, head, h, [len(tails[i]) for i in b], cands, b, out, device)
|
| 594 |
+
|
| 595 |
+
|
| 596 |
@_gpu(GPU_DURATION)
|
| 597 |
def _forward(key: int, seqs: list[list[int]], cands: list[list[int]]):
|
| 598 |
"""Right-padded maskless prefill; every layer is causal, so padding cannot reach the
|
| 599 |
+
last real position. Label rows of lm_head are applied in FP32. With several questions over
|
| 600 |
+
a long shared prefix, the prefix is encoded once (see PREFIX_CACHE)."""
|
| 601 |
import torch
|
| 602 |
|
| 603 |
engine = _LIVE[key]
|
| 604 |
model = engine.model
|
| 605 |
device = next(model.parameters()).device
|
| 606 |
head = model.lm_head.weight
|
| 607 |
+
budget = getattr(engine, "token_budget", TOKEN_BUDGET)
|
| 608 |
out: list = [None] * len(seqs)
|
| 609 |
t0 = time.perf_counter()
|
| 610 |
with torch.no_grad():
|
| 611 |
+
shared = _shared_prefix_len(seqs) if getattr(engine, "prefix_cache", PREFIX_CACHE) else 0
|
| 612 |
+
if shared > 0 and shared >= PREFIX_MIN_TOKENS:
|
| 613 |
+
_forward_shared(model, head, seqs, cands, shared, engine.pad_id, budget, out, device)
|
| 614 |
+
else:
|
| 615 |
+
for b in _batches([len(s) for s in seqs], budget):
|
| 616 |
+
ids = _padded(seqs, b, engine.pad_id).to(device)
|
| 617 |
+
h = model.model(input_ids=ids, use_cache=False).last_hidden_state
|
| 618 |
+
_read(model, head, h, [len(seqs[i]) for i in b], cands, b, out, device)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 619 |
return out, round((time.perf_counter() - t0) * 1000, 1)
|
| 620 |
|
| 621 |
|
|
|
|
| 761 |
model_ms = getattr(getattr(eng, "live", eng), "last_model_ms", None)
|
| 762 |
if model_ms is not None:
|
| 763 |
meta["model_ms"] = model_ms # forward passes only; excludes any wait for a device
|
| 764 |
+
prefill = getattr(getattr(eng, "live", eng), "last_prefill_tokens", None)
|
| 765 |
+
if prefill is not None:
|
| 766 |
+
meta["prefill_tokens"] = prefill # tokens computed; input_tokens counts every question's full prompt
|
| 767 |
return {"answers": answers, "meta": meta}
|
| 768 |
+
|
| 769 |
+
|
| 770 |
+
def decide_many(requests: list, temperature: float | None = None, model: str | None = None) -> list:
|
| 771 |
+
"""Several {state, questions} requests in one GPU call (E1 arm H1). Returns one item per request, in order: the
|
| 772 |
+
same {answers, meta} dict decide() returns, or the exception that request raised (a BlinkError when blink can't
|
| 773 |
+
answer it). Each request is validated and rendered on its own, so one request's failure never reaches another;
|
| 774 |
+
if the packed forward itself fails, every request is answered alone instead. Engines without logits_many
|
| 775 |
+
answer one request at a time."""
|
| 776 |
+
eng = engine(model)
|
| 777 |
+
T = float(temperature if temperature is not None else eng.temperature)
|
| 778 |
+
if T <= 0:
|
| 779 |
+
raise BlinkError("temperature must be positive")
|
| 780 |
+
live = getattr(eng, "live", eng)
|
| 781 |
+
packed = hasattr(live, "logits_rendered")
|
| 782 |
+
out: list = [None] * len(requests)
|
| 783 |
+
todo, works = [], []
|
| 784 |
+
for i, (state, questions) in enumerate(requests):
|
| 785 |
+
try:
|
| 786 |
+
validate(questions)
|
| 787 |
+
if packed:
|
| 788 |
+
works.append(live.render(state, questions))
|
| 789 |
+
todo.append(i)
|
| 790 |
+
except Exception as exc: # noqa: BLE001 - this request's own error (422 for a BlinkError, else 500)
|
| 791 |
+
out[i] = exc
|
| 792 |
+
t0 = time.perf_counter()
|
| 793 |
+
got: dict = {}
|
| 794 |
+
if packed and todo:
|
| 795 |
+
try:
|
| 796 |
+
got = dict(zip(todo, live.logits_rendered(works)))
|
| 797 |
+
except Exception: # noqa: BLE001 - e.g. out of memory on the packed batch: answer each request alone below
|
| 798 |
+
got = {}
|
| 799 |
+
n_packed = len(got)
|
| 800 |
+
for i in todo:
|
| 801 |
+
if i not in got:
|
| 802 |
+
try:
|
| 803 |
+
got[i] = (live if packed else eng).logits(*requests[i])
|
| 804 |
+
except Exception as exc: # noqa: BLE001
|
| 805 |
+
got[i] = exc
|
| 806 |
+
latency_ms = round((time.perf_counter() - t0) * 1000, 1)
|
| 807 |
+
for i in todo:
|
| 808 |
+
res = got[i]
|
| 809 |
+
if isinstance(res, Exception):
|
| 810 |
+
out[i] = res
|
| 811 |
+
continue
|
| 812 |
+
try:
|
| 813 |
+
raw, n_tokens = res
|
| 814 |
+
questions = requests[i][1]
|
| 815 |
+
answers = {q: answer_for(spec, [k for k, _ in question_options(spec)], softmax(raw[q], T))
|
| 816 |
+
for q, spec in questions.items()}
|
| 817 |
+
out[i] = {"answers": answers, "meta": {"model": getattr(eng, "model_id", MODEL_ID), "engine": eng.name,
|
| 818 |
+
"temperature": T, "input_tokens": n_tokens, "generated_tokens": 0,
|
| 819 |
+
"latency_ms": latency_ms, "batched_requests": n_packed or 1}}
|
| 820 |
+
except Exception as exc: # noqa: BLE001
|
| 821 |
+
out[i] = exc
|
| 822 |
+
return out
|
| 823 |
+
|
| 824 |
+
|
| 825 |
+
class BlinkBusy(RuntimeError):
|
| 826 |
+
"""The batching queue is full: the server is over capacity (HTTP 503; retry shortly)."""
|
| 827 |
+
|
| 828 |
+
|
| 829 |
+
class Batcher:
|
| 830 |
+
"""Cross-request batching for a serving process (E1 arm H1). One worker thread owns the engine: it takes the
|
| 831 |
+
oldest waiting request, waits up to `window_s` for more (at most `max_requests` in all) and decides them together
|
| 832 |
+
with decide_many. At most `max_queued` requests wait; past that, submit() raises BlinkBusy at once. submit()
|
| 833 |
+
blocks the calling thread until its own result is ready and returns it, or raises that request's own error."""
|
| 834 |
+
|
| 835 |
+
MAX_WINDOW_S = 1.0
|
| 836 |
+
MAX_REQUESTS = 64 # the largest group E1 measured
|
| 837 |
+
|
| 838 |
+
def __init__(self, window_s: float, max_requests: int = 16, model: str | None = None, max_queued: int = 64):
|
| 839 |
+
import math
|
| 840 |
+
import queue
|
| 841 |
+
import threading
|
| 842 |
+
|
| 843 |
+
window_s, max_requests, max_queued = float(window_s), int(max_requests), int(max_queued)
|
| 844 |
+
if not (math.isfinite(window_s) and 0 < window_s <= self.MAX_WINDOW_S):
|
| 845 |
+
raise ValueError(f"batch window must be more than 0 and at most {self.MAX_WINDOW_S:g} s")
|
| 846 |
+
if not 1 <= max_requests <= self.MAX_REQUESTS:
|
| 847 |
+
raise ValueError(f"max_requests must be 1-{self.MAX_REQUESTS}")
|
| 848 |
+
if max_queued < 1:
|
| 849 |
+
raise ValueError("max_queued must be at least 1")
|
| 850 |
+
self.window_s, self.max_requests, self.max_queued, self.model = window_s, max_requests, max_queued, model
|
| 851 |
+
self.q = queue.Queue(maxsize=max_queued)
|
| 852 |
+
self.worker = threading.Thread(target=self._loop, name="blink-batcher", daemon=True)
|
| 853 |
+
self.worker.start()
|
| 854 |
+
|
| 855 |
+
def alive(self) -> bool:
|
| 856 |
+
return self.worker.is_alive()
|
| 857 |
+
|
| 858 |
+
def submit(self, state, questions: dict) -> dict:
|
| 859 |
+
import concurrent.futures
|
| 860 |
+
import queue
|
| 861 |
+
|
| 862 |
+
validate(questions) # a malformed request fails at once and never takes a place in the queue
|
| 863 |
+
if not self.worker.is_alive():
|
| 864 |
+
raise RuntimeError("the batching worker has stopped")
|
| 865 |
+
fut = concurrent.futures.Future()
|
| 866 |
+
try:
|
| 867 |
+
self.q.put_nowait((state, questions, fut))
|
| 868 |
+
except queue.Full:
|
| 869 |
+
raise BlinkBusy(f"{self.max_queued} requests are already waiting; retry shortly") from None
|
| 870 |
+
while not fut.done():
|
| 871 |
+
concurrent.futures.wait([fut], timeout=1.0)
|
| 872 |
+
if not fut.done() and not self.worker.is_alive():
|
| 873 |
+
raise RuntimeError("the batching worker has stopped")
|
| 874 |
+
return fut.result()
|
| 875 |
+
|
| 876 |
+
def _loop(self) -> None:
|
| 877 |
+
import queue
|
| 878 |
+
|
| 879 |
+
while True:
|
| 880 |
+
batch = [self.q.get()]
|
| 881 |
+
results = None
|
| 882 |
+
try:
|
| 883 |
+
deadline = time.monotonic() + self.window_s
|
| 884 |
+
while len(batch) < self.max_requests:
|
| 885 |
+
left = deadline - time.monotonic()
|
| 886 |
+
if left <= 0:
|
| 887 |
+
break
|
| 888 |
+
try:
|
| 889 |
+
batch.append(self.q.get(timeout=left))
|
| 890 |
+
except queue.Empty:
|
| 891 |
+
break
|
| 892 |
+
results = decide_many([(s, q) for s, q, _ in batch], model=self.model)
|
| 893 |
+
except Exception as exc: # noqa: BLE001 - a failure outside any one request (e.g. the engine): each caller gets it
|
| 894 |
+
results = [exc] * len(batch)
|
| 895 |
+
finally:
|
| 896 |
+
if results is None or len(results) != len(batch):
|
| 897 |
+
results = [None] * len(batch)
|
| 898 |
+
for (_, _, fut), res in zip(batch, results):
|
| 899 |
+
if res is None:
|
| 900 |
+
fut.set_exception(RuntimeError("the batching worker lost this request"))
|
| 901 |
+
elif isinstance(res, BaseException):
|
| 902 |
+
fut.set_exception(res)
|
| 903 |
+
else:
|
| 904 |
+
fut.set_result(res)
|
serve.py
CHANGED
|
@@ -12,6 +12,10 @@ questions per request) gets HTTP 422 with the reason; nothing is truncated. With
|
|
| 12 |
the weights, every listed file is hashed before serving (weights_verified). Serving a local folder switches
|
| 13 |
the Hugging Face libraries to offline mode before any of them loads (hub_offline reports the setting the
|
| 14 |
libraries actually use); the server does not otherwise restrict the network.
|
|
|
|
|
|
|
|
|
|
|
|
|
| 15 |
"""
|
| 16 |
|
| 17 |
from __future__ import annotations
|
|
@@ -67,12 +71,44 @@ def versions() -> dict:
|
|
| 67 |
return out
|
| 68 |
|
| 69 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 70 |
def main() -> None:
|
| 71 |
ap = argparse.ArgumentParser(description="Serve blink over a Jev-compatible HTTP API.")
|
| 72 |
ap.add_argument("--model", default=os.environ.get("BLINK_MODEL", "thegovind/blink-4b"))
|
| 73 |
ap.add_argument("--revision", default=os.environ.get("BLINK_REVISION"))
|
| 74 |
ap.add_argument("--host", default="127.0.0.1")
|
| 75 |
ap.add_argument("--port", type=int, default=8000)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 76 |
a = ap.parse_args()
|
| 77 |
|
| 78 |
local = os.path.isdir(a.model)
|
|
@@ -84,6 +120,8 @@ def main() -> None:
|
|
| 84 |
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
| 85 |
import blink
|
| 86 |
|
|
|
|
|
|
|
| 87 |
if local:
|
| 88 |
root = a.model
|
| 89 |
else:
|
|
@@ -109,6 +147,17 @@ def main() -> None:
|
|
| 109 |
"hub_offline": hub_offline, "warmup": {"repeat_identical": repeat_identical}, "kernels": kernels,
|
| 110 |
"versions": ver}
|
| 111 |
lock = threading.Lock()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 112 |
|
| 113 |
class Handler(BaseHTTPRequestHandler):
|
| 114 |
protocol_version = "HTTP/1.1"
|
|
@@ -122,17 +171,19 @@ def main() -> None:
|
|
| 122 |
# on the client's delayed ACK, a flat ~40 ms on every request of a kept-alive connection
|
| 123 |
self.connection.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
|
| 124 |
|
| 125 |
-
def _send(self, code: int, obj: dict) -> None:
|
| 126 |
body = json.dumps(obj, ensure_ascii=False).encode("utf-8")
|
| 127 |
self.send_response(code)
|
| 128 |
self.send_header("Content-Type", "application/json")
|
| 129 |
self.send_header("Content-Length", str(len(body)))
|
|
|
|
|
|
|
| 130 |
self.end_headers()
|
| 131 |
self.wfile.write(body)
|
| 132 |
|
| 133 |
def do_GET(self):
|
| 134 |
if self.path.rstrip("/") in ("/healthz", "/health"):
|
| 135 |
-
return self._send(200,
|
| 136 |
return self._send(404, {"error": "not found"})
|
| 137 |
|
| 138 |
def do_POST(self):
|
|
@@ -146,10 +197,15 @@ def main() -> None:
|
|
| 146 |
if not isinstance(req, dict):
|
| 147 |
return self._send(400, {"error": "the body must be a JSON object"})
|
| 148 |
try:
|
| 149 |
-
|
| 150 |
-
out =
|
|
|
|
|
|
|
|
|
|
| 151 |
except blink.BlinkError as exc:
|
| 152 |
return self._send(422, {"error": str(exc)})
|
|
|
|
|
|
|
| 153 |
except Exception as exc: # noqa: BLE001 - report, keep serving
|
| 154 |
return self._send(500, {"error": f"{type(exc).__name__}: {exc}"})
|
| 155 |
return self._send(200, {
|
|
|
|
| 12 |
the weights, every listed file is hashed before serving (weights_verified). Serving a local folder switches
|
| 13 |
the Hugging Face libraries to offline mode before any of them loads (hub_offline reports the setting the
|
| 14 |
libraries actually use); the server does not otherwise restrict the network.
|
| 15 |
+
|
| 16 |
+
Opt-in cross-request batching (--batch-window-ms, default 0 = off): requests that arrive within the window are
|
| 17 |
+
decided together, up to --max-batch-requests; at most --max-queued-requests wait, and past that a request gets
|
| 18 |
+
HTTP 503 with Retry-After. Each request still gets its own answers or its own error.
|
| 19 |
"""
|
| 20 |
|
| 21 |
from __future__ import annotations
|
|
|
|
| 71 |
return out
|
| 72 |
|
| 73 |
|
| 74 |
+
def window_ms(value: str) -> float:
|
| 75 |
+
"""--batch-window-ms: 0 (off) or a finite window of at most 1000 ms."""
|
| 76 |
+
try:
|
| 77 |
+
ms = float(value)
|
| 78 |
+
except ValueError:
|
| 79 |
+
raise argparse.ArgumentTypeError(f"not a number: {value!r}") from None
|
| 80 |
+
if not 0 <= ms <= 1000: # also rejects nan and inf
|
| 81 |
+
raise argparse.ArgumentTypeError("must be 0 (off) to 1000 ms")
|
| 82 |
+
return ms
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
def int_range(lo: int, hi: int):
|
| 86 |
+
def parse(value: str) -> int:
|
| 87 |
+
try:
|
| 88 |
+
n = int(value)
|
| 89 |
+
except ValueError:
|
| 90 |
+
raise argparse.ArgumentTypeError(f"not an integer: {value!r}") from None
|
| 91 |
+
if not lo <= n <= hi:
|
| 92 |
+
raise argparse.ArgumentTypeError(f"must be {lo}-{hi}")
|
| 93 |
+
return n
|
| 94 |
+
|
| 95 |
+
return parse
|
| 96 |
+
|
| 97 |
+
|
| 98 |
def main() -> None:
|
| 99 |
ap = argparse.ArgumentParser(description="Serve blink over a Jev-compatible HTTP API.")
|
| 100 |
ap.add_argument("--model", default=os.environ.get("BLINK_MODEL", "thegovind/blink-4b"))
|
| 101 |
ap.add_argument("--revision", default=os.environ.get("BLINK_REVISION"))
|
| 102 |
ap.add_argument("--host", default="127.0.0.1")
|
| 103 |
ap.add_argument("--port", type=int, default=8000)
|
| 104 |
+
# a string default is parsed only when the flag is absent, so a stray BLINK_BATCH_WINDOW_MS can't block an explicit 0
|
| 105 |
+
ap.add_argument("--batch-window-ms", type=window_ms, default=os.environ.get("BLINK_BATCH_WINDOW_MS", "0"),
|
| 106 |
+
help="opt-in cross-request batching: requests arriving within this window are decided in one "
|
| 107 |
+
"GPU call (0 = one request at a time, the evaluated default; at most 1000)")
|
| 108 |
+
ap.add_argument("--max-batch-requests", type=int_range(1, 64), default=16,
|
| 109 |
+
help="most requests decided together (1-64)")
|
| 110 |
+
ap.add_argument("--max-queued-requests", type=int_range(1, 1024), default=64,
|
| 111 |
+
help="most requests waiting for a batch; past this a request gets HTTP 503 (1-1024)")
|
| 112 |
a = ap.parse_args()
|
| 113 |
|
| 114 |
local = os.path.isdir(a.model)
|
|
|
|
| 120 |
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
| 121 |
import blink
|
| 122 |
|
| 123 |
+
busy = getattr(blink, "BlinkBusy", ()) # an older blink.py has no batching (and no BlinkBusy)
|
| 124 |
+
|
| 125 |
if local:
|
| 126 |
root = a.model
|
| 127 |
else:
|
|
|
|
| 147 |
"hub_offline": hub_offline, "warmup": {"repeat_identical": repeat_identical}, "kernels": kernels,
|
| 148 |
"versions": ver}
|
| 149 |
lock = threading.Lock()
|
| 150 |
+
batcher = (blink.Batcher(a.batch_window_ms / 1000.0, a.max_batch_requests, max_queued=a.max_queued_requests)
|
| 151 |
+
if a.batch_window_ms > 0 else None)
|
| 152 |
+
health["batching"] = ({"window_ms": a.batch_window_ms, "max_requests": a.max_batch_requests,
|
| 153 |
+
"max_queued": a.max_queued_requests} if batcher else None)
|
| 154 |
+
|
| 155 |
+
def current_health() -> dict:
|
| 156 |
+
if batcher is None:
|
| 157 |
+
return health
|
| 158 |
+
alive = batcher.alive()
|
| 159 |
+
return {**health, "ok": health["ok"] and alive,
|
| 160 |
+
"batching": {**health["batching"], "worker_alive": alive, "queued": batcher.q.qsize()}}
|
| 161 |
|
| 162 |
class Handler(BaseHTTPRequestHandler):
|
| 163 |
protocol_version = "HTTP/1.1"
|
|
|
|
| 171 |
# on the client's delayed ACK, a flat ~40 ms on every request of a kept-alive connection
|
| 172 |
self.connection.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
|
| 173 |
|
| 174 |
+
def _send(self, code: int, obj: dict, headers: dict | None = None) -> None:
|
| 175 |
body = json.dumps(obj, ensure_ascii=False).encode("utf-8")
|
| 176 |
self.send_response(code)
|
| 177 |
self.send_header("Content-Type", "application/json")
|
| 178 |
self.send_header("Content-Length", str(len(body)))
|
| 179 |
+
for name, value in (headers or {}).items():
|
| 180 |
+
self.send_header(name, value)
|
| 181 |
self.end_headers()
|
| 182 |
self.wfile.write(body)
|
| 183 |
|
| 184 |
def do_GET(self):
|
| 185 |
if self.path.rstrip("/") in ("/healthz", "/health"):
|
| 186 |
+
return self._send(200, current_health())
|
| 187 |
return self._send(404, {"error": "not found"})
|
| 188 |
|
| 189 |
def do_POST(self):
|
|
|
|
| 197 |
if not isinstance(req, dict):
|
| 198 |
return self._send(400, {"error": "the body must be a JSON object"})
|
| 199 |
try:
|
| 200 |
+
if batcher is not None:
|
| 201 |
+
out = batcher.submit(req.get("state"), req.get("questions"))
|
| 202 |
+
else:
|
| 203 |
+
with lock:
|
| 204 |
+
out = blink.decide(req.get("state"), req.get("questions"))
|
| 205 |
except blink.BlinkError as exc:
|
| 206 |
return self._send(422, {"error": str(exc)})
|
| 207 |
+
except busy as exc:
|
| 208 |
+
return self._send(503, {"error": str(exc)}, {"Retry-After": "1"})
|
| 209 |
except Exception as exc: # noqa: BLE001 - report, keep serving
|
| 210 |
return self._send(500, {"error": f"{type(exc).__name__}: {exc}"})
|
| 211 |
return self._send(200, {
|