mirror of
https://github.com/NandhaKishorM/laya.git
synced 2026-09-28 16:02:56 +08:00
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:
+24
-9
@@ -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}
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user