文件導航

API 參考

這一頁的所有東西都能從 laya(常見名字)或 laya.hooks(完整表面)匯入。

from laya import PredictContext, PredictHook, Hook
from laya.hooks import HOOK_EVENTS, normalise_hooks, dispatch, aggregate_usage

PredictContext

每次公開呼叫建立一次 PredictContext,並傳給那次呼叫的每個鉤子。Router.predict_batch 是唯一 的例外:它為每個請求各建立一個,就像該請求的一次 Router.predict 呼叫會做的那樣,所以它的 Router 級鉤子每個請求觸發一次,每次帶一個狀態。它是可變的:鉤子可以改寫 states、 questions 和 results,on_route 可以改寫 decision。它用身份相等(eq=False),所以一個 上下文可雜湊,兩個上下文永不相等。

@dataclass(eq=False)
class PredictContext:
    states: List[Any]
    questions: Dict[str, Any]
    run_id: str = <uuid4 hex>
    results: Optional[List[Dict[str, Any]]] = None
    decision: Optional[Dict[str, Any]] = None
    model: Optional[str] = None
    agent: Any = None
    router: Any = None
    max_len: Optional[int] = None
    head_max_len: Optional[int] = None
    usage: Optional[Dict[str, int]] = None
    started_at: float = <perf_counter()>
    elapsed_ms: Optional[float] = None
    error: Optional[BaseException] = None
欄位 型別 何時設定 可變 含義
states list 總是 是(start) 這次呼叫的狀態。system_one/Router.predict 傳一個;Agent.predict_batch 傳多個;Router.predict_batch 每個請求傳一個。start 鉤子可以替換這個列表。
questions dict 總是 是(start) 問題。start 鉤子可以替換這個 dict。
run_id str 總是 否 這次呼叫的每個鉤子共享的一個唯一 id。用它來關聯事件和 span。
results list | None end(以及 skip 時) 是(end) 逐狀態的結果 dict,每個的形狀和 system_one 的返回一樣。推理完成前為 None。
decision dict | None 僅 Router 是(route) 選定 checkpoint 的那個 RouteDecision(一個 dict)。
model str | None 總是 否 checkpoint id:Agent 是 Agent.model_id,Router 是解析後的別名(例如 "english")。
agent Agent | ONNXAgent | None predict 事件 否 回答這次呼叫的執行時。
router Router | None Router 事件 否 在起作用時的 Router。
max_len int | None 總是 是(start) encoder 的逐呼叫 token 預算。None 用 agent 配置。
head_max_len int | None 總是 是(start) 問題 head 的逐呼叫 token 預算。None 用 agent 配置。
usage dict | None end 是(end) {"input_tokens", "output_tokens"},對這次呼叫的各個狀態求和。
started_at float 總是 否 呼叫開始時的 time.perf_counter()。
elapsed_ms float | None end 否 整次呼叫的牆鍾時間,毫秒。
error BaseException | None 失敗路徑 否 異常,在 on_error 和 on_predict_end 之前設定。

PredictContext.skip(results)

短路推理。從 on_predict_start 呼叫它,它會設定 ctx.results,於是前向傳播被跳過; on_predict_end 仍然執行,並返回提供的那些結果。

key 必須覆蓋答案所依賴的一切,鉤子必須覆蓋這次呼叫所攜帶的一切:鉤子每次呼叫觸發一次,而 predict_batch 會一次性用所有狀態呼叫它們。

CACHE = {}

def key(ctx, index):
    # Not sort_keys=True: criteria order is positional, so two orders are two questions,
    # and the checkpoint and token budget change the answer too.
    payload = json.dumps([ctx.states[index], ctx.questions, ctx.model,
                          ctx.max_len, ctx.head_max_len], default=str)
    return hashlib.sha256(payload.encode()).hexdigest()

def cache_read(ctx):
    hits = [CACHE.get(key(ctx, i)) for i in range(len(ctx.states))]
    if all(hit is not None for hit in hits):
        ctx.skip(hits)   # one entry per state, same shape as predict_batch's return

def cache_write(ctx):
    for i, result in enumerate(ctx.results or []):
        CACHE[key(ctx, i)] = result

laya.load("convaiinnovations/laya", on_predict_start=cache_read, on_predict_end=cache_write)

tests/test_hooks_api.py 會執行這段程式碼、examples/hooks/cache.py,以及 docs/hooks/patterns.md 和 docs/hooks/examples.md 裡的快取塊,並斷言這四處的 key 算得一樣, 所以一頁不會教一個示例已經改掉的 key。

