文件導航

生命週期

這一頁講每個入口點的事件確切順序。如果只看一張圖,就看 Router 那張;它是 超集。

Agent.predict_batch

predict_batch 是唯一的實現;system_one 和 predict 用一個狀態呼叫它。

predict_batch(states, questions, batch_size=..., hooks=..., ...)
  │
  ├─ active  = installed hooks + per-call hooks         (installed first)
  ├─ ctx     = PredictContext(states, questions, model=self.model_id, agent=self)
  │
  ├─ try:
  │    │
  │    ├─ on_predict_start ─────────────────────────────┐
  │    │      a hook may:                                │
  │    │        • rewrite ctx.states / ctx.questions     │
  │    │        • set ctx.max_len / ctx.head_max_len     │
  │    │        • ctx.skip(results) ─────────────┐       │
  │    │        • raise (aborts; see errors)     │       │
  │    │                                         │      │
  │    ├─ if ctx.results is not None:  ◄─────────┘      │  cache hit
  │    │      skip tokenization and forward             │
  │    ├─ else:                                         │
  │    │      validate states is a list                 │
  │    │      for each batch chunk:                     │
  │    │        _encode_state ─► collate ─► _forward    │
  │    │      _decode_answers                           │
  │    │      ctx.results = [...]                       │
  │    │                                                │
  │    └─ (any failure here) ──► except BaseException:  │
  │              ctx.error = exc                        │
  │              on_error                               │
  │              re-raise                               │
  │                                                     │
  │   finally:                                          │
  │     ctx.elapsed_ms = now - ctx.started_at           │
  │     if ctx.results: ctx.usage = aggregate_usage(...)│
  │     on_predict_end ─────────────────────────────────┘
  │
  └─ return ctx.results

同一個 ctx 物件貫穿 start、error 和 end,所以 run_id 把它們關聯起來,end 鉤子可以讀 ctx.error。

Agent.system_one / predict

system_one(state, questions, hooks=..., ...)
  └─ predict_batch([state], questions, hooks=..., ...)[0]

所以 system_one 繼承每個鉤子和同樣的生命週期,ctx.states == [state]。

Router.predict

Router.predict(state, questions, model=..., hooks=..., on_predict_start=..., on_predict_end=...)
  │
  ├─ active = installed hooks + per-call hooks
  │
  ├─ route(state, questions, ..., hooks=per-call, hooks_raise=...)
  │    │
  │    ├─ _route(...)                      detect script / language / workflow
  │    ├─ on_route  ──► ctx.decision       a hook may replace the decision
  │    └─ return ctx.decision
  │
  ├─ load(decision["model"])
  │    │
  │    ├─ already resident? ──► return it
  │    ├─ else build Agent(...) ──► on_load   (after the Router lock is released)
  │    └─ evict LRU checkpoints ──► on_evict  (after the Router lock is released)
  │
  ├─ ctx = PredictContext(states=[state], questions, decision, model=decision.model,
  │                       agent=agent, router=self)
  ├─ try:
  │    ├─ on_predict_start
  │    ├─ if ctx.results is None:
  │    │      result = agent.system_one(ctx.states[0], ctx.questions)
  │    │        └─ the Agent's own hooks run here (start / forward / end)
  │    │      result["routing"] = decision
  │    │      ctx.results = [result]
  │    └─ else:
  │           for each cached result: result.setdefault("routing", decision)
  │    └─ (any failure) ──► except: on_error, re-raise
  │    └─ finally: elapsed_ms, usage, on_predict_end
  │
  └─ return ctx.results[0]

要點:

  • on_route 在模型載入之前執行,所以鉤子可以釘住一個 checkpoint,避免載入另一個。
  • Router 級的 predict 鉤子包住整次呼叫。它們不會被轉發進 Agent;一個掛載的 Agent 若有自己 的鉤子,也會執行它們,這是預期內的。
  • Router 級的 ctx.skip() 仍然會加上 routing,所以返回形狀穩定。

Router.predict_batch

每個結果就是 predict 對該請求返回的東西,所以 Router 級的 predict 鉤子在這裡也是每請求執行一次: 每個請求都得到自己的 PredictContext、run_id 和 elapsed_ms。

Router.predict_batch(requests, batch_size=..., hooks=...)
  │
  ├─ active = installed hooks + per-call hooks          (installed first; None and [] add nothing)
  ├─ route_batch(requests, hooks=per-call) ──► on_route, once per request   (no checkpoint loaded yet)
  │
  └─ for each checkpoint, in order of first appearance:
       │
       ├─ load(checkpoint) ──► on_load / on_evict
       ├─ for each request of this checkpoint, in input order:
       │      ctx = PredictContext(states=[state], questions, decision, model, agent, router,
       │                           max_len=request.get("max_len"),
       │                           head_max_len=request.get("head_max_len"))
       │      on_predict_start       a hook may redact, rewrite, set a token budget or skip
       ├─ group the requests left to infer by (questions, ctx.max_len, ctx.head_max_len)
       │      agent.predict_batch(states, questions, ...)  ──► one shared forward pass per group
       │      result["routing"] = decision;  ctx.results = [result]
       ├─ ctx.usage, once per request
       ├─ (any failure) ──► for every started request, in reverse input order:
       │                    ctx.error = exc, on_error, on_predict_end;  then re-raise
       └─ on_predict_end, once per request of this checkpoint, in reverse input order

