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)