在 Router 上,一個被跳過的載荷會加上 routing 鍵(不覆蓋它已有的一個),所以 Router.predict 保持它文件化的返回形狀。

Hook 協議

Hook 是一個 typing.Protocol。實現方法的任意子集;其餘的被跳過。

class Hook(Protocol):
    def on_predict_start(self, ctx: PredictContext) -> None: ...
    def on_predict_end(self, ctx: PredictContext) -> None: ...
    def on_route(self, ctx: PredictContext) -> None: ...
    def on_load(self, ctx: PredictContext) -> None: ...
    def on_evict(self, ctx: PredictContext) -> None: ...
    def on_error(self, ctx: PredictContext) -> None: ...
事件 位置 執行時機 能改什麼
on_predict_start Agent、Router tokenization/前向之前 states、questions,或 skip()
on_predict_end Agent、Router 結果存在之後,無論成功還是失敗 results
on_route Router 檢測之後、載入之前 decision
on_load Router 一個 checkpoint 構建之後 無(觀察)
on_evict Router 一個 checkpoint 釋放之後 無(觀察)
on_error Agent、Router 一次 predict 呼叫失敗時 無(觀察)

一個鉤子可以自由地定義額外的屬性和方法;只有這六個事件名會被讀取。如果一個鉤子把這六個之一定義 成不可呼叫的東西,配置會快速失敗(見校驗)。

BaseHook

BaseHook 是這個協議的具象對應物:一個每個事件都是空實現的類。繼承它,只覆蓋你需要的事件。

from laya import BaseHook

class Audit(BaseHook):
    def on_predict_end(self, ctx):
        ship(ctx.run_id, ctx.results)

想要結構化型別(任何有合適方法的物件)時 Hook 最好;想要一個顯式的基類去繼承並呼叫 super() 時 BaseHook 最好。

便捷型別

PredictHook = Callable[[PredictContext], None]

PredictHook 是和 on_predict_start= / on_predict_end= 一起用的普通可呼叫物件的型別。傳一個 可呼叫物件或一串它們;每一個都被包成一個最小鉤子。

執行時註冊

每個執行時都混入 HookRegistry,所以鉤子可以在構造之後新增、移除或限定作用域。變更執行緒安全; 一次呼叫讀取列表的一個快照,所以新增或移除鉤子永遠不會打擾在途的一次呼叫。

agent.add_hook(tracer)              # one hook or a sequence; returns self for chaining
agent.remove_hook(tracer)           # by identity; True if it was installed

with agent.hooks_installed(debug):  # installed for the block, removed on exit
    agent.system_one(state, questions)

add_hook 接受和 hooks= 一樣的物件(不是普通可呼叫物件)。hooks_installed 接受任意數量的 鉤子物件或序列,並在退出時恢復之前的列表,程式碼塊拋異常時也一樣。

程序級預設值

laya.hooks 維護一個小巧的程序級登錄檔,所以一個 tracer、指標鉤子或租戶標記器不必穿過每個 Agent 和 Router。預設值最先執行,然後是例項上安裝的鉤子,然後是逐呼叫鉤子。

from laya import hooks

hooks.set_default_hooks(hooks=[Tracer()])          # replaces the set, accepts the hooks= arguments
hooks.add_default_hook(Metrics())                  # appends
hooks.clear_default_hooks()                        # removes everything
hooks.default_hooks()                              # a copy of the current list

hooks.compose_hooks(agent.hooks)                   # defaults + installed (advanced)

預設值適用於每個事件,包括 Router 生命週期事件 on_load 和 on_evict。登錄檔在呼叫時讀取, 所以在 Agent 或 Router 構建之後設定的鉤子仍然生效。沒有按例項的退出開關;呼叫 clear_default_hooks() 來關掉程序級的那一組。

非同步鉤子

一個事件可以是一個協程。把鉤子包進 AsyncHook,它的 async def 方法就會在同步核心裡執行到 完成:

from laya import AsyncHook

class Remote:
    async def on_predict_end(self, ctx):
        await ship(ctx.results)

agent = laya.load("convaiinnovations/laya", hooks=[AsyncHook(Remote())])

傳給 on_predict_start= / on_predict_end= 的普通非同步可呼叫物件也行,因為 dispatch 會執行 鉤子返回的任何可等待物件。