要點:

  • 一個 checkpoint 的各個請求,其每個 start 鉤子都在它們任何一個 end 鉤子之前執行,因為它們共享 前向傳播。所以在 on_predict_end 裡填充的快取無法在同一個 checkpoint 組內服務一個重複的狀態; 跨呼叫可以。
  • 同樣的原因,請求結束的順序和它們開始的順序相反,所以一個在 start 裡設定、在 end 裡重置東西的 鉤子(一個 contextvars 值,一次 OpenTelemetry 的 context.attach / detach)會展開回它發現 時的值。
  • 一個替換 ctx.states、ctx.questions 或 token 預算的 start 鉤子只改變它自己的請求:請求是在 它們的 start 鉤子執行之後才為前向傳播分組的。原地改動一個共享的 questions dict 則不同,也不是 predict 的做法:這個組裡沒有任何東西會被推理,直到它的所有 start 鉤子都執行過,所以改動會抵達 每一個共享該 dict 的請求,包括那些 start 鉤子更早執行的請求,也包括呼叫方。請改為給 ctx.questions 賦一個新的 dict。
  • checkpoint 名在分組之前解析,所以一個用別名("ml")釘住某個請求的 on_route 鉤子會共享那個 checkpoint 的前向傳播,而 ctx.model 是解析後的名字。
  • 一個請求可以攜帶它自己的 max_len / head_max_len,那是 predict 作為呼叫參數接受的 token 預算的逐請求形式。一個設定 ctx.max_len 的 start 鉤子會覆蓋它,因為鉤子在上下文構建之後執行。 要求不同預算的請求無法共享一次前向傳播,所以一個混用預算的批次會按預算各做一次 agent.predict_batch 呼叫。
  • 每個已開始的請求都恰好得到一個 on_predict_end,即便另一個請求的 end 鉤子拋異常;第一個這樣的 錯誤在所有它們都執行完之後才丟擲。
  • 如果一個 checkpoint 組失敗,它每一個已開始的請求都以那個異常失敗:每個都拿到 on_error, ctx.error 設為它,然後是 on_predict_end。這包括一次快取命中,以及一個其問題組已經跑過的 請求,因為呼叫方為它們中任何一個都拿到異常、拿不到結果;這也意味著 ctx.error 可能是另一個請求 的失敗(一個為某個請求拋異常的 start 鉤子會讓它所在的組失敗)。已經完成的 checkpoint 組的請求 已經帶著它們的結果結束了,就像更早的 [router.predict(...) for ...] 呼叫本會做的那樣。

模型生命週期

on_load 在一個 checkpoint 被構建時觸發;on_evict 在一個被釋放時觸發。兩者都在 Router 內部 鎖釋放之後執行,所以鉤子可以安全地回撥進 Router。

load("multilingual")
  │
  ├─ [lock]
  │    build Agent(...)          (seconds: download + weights)
  │    register in _agents / _order
  │    evict LRU if over max_loaded ──► evicted = ["english"]
  ├─ [unlock]
  ├─ on_evict("english")
  └─ on_load("multilingual")

unload("english")
  ├─ [lock] remove from _agents / _order
  ├─ [unlock]
  └─ on_evict("english")

attach(name, agent) 註冊一個已有的 agent,且不會觸發 on_load,因為沒有構建任何 checkpoint。

用 skip 做快取

on_predict_start
  ├─ cache hit?  ctx.skip([cached_result])
  │     └─ forward pass skipped
  │     └─ on_predict_end still runs
  │     └─ Router adds `routing` if missing
  └─ cache miss? nothing
        └─ forward pass runs
        └─ on_predict_end can store the result

一個能用的快取見 examples/hooks/cache.py。

空輸入

鉤子仍會觸發,所以審計能看到每一次呼叫:

輸入 end 時的 ctx.results
predict_batch([]) []
predict_batch(states, {})(沒有 questions) 每個狀態一個空答案載荷
system_one(state, {}) 單個空答案載荷

這些情況下不發生 tokenization 或前向傳播,但 on_predict_start 和 on_predict_end 會執行。

順序規則

  1. 已安裝的鉤子總在逐呼叫鉤子之前執行。
  2. 在一個列表內,鉤子按列表順序執行。
  3. 對於一個事件,實現它的每個鉤子都按那個順序執行,然後才輪到下一個事件。
  4. 在失敗路徑上 on_error 在 on_predict_end 之前執行。
  5. 當一次 load 同時驅逐和構建時,on_evict 在 on_load 之前執行。
installed: [A, B]   per-call: [C]
on_predict_start: A, B, C
on_predict_end:   A, B, C

併發

Agent 和 Router 可以從很多執行緒安全地呼叫。每次呼叫建立它自己的 PredictContext,所以上下文 永不跨請求洩漏。唯一共享的狀態就是鉤子物件本身,所以一個非執行緒安全的鉤子必須要麼自己守住它的 狀態,要麼用 hooks_concurrent=False 安裝。

hooks_concurrent=True (default)      hooks_concurrent=False
  thread 1 ─┐                          thread 1 ─┐
  thread 2 ─┼─ hooks run in parallel   thread 2 ─┼─ one hook at a time
  thread 3 ─┘                          thread 3 ─┘   (RLock)

hooks_concurrent=False 序列化的是每次鉤子呼叫,不是整次呼叫:兩次呼叫仍然可以在事件之間交錯。 它用一把可重入鎖,所以一個鉤子可以回撥進同一個 Agent/Router 而不會死鎖。

另見