You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch中使用Functorch计算PINN梯度报错的解决咨询

解决Functorch计算PINN批量导数的RuntimeError问题

问题根源

Functorch的grad函数默认要求被求导的函数返回标量张量,但PINN场景中我们通常输入批量的(t,x,y)三维张量,网络输出的是批量温度u(形状如[batch_size,1]或[batch_size]),这直接触发了"返回非标量张量"的RuntimeError。


解决方案1:vmap + grad 批量求一阶导数

先定义处理单样本的函数(确保返回标量),再用vmap将其映射到批量输入,配合grad实现批量自动微分:

import torch
import functorch as ft

# 定义PINN模型
class PINN(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.layers = torch.nn.Sequential(
            torch.nn.Linear(3, 64),
            torch.nn.Tanh(),
            torch.nn.Linear(64, 64),
            torch.nn.Tanh(),
            torch.nn.Linear(64, 1)
        )
    
    def forward(self, x):
        # 单样本输入为[3],批量输入为[batch_size,3]
        return self.layers(x).squeeze(-1)  # 单样本返回标量,批量返回[batch_size]

model = PINN()

# 单样本前向函数:输入[3]张量,返回标量u
def single_sample_forward(txy):
    return model(txy)

# 用vmap批量化grad操作,得到批量一阶导数计算函数
batch_grad = ft.vmap(ft.grad(single_sample_forward))

# 测试批量输入
batch_txy = torch.randn(100, 3, requires_grad=True)
du_dinput = batch_grad(batch_txy)
# du_dinput形状为[100,3],对应每个样本的du/dt、du/dx、du/dy

解决方案2:vmap + jacrev/hessian 求高阶导数

如果需要计算二阶导数(如PINN中的拉普拉斯项),可以用jacrev(计算雅可比矩阵)配合vmap,再嵌套求导得到海森矩阵:

# 单样本二阶导数计算函数
def compute_single_hessian(txy):
    # 先求一阶雅可比矩阵([3])
    first_jac = ft.jacrev(single_sample_forward)(txy)
    # 再对一阶导数求雅可比,得到二阶导数矩阵([3,3])
    return ft.jacrev(lambda x: first_jac)(txy)

# 批量计算二阶导数
batch_hessian = ft.vmap(compute_single_hessian)(batch_txy)
# batch_hessian形状为[100,3,3],对应每个样本的二阶导数矩阵

关于torch.autograd.grad的适配问题

torch.autograd.grad需要手动处理批量逻辑:要么循环遍历每个样本求导(效率极低),要么用torch.autograd.functional.jacobian但需手动调整输入维度,相比之下Functorch的vmap可以更简洁地实现批量自动微分,无需修改核心求导逻辑。

内容的提问来源于stack exchange,提问作者Bo van Hasselt

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.19 07:02:40