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

PyTorch中PINN求解器导数高效计算及JIT/编译报错问题

物理信息神经网络(PINNs)PyTorch优化问题及解决方案

问题背景

正在实现物理信息神经网络(PINNs),需计算模型输出对输入的高阶偏导数以构建热方程、伯格斯方程等PDE残差。通过torch.autograd.grad实现了简洁的导数计算函数diff,但尝试torch.jit.script、torch.jit.trace和torch.compile优化训练速度时遇到以下报错:

  • JIT Trace报错:element 0 of tensors does not require grad and does not have a grad_fn
  • torch.compile报错:Unsupported: inlining disallowed: <function grad at 0x7f36642888b0>

有TensorFlow实现经验,作为PyTorch新手,需明确优化方案选择、是否尝试TorchDynamo,以及获取报错解决方法和设计建议。

相关代码

自定义导数计算函数

def diff(y, xs):
    grad = y
    ones = torch.ones_like(y)
    for x in xs:
        grad = torch.autograd.grad(grad, x, grad_outputs=ones, create_graph=True)[0]
    return grad

JIT Script版导数函数

@torch.jit.script
def D(y: torch.Tensor, xs: List[torch.Tensor]) -> torch.Tensor:
    func: torch.Tensor = y
    ones: List[Optional[torch.Tensor]] = [torch.ones_like(y)]
    for x in xs:
        grad = torch.autograd.grad(
            [func],
            [x],
            grad_outputs=ones,
            allow_unused=True,
            create_graph=True)[0]
        func = grad if grad is not None else torch.zeros_like(x)
    return func

测试用神经网络

class NeuralNet(nn.Module):
    def __init__(self, n_input):
        super().__init__()
        self.linear_sigmoid_stack = nn.Sequential(
            nn.Linear(n_input, 32),
            nn.Sigmoid(),
            nn.Linear(32, 32),
            nn.Sigmoid(),
            nn.Linear(32, 1)
        )
    def forward(self, x):
        x = torch.hstack(x)
        return self.linear_sigmoid_stack(x)

报错原因及解决方案

1. JIT Trace报错

JIT Trace依赖记录静态计算图实现优化,但torch.autograd.grad是动态梯度计算操作,Trace无法捕获其动态逻辑;若输入张量未开启requires_grad=True,会直接触发该错误。

  • 解决方法:
    • 放弃用Trace优化包含autograd.grad的代码,改用torch.jit.script;
    • 确保传入导数函数的所有xs张量都设置requires_grad=True;
    • 若需Trace模型,将导数计算逻辑整合到模型forward方法中,且Trace时传入带requires_grad=True的输入张量。

2. torch.compile报错

当前torch.compile对torch.autograd.grad(尤其是create_graph=True场景)的支持有限,Dynamo难以追踪这类动态梯度计算的反向图结构,导致无法内联优化。

  • 解决方法:
    • 优先用JIT Script替代torch.compile处理PINNs的导数计算逻辑;
    • 用torch.func系列API(如jacrev、hessian)重构导数计算,该API更兼容torch.compile;
    • 若坚持用torch.compile,可尝试设置mode="reduce-overhead"或fullgraph=False,但性能提升可能有限。

优化方案选择建议

  • 优先使用torch.jit.script:Script能更好地处理动态循环和autograd.grad操作,只要保证类型标注正确(如你已实现的D函数),即可正常编译,能带来稳定的性能提升;
  • TorchDynamo(torch.compile):目前对PINNs高阶导数计算的支持不完善,建议等待PyTorch版本更新,或改用torch.func重构代码后再尝试;
  • 避免JIT Trace:Trace仅适合静态计算图模型,PINNs的导数计算是动态逻辑,Trace无法完整捕获,易引发错误。

额外设计建议

  • 用torch.func重构导数计算:torch.func.jacrev可直接计算雅可比矩阵,高阶导数可通过嵌套调用实现,代码更简洁且兼容现代PyTorch优化工具,示例:
import torch.func as tf

def compute_derivative(model, xs):
    # 定义输入到输出的函数
    def model_func(*input_tensors):
        x = torch.hstack(input_tensors)
        return model(x)
    # 计算一阶导数(雅可比)
    jacobian = tf.jacrev(model_func)(*xs)
    # 计算二阶导数可嵌套tf.jacrev或用tf.hessian
    return jacobian
  • 训练时确保所有输入张量开启requires_grad=True,避免梯度计算时出现无梯度张量的问题;
  • 将PDE残差的完整计算逻辑封装为单个函数,再用JIT Script编译该函数,而非单独编译导数函数,优化效果更显著。

内容的提问来源于stack exchange,提问作者Salih Taşdelen

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 05:42:52