快速开始#

本例生成参数范数可观测量,执行一个训练 step,在 backward 后计算数值,将其存到本地,再读回其中一个值。

import torch
import observable_library as ol

torch.manual_seed(0)
model = torch.nn.Sequential(torch.nn.Linear(2, 1))
optimizer = torch.optim.SGD(model.parameters(), lr=0.1)
observables = ol.generate(model, reductions=["l2_norm"])

source = ol.HookSource(model)
source.attach()
storage = ol.LocalStorage("./run")
runtime = ol.Runtime(
    observables=ol.Pack(observables),
    source=source,
    sink=storage,
    budget=ol.Budget(max_compute_ms=10.0),
)

inputs = torch.tensor([[1.0, -1.0], [0.5, 0.25]])
targets = torch.tensor([[0.5], [-0.25]])
optimizer.zero_grad()
predictions = model(inputs)
loss = torch.nn.functional.mse_loss(predictions, targets)
loss.backward()

values = runtime.observe(step=0)
assert values
observable_id, value = next(iter(values.items()))
stored_value = ol.query(storage, observable_id, step=0)
assert stored_value == value
print(f"{observable_id}: {stored_value}")

optimizer.step()
source.detach()

发生了什么#

  1. ol.generate() 检查 model.named_parameters(),为每个参数生成一个 l2_norm 可观测量。

  2. HookSource 提供当前参数张量。参数读取是惰性的,因此这个仅含生成参数的示例严格来说不需要 hooks;这里展示 attach(),因为 activation 和 gradient source 需要相同的生命周期。

  3. loss.backward()Runtime.observe() 之前运行。自定义可观测量读取 gradient 时必须保持这一顺序。

  4. Runtime.observe() 返回以各可观测量稳定的 16 字符 spec id 为键的数值。

  5. LocalStorage 写入 SQLite 元数据和 NumPy NPZ payload。query() 按精确 id 和 step 读回一个值。

保持显式生命周期#

对于 activation 或 gradient source,应在对应的 forward/backward 之前 attach,在 optimizer 修改参数之前 observe,并在结束后始终 detach。HookSource 保留最近捕获的张量,不强制检查每个 step 的新鲜度。

继续阅读核心概念或面向任务的操作指南