文档导航

API 参考

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.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 和 predict 上的逐调用 hooks= 作用于整次调用,包括 on_route。
  • 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)

另见