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的输入张量。
- 放弃用Trace优化包含
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,但性能提升可能有限。
- 优先用JIT Script替代
优化方案选择建议
- 优先使用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
相关产品推荐
相关产品推荐

