注意
跳转至页面底部下载完整示例代码。
torch.vmap#
本教程介绍了 PyTorch 操作的自动向量化工具 torch.vmap。torch.vmap 目前是一个原型功能,尚无法处理许多用例;然而,我们希望能收集相关的用例以优化设计。如果您正在考虑使用 torch.vmap,或者认为它在某些方面非常有前景,请通过 pytorch/pytorch#42368 联系我们。
那么,什么是 vmap?#
vmap 是一个高阶函数。它接收一个函数 func,并返回一个新的函数,该函数将 func 映射到输入的某个维度上。它的设计深受 JAX 的 vmap 启发。
从语义上讲,vmap 将“映射”(map)操作下推到 func 调用的 PyTorch 操作中,从而有效地对这些操作进行向量化。
# NB: vmap is only available on nightly builds of PyTorch.
# You can download one at pytorch.org if you're interested in testing it out.
vmap 的第一个用途是简化代码中批处理维度(batch dimension)的处理。我们可以编写一个处理单个示例的函数 func,然后使用 vmap(func) 将其提升为一个能够处理批量示例的函数。然而,func 受到许多限制:
它必须是函数式的(不能在内部修改 Python 数据结构),但就地(in-place)的 PyTorch 操作除外。
批量示例必须以张量(Tensors)形式提供。这意味着 vmap 开箱即用时无法处理变长序列。
使用 vmap 的一个例子是计算批量点积。PyTorch 没有提供批量的 torch.dot API;与其在文档中苦苦搜索无果,不如使用 vmap 构建一个新函数:
vmap 有助于隐藏批处理维度,从而带来更简洁的模型编写体验。
# Note that model doesn't work with a batch of feature vectors because
# torch.dot must take 1D tensors. It's pretty easy to rewrite this
# to use `torch.matmul` instead, but if we didn't want to do that or if
# the code is more complicated (e.g., does some advanced indexing
# shenanigins), we can simply call `vmap`. `vmap` batches over ALL
# inputs, unless otherwise specified (with the in_dims argument,
# please see the documentation for more details).
vmap 还可以帮助对以前难以或无法进行批处理的计算进行向量化。这引出了我们的第二个用例:批量梯度计算。
PyTorch 自动微分引擎计算的是 vjp(向量-雅可比乘积)。利用 vmap,我们可以计算(批量向量)- 雅可比乘积。
一个例子是计算完整的雅可比矩阵(这也可应用于计算完整的 Hessian 矩阵)。计算某个函数 f: R^N -> R^N 的完整雅可比矩阵通常需要调用 N 次 autograd.grad,每一行雅可比矩阵调用一次。
# Setup
# Sequential approach
# Using `vmap`, we can vectorize the whole computation, computing the
# Jacobian in a single call to `autograd.grad`.
vmap 的第三个主要用例是计算逐样本梯度(per-sample-gradients)。这是 vmap 原型目前无法高效处理的部分。我们尚不确定计算逐样本梯度的 API 应该是什么样的,如果您有任何想法,请在 pytorch/pytorch#7786 中发表评论。
# The following doesn't actually work in the vmap prototype. But it
# could be an API for computing per-sample-gradients.
# batch_of_samples = torch.randn(64, 5)
# vmap(grad_sample)(batch_of_samples)
# %%%%%%RUNNABLE_CODE_REMOVED%%%%%%
脚本总运行时间:(0 分 0.002 秒)