協程在哪裡執行:

  • 如果呼叫執行緒沒有正在執行的事件迴圈,用 asyncio.run 執行它。
  • 如果它已經有一個(呼叫方在一個非同步函式內部),它在一個專用的後臺迴圈上執行,於是呼叫執行緒可以 阻塞而不會死鎖。傳 AsyncHook(hook, loop=...) 讓它匯入某個特定的迴圈;那個迴圈必須正在執行, 且不能是呼叫執行緒自己的迴圈。兩者都會被檢查:一個已停止的迴圈和呼叫方自己的迴圈都拋 ValueError,而不是永遠阻塞。

一個沒有 async 方法的鉤子不受影響。

配置表面

每個入口點都接受同樣的鉤子參數。hooks 接受一個物件或一串物件;on_predict_start / on_predict_end 接受一個可呼叫物件或一串。

參數 型別 預設值 含義
hooks Hook | Sequence[Hook] | None None 生命週期鉤子(六個事件中的任意個)。
on_predict_start PredictHook | Sequence[PredictHook] | None None 某個事件的便捷可呼叫物件。
on_predict_end PredictHook | Sequence[PredictHook] | None None 某個事件的便捷可呼叫物件。
hooks_raise bool True True:鉤子異常向外傳播。False:警告並繼續。
hooks_concurrent bool True False:在鎖下逐個派發鉤子。
hooks_timeout float | None None 每個鉤子的時間限制,秒;None 表示不限。

Agent

Agent(
    model_id_or_path="convaiinnovations/laya",
    device=None, token=None, subfolder=None, fast=False, compile=False,
    hooks=None, on_predict_start=None, on_predict_end=None,
    hooks_raise=True, hooks_concurrent=True, hooks_timeout=None,
)

load(..., hooks=None, on_predict_start=None, on_predict_end=None,
     hooks_raise=True, hooks_concurrent=True, hooks_timeout=None)

agent.predict_batch(states, questions, batch_size=None,
                    hooks=None, on_predict_start=None, on_predict_end=None, hooks_raise=None,
                    hooks_timeout=None, max_len=None, head_max_len=None, sort_by_length=False)

agent.system_one(state, questions,
                 hooks=None, on_predict_start=None, on_predict_end=None, hooks_raise=None,
                 hooks_timeout=None, max_len=None, head_max_len=None)

agent.predict_long(state, questions, window=None, stride=None, aggregate="auto",
                   batch_size=None, lang=None,
                   hooks=None, on_predict_start=None, on_predict_end=None, hooks_raise=None,
                   hooks_timeout=None)

agent.predict(...)          # alias of system_one
  • 逐呼叫方法上的 hooks_raise 和 hooks_timeout 預設是 None,意思是「用例項的值」。

  • hooks_concurrent 只在例項級設定。

  • 在 predict_long 上,鉤子包住回答該狀態的推理,而對一份需要多個視窗的文件來說,那就是覆蓋 它們的那次共享的 predict_batch:on_predict_start 觸發一次,ctx.states 以掃描順序儲存解碼 出的視窗文本,而不是呼叫方的狀態(那個狀態是為了產出它們而被 tokenize 的)。states 從 start 鉤子起可變,所以到達推理的那次掃描不必是 predict_long 算出的切分。改變的是答案能聲稱什麼:

    start 鉤子做了什麼 usage["windows"] answer["window"]
    用 ctx.skip([result]) 作答 0 缺席 —— 沒有視窗給它打分
    讓掃描保持原樣 N 存在 —— index、token_start/token_end 指明決定性的那段
    以任何方式替換了掃描 被打分的那些狀態 缺席 —— 這些偏移描述的是 predict_long 的視窗,不是被打分的文本

    usage["windows"] 在每條路徑上都是合計,包括狀態本來就能裝進單個視窗的那條路徑,所以一個快取 答案永遠不會讀成模型讀過的一個視窗。

Router

Router(
    models=None, device=None, token=None, max_loaded=2, default="english",
    auto_task_detection=False, standalone_repos=False, preload=False, lang_guess=None,
    hooks=None, on_predict_start=None, on_predict_end=None,
    hooks_raise=True, hooks_concurrent=True,
)

router.route(state, questions=None, model=None, task=None, lang=None, lang_guess=None,
             hooks=None, hooks_raise=None)

router.predict(state, questions, model=None, task=None, lang=None, lang_guess=None,
               hooks=None, on_predict_start=None, on_predict_end=None, hooks_raise=None,
               hooks_timeout=None, max_len=None, head_max_len=None)

router.predict_batch(requests, batch_size=None, hooks_timeout=None, min_confidence=None,
                     sort_by_length=False, hooks=None, on_predict_start=None, on_predict_end=None,
                     hooks_raise=None)

