如何高效对NumPy数组切片?(有限差分法场景)
NumPy有限差分法(Heston模型)速度优化方案
你的代码速度慢的核心原因是过度依赖Python原生嵌套循环,没有利用NumPy的向量化运算优势——Python循环的开销远高于Matlab的向量化操作,这是两者速度差异的关键。以下是针对性的优化方案:
1. 初始化阶段向量化
将原来的双重循环初始化改为广播机制,一次性生成整个数组:
U = np.maximum(S[:, np.newaxis] - K, 0)
2. 边界条件设置向量化
用切片赋值替代循环设置边界:
# 左边界s=0、右边界s=Ns-1(v范围0到Nv-2) U[0, :Nv-1] = 0 U[-1, :Nv-1] = np.maximum(Smax - K, 0) # 上边界v=Nv-1(所有s) U[:, -1] = np.maximum(S - K, 0)
3. v=0层计算向量化
将单元素循环改为数组批量运算:
s_range = np.arange(1, Ns-1) derV = (u[s_range, 1] - u[s_range, 0]) / (V[1] - V[0]) derS = (u[s_range+1, 0] - u[s_range-1, 0]) / (S[s_range+1] - S[s_range-1]) U[s_range, 0] = u[s_range, 0] + dt * (-r*u[s_range, 0] + (r-q)*S[s_range]*derS + kappa*theta*derV)
4. 核心双重循环完全向量化
这是最耗时的部分,用切片运算彻底消除嵌套循环,一次性完成所有元素的导数计算和更新:
# 提取所需切片 u_center = u[1:-1, 1:-1] u_s_plus = u[2:, 1:-1] u_s_minus = u[:-2, 1:-1] u_v_plus = u[1:-1, 2:] u_v_minus = u[1:-1, :-2] u_sv_plusplus = u[2:, 2:] u_sv_minusminus = u[:-2, :-2] u_sv_minusplus = u[:-2, 2:] u_sv_plusminus = u[2:, :-2] # 计算各阶导数 derS = 0.5 * (u_s_plus - u_s_minus) / (S[2:] - S[:-2])[:, np.newaxis] derV = 0.5 * (u_v_plus - u_v_minus) / (V[2:] - V[:-2])[np.newaxis, :] derSS = (u_s_plus - 2*u_center + u_s_minus) / ((S[2:] - S[1:-1]) * (S[1:-1] - S[:-2]))[:, np.newaxis] derVV = (u_v_plus - 2*u_center + u_v_minus) / ((V[2:] - V[1:-1]) * (V[1:-1] - V[:-2]))[np.newaxis, :] derSV = (u_sv_plusplus + u_sv_minusminus - u_sv_minusplus - u_sv_plusminus) / (4 * (S[2:] - S[:-2])[:, np.newaxis] * (V[2:] - V[:-2])[np.newaxis, :]) # 计算各项系数 V_slice = V[1:-1][np.newaxis, :] S_slice = S[1:-1][:, np.newaxis] A = 0.5 * V_slice * (S_slice**2) * derSS B = rho * sigma * V_slice * S_slice * derSV C = 0.5 * (sigma**2) * V_slice * derVV D = 0.5 * (r - q) * S_slice * derS E = kappa * (theta - V_slice) * derV L = dt * (A + B + C + D + E - r*u_center) # 更新核心区域 U[1:-1, 1:-1] = u_center + L
优化后完整代码
import numpy as np from scipy import interpolate def heston_explicit_nonuniform(S, V, K, T, Ns, Nv, kappa, theta, rho, sigma, r, q, dt, dv, ds, Smax): # 初始化U,向量化替代双重循环 U = np.maximum(S[:, np.newaxis] - K, 0) for t in range(Nt - 1): # 边界条件向量化设置 U[0, :Nv-1] = 0 U[-1, :Nv-1] = np.maximum(Smax - K, 0) U[:, -1] = np.maximum(S - K, 0) u = np.copy(U) # v=0层计算向量化 s_range = np.arange(1, Ns-1) derV = (u[s_range, 1] - u[s_range, 0]) / (V[1] - V[0]) derS = (u[s_range+1, 0] - u[s_range-1, 0]) / (S[s_range+1] - S[s_range-1]) U[s_range, 0] = u[s_range, 0] + dt * (-r*u[s_range, 0] + (r-q)*S[s_range]*derS + kappa*theta*derV) u = np.copy(U) # 核心区域向量化计算 u_center = u[1:-1, 1:-1] u_s_plus = u[2:, 1:-1] u_s_minus = u[:-2, 1:-1] u_v_plus = u[1:-1, 2:] u_v_minus = u[1:-1, :-2] u_sv_plusplus = u[2:, 2:] u_sv_minusminus = u[:-2, :-2] u_sv_minusplus = u[:-2, 2:] u_sv_plusminus = u[2:, :-2] derS = 0.5 * (u_s_plus - u_s_minus) / (S[2:] - S[:-2])[:, np.newaxis] derV = 0.5 * (u_v_plus - u_v_minus) / (V[2:] - V[:-2])[np.newaxis, :] derSS = (u_s_plus - 2*u_center + u_s_minus) / ((S[2:] - S[1:-1]) * (S[1:-1] - S[:-2]))[:, np.newaxis] derVV = (u_v_plus - 2*u_center + u_v_minus) / ((V[2:] - V[1:-1]) * (V[1:-1] - V[:-2]))[np.newaxis, :] derSV = (u_sv_plusplus + u_sv_minusminus - u_sv_minusplus - u_sv_plusminus) / (4 * (S[2:] - S[:-2])[:, np.newaxis] * (V[2:] - V[:-2])[np.newaxis, :]) V_slice = V[1:-1][np.newaxis, :] S_slice = S[1:-1][:, np.newaxis] A = 0.5 * V_slice * (S_slice**2) * derSS B = rho * sigma * V_slice * S_slice * derSV C = 0.5 * (sigma**2) * V_slice * derVV D = 0.5 * (r - q) * S_slice * derS E = kappa * (theta - V_slice) * derV L = dt * (A + B + C + D + E - r*u_center) U[1:-1, 1:-1] = u_center + L if t % 100 == 0: print(t) return U # 参数设置 Vmax=0.5 Vmin=0 Smin=0 T=0.15 K=100. Smax = 2*K r= 0.02 q=0.05 rho=-0.9 v0=0.05 sigma=0.3 kappa=1.5 theta=0.04 Nv=39 Ns=79 Nt = 3000 dt = T/Nt Vi = np.linspace(0, Vmax, Nv) Si = np.linspace(0, Smax, Ns) dv = (Vmax-Vmin)/Nv ds = (Smax-Smin)/Ns # 调用优化后的函数 call = heston_explicit_nonuniform(Si, Vi, K, T, Ns, Nv, kappa, theta, rho, sigma, r, q, dt, dv, ds, Smax) data = interpolate.RectBivariateSpline(Si, Vi, call) z_new = data(101.52, 0.05412) print('FDM: ', z_new[0,0])
优化后代码完全利用NumPy的C底层加速,计算速度会接近Matlab的水平。若需进一步提速,可考虑用Numba对核心函数进行JIT编译。
内容的提问来源于stack exchange,提问作者THATS MY QUANT MY QUANTITATIVE
相关产品推荐
相关产品推荐

