使用Scipy RectBivariateSpline插值时,如何用MyGrad正确追踪梯度?
问题:Scipy双变量样条插值与MyGrad自动微分的梯度追踪失效
我正在进行一个项目,需使用scipy.interpolate.RectBivariateSpline对焓值进行插值,之后通过mygrad实现自动微分。但我遇到了插值过程完全无法追踪梯度的问题,以下是简化后的代码:
import numpy as np from scipy.interpolate import RectBivariateSpline import CoolProp.CoolProp as CP import mygrad as mg from mygrad import tensor # Define the refrigerant refrigerant = 'R134a' # Constant temperature (e.g., 20°C) T = 20 + 273.15 # Convert to Kelvin # Get saturation pressures P_sat = CP.PropsSI('P', 'T', T, 'Q', 0, refrigerant) # Define a pressure range around the saturation pressure P_min = P_sat * 0.5 P_max = P_sat * 1.5 P_values = np.linspace(P_min, P_max, 100) # Define a temperature range around the constant temperature T_min = T - 10 T_max = T + 10 T_values = np.linspace(T_min, T_max, 100) # Generate enthalpy data h_values = [] for P in P_values: h_row = [] for T in T_values: try: h = CP.PropsSI('H', 'P', P, 'T', T, refrigerant) h_row.append(h) except: h_row.append(np.nan) h_values.append(h_row) # Convert lists to arrays h_values = np.array(h_values) # Fit spline for enthalpy h_spline = RectBivariateSpline(P_values, T_values, h_values) # Function to interpolate enthalpy def h_interp(P, T): return tensor(h_spline.ev(P, T)) # Example function using the interpolated enthalpy with AD def example_function(P): h = h_interp(P, T) result = h**2 # Example calculation return result # Define a pressure value for testing P_test = tensor(P_sat, ) # Compute the example function and its gradient result = example_function(P_test) result.backward() # Print the result and the gradient print(f"Result: {result.item()}") print(f"Gradient: {P_test.grad}")
请问这是RectBivariateSpline或mygrad的问题吗?其他自动微分库能否解决该问题?我是否应该采用插值以外的方法?
回答
1. 问题根源:计算图断开
这既不是RectBivariateSpline的问题,也不是mygrad的问题。核心原因是:
- Scipy的插值函数基于纯NumPy实现,没有为自动微分(AD)框架提供梯度追踪接口。
- 当你将
h_spline.ev(P, T)的结果转换为mygrad tensor时,仅将数值结果传入了AD计算图,插值的整个计算过程完全脱离了框架的追踪范围,反向传播时自然无法计算输入P到输出h的梯度。
2. 其他自动微分库的解决方案
主流AD库(如PyTorch、TensorFlow)同样无法直接对Scipy插值函数追踪梯度,但有两种可行的处理方式:
- 手动实现可微分插值:用框架的原生张量操作重写双变量样条插值逻辑(比如基于B样条基函数的展开),让整个插值过程落在AD计算图内。
- 使用JAX的封装接口:JAX对部分Scipy函数做了可微分封装,
jax.scipy.interpolate.RectBivariateSpline支持自动微分,可以直接替换使用。 - 数值差分近似梯度:如果精度要求不高,可以用AD库的数值微分工具(比如PyTorch的
torch.autograd.functional.jacobian配合数值差分)来近似梯度。
3. 插值以外的替代方案
如果核心需求是获得焓值的可微分计算,完全可以绕过插值:
- 直接使用CoolProp的导数接口:CoolProp本身提供了计算焓对压力/温度偏导数的函数,比如
CP.PropsSI('dHdP', 'P', P, 'T', T, refrigerant),可以直接用这些导数配合AD框架,避免插值带来的误差和梯度问题。 - 改用可微分热力学库:如果项目允许,使用基于PyTorch/TensorFlow实现的可微分热力学库,这类库的所有操作都能被AD框架直接追踪,无需额外处理。
内容的提问来源于stack exchange,提问作者Andreas Schuldei
相关产品推荐
相关产品推荐