router.system_one(...)      # alias of predict
router.load(name)           # builds on first use; fires on_load
router.preload(names=None)  # builds several; fires on_load per build
router.unload(name=None)    # frees one or all; fires on_evict
router.attach(name, agent)  # registers an existing agent; does not fire on_load
router.loaded               # list of resident checkpoint names
  • route、route_batch、predict 和 predict_batch 上的逐呼叫 hooks= 作用於整次呼叫,包括 on_route。在 predict_batch 上,列表的組裝方式與 predict 相同(先已安裝的鉤子,再逐呼叫列表;None 和 [] 不新增任何東西),並且每個請求執行一次。
  • route() 是公開的:呼叫它會用已安裝的鉤子外加任何逐呼叫 hooks 派發 on_route。

ONNXAgent

ONNXAgent(model_id_or_path, onnx_path="laya.onnx", subfolder=None,
          hooks=None, on_predict_start=None, on_predict_end=None,
          hooks_raise=True, hooks_concurrent=True)

onnx_agent.system_one(state, questions,
                      hooks=None, on_predict_start=None, on_predict_end=None, hooks_raise=None,
                      hooks_timeout=None, max_len=None, head_max_len=None)

onnx_agent.predict(...)     # alias of system_one

ONNXAgent 沒有 Router,所以它只暴露 predict 級事件。

事件載荷

按事件和執行時,哪些欄位會被填充:

事件 執行時 states questions decision model agent router results usage elapsed_ms error
on_predict_start Agent ✓ ✓ – ✓ ✓ – – – – –
on_predict_start Router ✓ ✓ ✓ ✓ ✓ ✓ – – – –
on_predict_end Agent ✓ ✓ – ✓ ✓ – ✓ 成功時 ✓ 失敗時
on_predict_end Router ✓ ✓ ✓ ✓ ✓ ✓ ✓ 成功時 ✓ 失敗時
on_error 兩者 ✓ ✓ ✓(Router) ✓ ✓ ✓(Router) – – – ✓
on_route Router ✓ ✓ ✓ – – ✓ – – – –
on_load Router [] {} – ✓ ✓ ✓ – – – –
on_evict Router [] {} – ✓ – ✓ – – – –

計時細節:

  • on_predict_end 在成功路徑上能看到 results。失敗路徑上 results 是 None,除非某個 start 鉤子通過 skip() 設過它,而 usage 因此也是 None(它由 results 匯出);elapsed_ms 總是會被設定。
  • on_error 在計算 elapsed_ms 和 usage 的 finally 塊之前執行,所以那裡兩者都是 None。 計時和 usage 請從 on_predict_end 讀。
  • run_id 總是被填充。

校驗

配置在鉤子被規範化時校驗,安裝的鉤子在構造時發生,逐呼叫鉤子在呼叫時發生。這些會拋 TypeError:

情況 訊息
傳的是類而不是例項 hooks entries must be instances, not classes; ...
一個物件沒有實現這六個事件中的任何一個 hooks entries must implement at least one of ...
某個事件屬性不可呼叫 hooks entry X.on_predict_start must be callable, got int
on_predict_start= / on_predict_end= 不可呼叫 on_predict_start must be callable, got int

hooks= 不接受普通的可呼叫物件,因為一個裸的可呼叫物件說不清它是哪個事件的。那些請用 on_predict_start= / on_predict_end=。

進階輔助函式

這些在內部使用且穩定,但大多數使用者不需要它們。

HOOK_EVENTS          # tuple of the six event names, in dispatch order
normalise_hooks(hooks=None, on_predict_start=None, on_predict_end=None) -> list
dispatch(hooks, event, ctx, *, raise_errors=True, lock=None) -> None
aggregate_usage(results) -> {"input_tokens": int, "output_tokens": int}
dispatch(hooks, event, ctx, *, raise_errors=True, lock=None, timeout=None)
run_coroutine_sync(coro, loop=None)

normalise_hooks 把一個 hooks 物件/序列和那兩個可呼叫物件攤平成一個有序列表。dispatch 在 每個實現了 event 的鉤子上呼叫它,應用丟擲策略、鎖和超時,並在鉤子的結果可等待時執行它。 run_coroutine_sync 從同步程式碼裡把一個可等待物件執行到完成,在呼叫方自己的迴圈空閒時用它,呼叫方 已經有一個迴圈時用後臺迴圈。aggregate_usage 把逐狀態的 usage 塊加起來。

from laya.hooks import normalise_hooks, dispatch, PredictContext

hooks = normalise_hooks(on_predict_start=[log, redact])
ctx = PredictContext(states=["..."], questions={...})
dispatch(hooks, "on_predict_start", ctx)

另見