- 人工智能
- 机器学习
- 深度学习
- 概率编程
【免费下载链接】pyro
Deep universal probabilistic programming with Python and PyTorch
导读
pyro.settings是 Pyro 深度通用概率编程框架(基于 Python 与 PyTorch)中一套轻量而统一的全局配置机制:它通过"别名(alias)→ 真实变量路径(模块 + 点分隔深度名)"的注册映射,让用户可以用一致的 API 读写散布在不同模块中的全局常量与类属性。本文以 settings.rst 对应的 pyro/settings.py 文档字符串为骨架,结合仓库源码逐一拆解get、set、context、register四个接口的用法、内置设置清单及其底层实现原理,帮助你安全地在推理、采样与数值运算中切换验证开关、数值容差与近似阈值。
一、settings 是什么:别名驱动的全局变量注册表
Pyro 的全局设置并不集中存放在某一个字典里,而是散落在多个模块的模块级常量(如pyro/ops/tensor_utils.py中的CHOLESKY_RELATIVE_JITTER)和类属性(如pyro/distributions/torch.py中Binomial.approx_sample_thresh)中。pyro.settings通过一个全局注册表_REGISTRY(见 pyro/settings.py)把这些散落的变量统一映射为短小易记的别名:
_REGISTRY: Dict[str, Tuple[str, str, Optional[Callable]]] = {}注册表将每个别名映射为一个三元组:(modulename, deepname, validator),其中deepname可以是模块常量(如MY_CONSTANT),也可以是带点的类属性路径(如Binomial.approx_sample_thresh)。读取时先import_module(module),再沿deepname逐级getattr直达目标;写入时则沿路径逐级定位并用setattr覆盖最后一个名字。
所有设置在模块导入时即完成注册,因此pyro.settings.get()无参数调用即可返回当前全部设置的有序字典,而读取模块 docstring(即__doc__)时会动态渲染出"Default Settings"清单(见 pyro/settings.py)。
二、四个核心接口的完整用法
原文档(settings.rst)通过automodule:: pyro.settings引用上述 docstring,其中给出了完整的示例代码,下面逐一展开并补充源码依据。
1.get(alias=None):查询一个或全部设置
print(pyro.settings.get()) # 打印所有设置(dict) print(pyro.settings.get("cholesky_relative_jitter")) # 打印单个设置实现位于 pyro/settings.py:当alias is None时,对注册表中所有别名按字典序逐个get并组装成字典返回;传入别名时则按注册三元组定位真实变量并返回当前值。若别名未注册,会抛出KeyError(参见 tests/test_settings.py 的测试用例)。
2.set(**kwargs):一次性设置一个或多个配置
pyro.settings.set(cholesky_relative_jitter=0.5) # 设置单个 pyro.settings.set(**my_settings) # 批量设置实现见 pyro/settings.py。每个alias=value对都会先查注册表;若该设置注册了 validator,则先调用validator(value)做校验(校验失败即抛错,写入不生效),再沿模块路径写入目标变量。这意味着set是"校验 + 赋值"一体化的操作。
3.context(**kwargs):上下文管理器与装饰器双重形态
# 作为上下文管理器:临时覆盖,退出自动恢复 with pyro.settings.context(cholesky_relative_jitter=0.5): my_function() # 作为装饰器 fn = pyro.settings.context(cholesky_relative_jitter=0.5)(my_function) fn()实现见 pyro/settings.py:进入时先快照旧值并set(**kwargs),yield后通过finally恢复旧值。该机制在仓库中广泛用于"只在某段代码内放宽容差"的场景,例如流行病学模块 pyro/contrib/epidemiology/compartmental.py 在fit_mcmc、predict等方法上以装饰器形式临时设置approx_log_prob_tol=0.1与approx_sample_thresh=10000,从而在不污染全局状态的前提下加速大规模后验计算。测试用例 tests/test_settings.py 同时验证了上下文与装饰器两种用法的值隔离与恢复。
4.register(alias, modulename, deepname, validator=None):注册新设置
# 直接声明式注册 pyro.settings.register( "binomial_approx_sample_thresh", # alias "pyro.distributions.torch", # module "Binomial.approx_sample_thresh", # deep name ) # 以装饰器形式附加自定义校验器(每次 set 时被调用) @pyro.settings.register( "binomial_approx_sample_thresh", # alias "pyro.distributions.torch", # module "Binomial.approx_sample_thresh", # deep name ) def validate_thresh(thresh): assert isinstance(thresh, float) assert thresh > 0实现见 pyro/settings.py,参数约定如下:
| 参数 | 类型 | 说明 |
|---|---|---|
alias | str | 合法的 Python 标识符,推荐小写蛇形命名(如my_setting) |
modulename | str | 设置声明所在的模块名,通常是__name__ |
deepname | str | 点分隔的名称路径:模块常量用MY_CONSTANT,类属性用MyClass.my_attribute |
validator | callable | None | 可选校验器:输入一个值,可抛出校验错误,返回None |
register还具备三个附带行为:注册后立即将默认值刷新进模块 docstring(Default Settings区块);当validator is None时返回functools.partial支持装饰器二次调用;当提供 validator 时立即用当前值执行一次校验(见 pyro/settings.py),确保存量值也合法。
最佳实践:官方要求"设置在定义它的模块中声明"。仓库中的注册均遵循此约定,例如 pyro/distributions/torch.py、pyro/ops/tensor_utils.py、pyro/nn/module.py。
三、内置设置清单与默认值
pyro.settings在导入时自动注册以下设置(均为模块加载时声明,默认值以当前仓库代码为准):
| 别名 | 指向的真实变量 | 默认值 | 作用 |
|---|---|---|---|
validate_distributions_pyro | pyro.distributions.util._VALIDATION_ENABLED | True(跟随__debug__) | Pyro 自定义分布的验证开关,见 pyro/distributions/util.py |
validate_distributions_torch | torch.distributions.distribution.Distribution._validate_args | True(跟随__debug__) | Torch 底层分布参数校验开关 |
validate_poutine | pyro.poutine.util._VALIDATION_ENABLED | True(跟随__debug__) | 概率编程效应处理器(poutine)的验证开关,见 pyro/poutine/util.py |
validate_infer | pyro.infer.util._VALIDATION_ENABLED | True(跟随__debug__) | 推理算法(ELBO 等)的验证开关,见 pyro/infer/util.py |
cholesky_relative_jitter | pyro.ops.tensor_utils.CHOLESKY_RELATIVE_JITTER | 4.0(以finfo.eps为单位) | Cholesky 分解的自适应 jitter 强度,见 pyro/ops/tensor_utils.py |
binomial_approx_sample_thresh | pyro.distributions.torch.Binomial.approx_sample_thresh | math.inf | 二项分布采样切换为"截断泊松近似"的阈值(total_count超过该值启用),见 pyro/distributions/torch.py |
binomial_approx_log_prob_tol | pyro.distributions.torch.Binomial.approx_log_prob_tol | 0.0 | .log_prob()使用移位的 Stirling 近似计算 Beta 函数的容差(推荐 0.1~0.01),见 pyro/distributions/torch.py |
module_local_params | pyro.nn.module._MODULE_LOCAL_PARAMS | False | 是否启用 PyroModule 局部参数模式,见 pyro/nn/module.py |
其中四个验证类设置均初始化为 Python 全局__debug__,这与 pyro/primitives.py 中enable_validation的语义一致:默认开启验证,但以python -O优化模式运行时会自动关闭。tests/test_settings.py中test_settings(tests/test_settings.py)断言了这四类验证默认值均为True。
四、源码级原理:设置如何影响真实运行路径
4.1cholesky_relative_jitter与数值稳定性
该设置直接控制 pyro/ops/tensor_utils.py 中safe_cholesky的行为:当 jitter 非零时,先取矩阵绝对值的行最大值,计算jitter = CHOLESKY_RELATIVE_JITTER * finfo.eps * x_max加到对角线上,再调用torch.linalg.cholesky。对于半正定(PSD)但可能近奇异的协方差矩阵,适当调大 jitter(如pyro.settings.set(cholesky_relative_jitter=10.0))能显著提升分解稳定性,代价是引入轻微偏差;而 GP、高斯 HMM、MCMC 等大量依赖 Cholesky 的算子在数值不稳时可借助它收敛。
4.2 二项分布近似开关
binomial_approx_sample_thresh控制 pyro/distributions/torch.py 的Binomial.sample:当total_count > approx_sample_thresh时改用矩匹配的截断 Poisson 近似采样(torch.no_grad()下进行),适合超大规模计数场景;binomial_approx_log_prob_tol则切换log_prob到低开销的 Stirling 近似路径。两者都带 validator(要求thresh > 0、tol >= 0),保证非法值在set阶段即被拦截。流行病学模块的set_approx_sample_thresh/set_approx_log_prob_tol上下文管理器(pyro/contrib/epidemiology/distributions.py)正是围绕同一套类属性封装的安全临时覆盖工具。
4.3 验证开关与enable_validation
pyro.primitives.enable_validation(pyro/primitives.py)会同时联动dist、infer、poutine三套验证;对应地,settings.set(validate_distributions_torch=...)可直接改写 Torch 底层的Distribution._validate_args。推荐在开发阶段保持验证开启以便尽早暴露 NaN、非法参数等问题,在性能敏感的 jit 编译或大批量生产推理时再临时关闭(Pyro 官方在 pyro/primitives.py 也注明 jit 编译期间验证会被暂时禁用)。
4.4 设置读取点示例:module_local_params
pyro/nn/module.py 中_pyro_set_supermodule与__call__会读取pyro.settings.get("module_local_params")与validate_poutine,用于决定是否触发局部参数使用检查。这展示了 settings 的另一特征:同一设置在多个代码路径中被消费,通过别名统一入口即可全局生效。
五、扩展:注册你自己的全局设置
在业务代码(或自定义插件模块)中可按以下模板注册新设置:
import pyro # 声明式注册(指向模块常量) pyro.settings.register("my_setting", __name__, "MY_SETTING") # 带校验器的注册(每次 set 都会调用) @pyro.settings.register("my_tolerance", __name__, "MY_TOLERANCE") def _validate_tolerance(value): assert isinstance(value, float) assert 0 < value <= 1.0 # 使用 pyro.settings.set(my_tolerance=0.05) with pyro.settings.context(my_tolerance=0.5): run_experiment()注意:register的 docstring 校验要求alias必须是合法 Python 标识符(alias.isidentifier());validator 若被用作普通装饰器,须返回None。注册完成后,新设置会自动出现在pyro.settings.get()的返回字典与模块 docstring 的默认值清单中。
六、测试验证与使用注意事项
- 单元测试 tests/test_settings.py 完整覆盖了注册、
get/set、校验器拦截非法值(set(test_setting=-0.1)抛AssertionError)、上下文与装饰器隔离恢复四条路径,可作为自定义设置的回归参考。 - 批量读取与写入均为 O(n) 遍历注册表,调用频率极高时可先用
get()缓存结果。 context的恢复基于"进入前快照",支持嵌套使用,但同一别名嵌套覆盖时遵循标准的 LIFO 恢复语义。- 修改设置是全局副作用,多线程 / 多进程训练场景下建议优先使用
context限定作用域,避免跨任务串扰。 - 若设置未注册即调用
get/set,将触发KeyError,请在模块导入期(而非运行时)完成注册,与 pyro/settings.py 建议的声明位置保持一致。
通过本文介绍的四个接口与内置设置表,你可以在不修改任何源码的前提下,以统一、可校验、可临时恢复的方式调优 Pyro 的验证级别、Cholesky 数值稳定性、二项近似阈值与模块参数行为,从而在开发调试与生产性能之间灵活切换。
- 人工智能
- 机器学习
- 深度学习
- 概率编程
【免费下载链接】pyro
Deep universal probabilistic programming with Python and PyTorch
相关推荐
HumHub设置管理系统:全局配置、模块设置与容器级配置完整指南
HumHub设置管理系统:全局配置、模块设置与容器级配置完整指南 HumHub作为一款优秀的企业社交网络平台,其强大的设置管理系统让管理员能够轻松管理整个平台的
后端社交企业应用内容协同JAX 配置系统完全指南:三种设置方式与完整配置项解析
JAX 配置系统完全指南:三种设置方式与完整配置项解析 JAX 提供了统一的配置系统,用以控制从数值精度、后端平台选择到调试检查等各类运行时行为。本文以仓库文档
人工智能机器学习深度学习编译器高性能计算Dapper.SimpleCRUD属性详解:[Table]、[Key]、[Column]等注解完全指南
Dapper.SimpleCRUD属性详解: Table 、 Key 、 Column 等注解完全指南 Dapper.SimpleCRUD是一款为Dapper提
后端
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考