
如何用 torch.func 的 vmap 与 grad 计算 per-sample 梯度【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch训练时通常只能拿到整个 batch 平均后的一个梯度但有时需要每个样本各自的梯度per-sample gradient例如分析单个样本对损失的贡献。PyTorch 自带的 autograd 无法高效地做到这一点torch.func库原名 functorch文档注明目前处于 beta 阶段API 可能随反馈调整提供了grad和vmap两个可组合的函数变换其中vmap(grad(f))正是文档给出的 per-sample 梯度计算方式。本文基于 torch.func 快速入门 中的完整示例给出可运行的操作路径、结果验证方式和vmap的限制说明。grad 与 vmap 各自做什么grad(func)返回一个新函数用于计算func的梯度。它假设func返回一个单元素 Tensor并且默认计算func输出对第一个输入的梯度——这正好适合对权重求损失梯度的场景import torch from torch.func import grad x torch.randn([]) cos_x grad(lambda x: torch.sin(x))(x) assert torch.allclose(cos_x, x.cos()) # 高阶梯度 neg_sin_x grad(grad(lambda x: torch.sin(x)))(x) assert torch.allclose(neg_sin_x, -x.sin())vmap(func)返回一个新函数把func映射到每个输入 Tensor 的某个维度上默认第 0 维等价于给func里所有 Tensor 运算增加一个维度。文档给出的心智模型是对纯函数vmap(f)(x)等价于torch.stack([f(x_i) for x_i in x.unbind(0)])也就是说你可以先写一个只处理单个样本的函数再用vmap把它提升为处理整个 batch 的版本import torch from torch.func import vmap batch_size, feature_size 3, 5 weights torch.randn(feature_size, requires_gradTrue) def model(feature_vec): # 非常简单的线性模型带激活 assert feature_vec.dim() 1 return feature_vec.dot(weights).relu() examples torch.randn(batch_size, feature_size) result vmap(model)(examples)组合 vmap 与 grad 计算 per-sample 梯度下面是文档中计算 per-sample 梯度的完整示例。核心是vmap(grad(compute_loss), in_dims(None, 0, 0))compute_loss接收三个输入in_dims指定每个输入被 vmapped 的维度——weights是所有样本共享的参数传None表示不参与映射examples和targets按第 0 维batch 维逐个映射。文档原文示例未写出grad的导入行此处补全以便直接运行。import torch from torch.func import vmap, grad batch_size, feature_size 3, 5 def model(weights, feature_vec): # 非常简单的线性模型带激活 assert feature_vec.dim() 1 return feature_vec.dot(weights).relu() def compute_loss(weights, example, target): y model(weights, example) return ((y - target) ** 2).mean() # MSELoss weights torch.randn(feature_size, requires_gradTrue) examples torch.randn(batch_size, feature_size) targets torch.randn(batch_size) inputs (weights, examples, targets) grad_weight_per_example vmap(grad(compute_loss), in_dims(None, 0, 0))(*inputs)注意两点compute_loss返回单元素 Tensor((y - target) ** 2).mean()是标量满足grad的输入要求grad默认对第一个输入weights求梯度。按文档中 vmap 的语义grad(compute_loss)对每个样本各返回一个形状为(feature_size,)的梯度vmap 再把 3 个结果堆叠起来因此grad_weight_per_example的形状为(3, 5)即(batch_size, feature_size)。验证结果文档对 vmap 给出的等价定义纯函数下vmap(f)(x)等于逐样本调用后torch.stack直接构成一条验证路径把 vmap 的结果和逐样本循环调用grad的结果比对。文档自身的示例如上面的sin/cos例也采用assert torch.allclose作为判断方式manual_grads torch.stack([ grad(compute_loss)(weights, example, target) for example, target in zip(examples.unbind(0), targets.unbind(0)) ]) assert torch.allclose(grad_weight_per_example, manual_grads)若断言通过说明vmap(grad(compute_loss))的结果与逐样本计算一致循环版本只用于核对实际使用以 vmap 版本为准。vmap 带来的限制与文档给出的处理方式文档明确说明vmap是torch.func中限制最多的变换grad、vjp、jvp等梯度变换没有这些限制因此下面的限制都源自 vmap 这一层详见 UX 限制文档纯函数要求被变换的函数不应给全局变量赋值所有输出都必须 return 出来。需要额外返回中间量时改用grad(f, has_auxTrue)文档给出了改写示例。不支持的 in-place 操作把更多元素写入更少元素的 in-place 运算会报错如regular.add_(batched)。文档给出的通用修复是把工厂函数换成new_*等价形式例如把torch.zeros(...)换成vec.new_zeros(...)、torch.empty换成Tensor.new_empty使中间结果本身带上 batch 维。out关键字参数vmap 内不支持 PyTorch 操作的out参数遇到会报错。数据依赖的控制流if/while/for 的条件若是被 vmapped 的 Tensor 会报错条件不依赖 vmapped Tensor 值的控制流可以正常使用。.item()调用对 vmapped Tensor 调用.item()不支持需要改写代码去掉该调用。动态形状操作torch.nonzero、torch.is_nonzero等对不同样本可能返回不同形状的操作不支持因为输出无法堆叠成单一 Tensor。随机操作若函数内调用随机算子vmap 默认处于error模式会报错需要显式传randomnessdifferentbatch 内各元素取不同随机值或randomnesssamebatch 内取同一随机值。另一个常见场景是模型里带 BatchNorm它对running_mean/running_var的 in-place 更新与 vmap 冲突。Patch Batch Norm 文档 给出四种处理方式把BatchNorm2d换成GroupNorm需满足C % G 0取C G为兜底自建模块时设置BatchNorm2d(..., track_running_statsFalse)对 torchvision 模型如 resnet、regnet通过norm_layer参数注入 GroupNorm 或关闭 running stats 的 BatchNorm或者对模块调用replace_all_batch_norm_modules_(net)来自torch.func就地关闭 running stats。另外模型处于 eval 模式时 running stats 不更新model.eval()下vmap(model)(x)可以直接支持。这些选项的前提是任务不依赖 running stats 的更新。对 nn.Module 求参数梯度可选路径上面的示例把参数作为显式输入传入。如果模型是nn.ModuleAPI 参考 说明可以变换调用该模块的函数当变换需要对模块的参数进行时用torch.func.functional_call把参数变成函数的显式输入文档中的写法是model torch.nn.Linear(3, 3) def f(params, x): return torch.func.functional_call(model, params, x) x torch.randn(3) jacobian jacrev(f)(dict(model.named_parameters()), x)同样的模式可以用于 per-sample 梯度把f改为返回单元素损失再用vmap(grad(f), in_dims(None, 0, 0))套用本文的主路径。限制与延伸torch.func目前处于 beta 阶段功能整体可用但 API 可能调整且对 PyTorch 算子的覆盖不完整。变换遇到不支持的写法时会给出报错信息完整的变换 APIvmap、grad、grad_and_value、vjp、jvp、jacrev、jacfwd、hessian等见 func.api.mdvmap 限制的全部细节与示例见 func.ux_limitations.md。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考