thegovind commited on
Commit
abeb875
·
verified ·
1 Parent(s): 2751bf4

v1.1: opt-in cross-request batching in serve.py (--batch-window-ms, off by default); weights unchanged

Browse files
Files changed (3) hide show
  1. README.md +15 -0
  2. blink.py +300 -19
  3. serve.py +60 -4
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
- rows, self.last_model_ms = _forward(self.key, [w["ids"] for w in work], [w["cand"] for w in work])
 
 
431
  return (
432
  {w["qkey"]: r for w, r in zip(work, rows)},
433
- sum(len(w["ids"]) for w in work),
434
  )
435
 
436
-
437
- def _batches(lengths: list[int], budget: int):
438
- """Shortest first; a batch closes when (longest x count) would pass the padded-token budget.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- for b in _batches([len(s) for s in seqs], getattr(engine, "token_budget", TOKEN_BUDGET)):
467
- L = max(len(seqs[i]) for i in b)
468
- ids = torch.full((len(b), L), engine.pad_id, dtype=torch.long)
469
- for r, i in enumerate(b):
470
- ids[r, : len(seqs[i])] = torch.tensor(seqs[i], dtype=torch.long)
471
- ids = ids.to(device)
472
- h = model.model(input_ids=ids, use_cache=False).last_hidden_state
473
- last = torch.tensor([len(seqs[i]) - 1 for i in b], device=device)
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, health)
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
- with lock:
150
- out = blink.decide(req.get("state"), req.get("questions"))
 
 
 
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, {