Merge pull request #611 from aashish254/feat-demo-predict-batch

fix(examples): /predict/batch must call Router.predict_batch, not predict per state

# Conflicts:
#	tests/test_example_server_limits.py
This commit is contained in:
Nandakishor
2026-09-28 00:54:22 +05:30
2 changed files with 159 additions and 9 deletions
+24 -9
View File
@@ -207,7 +207,10 @@ def _router() -> Router:
def _predict(state: Any, questions: Dict[str, Any], **kw: Any) -> Dict[str, Any]:
"""The one place that calls Router.predict.
"""The one place that calls Router.predict for a single state.
`/predict/batch` goes straight to `Router.predict_batch` instead, so a batch shares forward
passes; it reaches `_predict` only as the per-state fallback when the batch call fails.
No lock needed here: Router's own model lifecycle (load/evict/LRU) is thread-safe as of
laya 0.3.5 (fixes #95), and inference is deliberately left outside Router's internal lock
@@ -344,14 +347,26 @@ def predict_batch(req: BatchRequest) -> Dict[str, Any]:
for state in req.states:
_check_request_limits(state, req.questions)
questions = _questions(req.questions)
results: List[Dict[str, Any]] = []
for i, state in enumerate(req.states):
try:
results.append(
_predict(state, questions, model=req.model, task=req.task, lang=req.lang)
)
except Exception as exc:
results.append({"index": i, "error": f"{type(exc).__name__}: {exc}"})
# Only the controls that were actually set: `route_batch` reads them with `.get`, so an
# omitted key and an explicit null mean the same thing, and the request dicts stay minimal.
controls = {key: value for key, value in (("model", req.model), ("task", req.task),
("lang", req.lang)) if value is not None}
requests = [{"state": state, "questions": questions, **controls} for state in req.states]
try:
# One call, not one per state: `Router.predict_batch` routes the whole batch, groups it by
# checkpoint and shares a forward pass across states that carry the same question schema --
# which is exactly this endpoint, since `BatchRequest` holds one `questions` map.
results: List[Dict[str, Any]] = list(_router().predict_batch(requests))
except Exception:
# The batch fails as a unit, so a single bad state would otherwise cost every other state
# its answer. Fall back to the per-state path to keep the documented envelope: N results,
# with `{"index": i, "error": ...}` only where a state genuinely failed.
results = []
for i, state in enumerate(req.states):
try:
results.append(_predict(state, questions, **controls))
except Exception as item_exc:
results.append({"index": i, "error": f"{type(item_exc).__name__}: {item_exc}"})
return {"count": len(results), "results": results}
+135
View File
@@ -371,6 +371,141 @@ def main():
sys.modules.pop("uvicorn", None)
with_cap(None) # leave the module as the rest of the file found it
# --- /predict/batch must make ONE Router call, not one per state ---------
#
# README teaches `Router.predict_batch` as the way to answer many states with shared forward
# passes ("routes the full workload first, groups requests by checkpoint ... results are
# restored to the original request order"), and this endpoint is the batch-shaped surface the
# README points at for trying Laya without writing code. The handler used to call
# `Router.predict` once per state anyway, so 64 states were 64 forwards. These drive the real
# app -- validation, `_questions()`, the handler body, the JSON envelope -- over a recording
# stand-in at `demo.ROUTER`, so no weights are needed to count the calls.
import inspect
from laya.router import Router as CoreRouter
class RecordingRouter:
"""Answers like the Router does, and remembers how it was asked."""
def __init__(self, fail_on=None):
self.predict_calls = []
self.batch_calls = []
self.fail_on = fail_on
def predict(self, state, questions, **kw):
self.predict_calls.append((state, dict(kw)))
return self._answer(state, questions)
def predict_batch(self, requests, **kw):
self.batch_calls.append((list(requests), dict(kw)))
return [self._answer(r["state"], r["questions"]) for r in requests]
def _answer(self, state, questions):
if self.fail_on and self.fail_on in str(state):
raise ValueError("simulated failure for " + self.fail_on)
return {"answers": {"a": {"type": "noul", "choice": False}}, "state": state,
"questions": questions}
states = ["ticket %d" % i for i in range(8)]
class LegacyRouter(RecordingRouter):
"""A Router-like object whose `predict_batch` predates the batch path entirely."""
predict_batch = None
# A Router that predates `predict_batch` must still be usable through the fallback.
older = LegacyRouter()
demo.ROUTER = older
legacy = TestClient(demo.app, raise_server_exceptions=False).post(
"/predict/batch", json={"states": states, "questions": one})
ok("a router without predict_batch still answers every state",
legacy.status_code == 200 and len(legacy.json()["results"]) == 8
and not [r for r in legacy.json()["results"] if "error" in r],
"%s / %s" % (legacy.status_code, json.dumps(legacy.json())[:200]))
router = RecordingRouter()
demo.ROUTER = router
batched = TestClient(demo.app, raise_server_exceptions=False).post(
"/predict/batch", json={"states": states, "questions": one})
body = batched.json()
ok("an 8-state batch is ONE predict_batch call and zero predict calls",
len(router.batch_calls) == 1 and not router.predict_calls,
"predict_batch=%d predict=%d" % (len(router.batch_calls), len(router.predict_calls)))
ok("the batch envelope is unchanged: count and one result per state, in order",
body.get("count") == 8 and [r.get("state") for r in body.get("results", [])] == states)
# Read defensively: an endpoint that never batches has nothing to inspect, and the checks below
# say so by name instead of letting this script die at the unpack.
requests_sent, call_kwargs = (router.batch_calls[0] if router.batch_calls else ([], {}))
ok("each request carries state + questions, and the questions map is the same one",
len(requests_sent) == 8
and all(sorted(r) == ["questions", "state"] for r in requests_sent)
and all(r["questions"] == requests_sent[0]["questions"] for r in requests_sent))
controls_sent = [k for r in requests_sent for k in r if k in ("model", "task", "lang")]
ok("unset controls are absent, not sent as null",
len(requests_sent) == 8 and not call_kwargs and not controls_sent,
"call kwargs=%r controls=%r" % (call_kwargs, controls_sent))
# Every key the endpoint puts in a request must be one core actually reads out of it, and the
# set of those keys is derived from `Router.route_batch`'s own source, so a rename or a new
# override in core shows up here rather than silently stopping reaching the router.
import re
route_src = inspect.getsource(CoreRouter.route_batch)
read = ({"state", "questions"}
| set(re.findall(r'request\["(\w+)"\]', route_src))
| set(re.findall(r'request\.get\("(\w+)"', route_src)))
sent_keys = {k for r in requests_sent for k in r}
ok("the request keys the endpoint sends are ones core reads",
len(requests_sent) == 8 and read >= sent_keys,
"core reads %r, endpoint sends %r" % (sorted(read), sorted(sent_keys)))
ok("predict_batch is still a one-positional-list call",
list(inspect.signature(CoreRouter.predict_batch).parameters)[1] == "requests")
pinned = RecordingRouter()
demo.ROUTER = pinned
TestClient(demo.app, raise_server_exceptions=False).post(
"/predict/batch", json={"states": states, "questions": one,
"model": "multilingual", "lang": "de"})
sent = pinned.batch_calls[0][0] if pinned.batch_calls else []
ok("a pinned model/lang travels with every request",
len(sent) == 8
and all(r["model"] == "multilingual" and r["lang"] == "de" for r in sent)
and not any("task" in r for r in sent),
json.dumps(sent[:1])[:200])
# One state failing must not cost its neighbours their answer: the endpoint's published
# contract is per-item errors inside a 200.
partial = RecordingRouter(fail_on="ticket 3")
demo.ROUTER = partial
poison = TestClient(demo.app, raise_server_exceptions=False).post(
"/predict/batch", json={"states": states, "questions": one}).json()
errors = [r for r in poison["results"] if "error" in r]
ok("one failing state yields one error entry, not an empty batch",
len(poison["results"]) == 8 and [r["index"] for r in errors] == [3]
and "simulated failure" in errors[0]["error"],
json.dumps(poison)[:240])
ok("the failing batch retried per state, so its neighbours still answered",
len(partial.batch_calls) == 1 and len(partial.predict_calls) == 8)
# The single-state surface must keep going through predict(), and `/predict/batch` with one
# state must still batch -- otherwise the two endpoints diverge on where hooks fire.
single = RecordingRouter()
demo.ROUTER = single
one_state = TestClient(demo.app, raise_server_exceptions=False)
one_state.post("/predict", json={"state": "ticket 0", "questions": one})
ok("/predict still calls Router.predict once",
len(single.predict_calls) == 1 and not single.batch_calls)
demo.ROUTER = None
not_ready = TestClient(demo.app, raise_server_exceptions=False).post(
"/predict/batch", json={"states": states, "questions": one})
ok("an unready router keeps answering 200 with per-item 503s",
not_ready.status_code == 200
and len(not_ready.json()["results"]) == 8
and all("503" in r.get("error", "") for r in not_ready.json()["results"]),
json.dumps(not_ready.json())[:200])
demo.ROUTER = None
print("\n%d passed, %d failed" % (len(PASS), len(FAIL)))
for f in FAIL:
print(" FAIL " + f)