numpy.apply_along_axis传参异常:t0参数重复赋值问题
修复
numpy.apply_along_axis传参错误的解决方案 错误根源
numpy.apply_along_axis的传参逻辑是:将数组指定轴上的每个子数组作为第一个位置参数传递给目标函数。如果你的sde_euler_1函数定义中t0是第一个位置参数,而调用apply_along_axis时又通过关键字参数传递t0,就会导致t0被重复赋值(一次是子数组,一次是你指定的关键字参数),从而抛出TypeError: sde_euler_1() got multiple values for argument 't0'。
解决方案
方案1:写参数适配的包装函数
如果不想修改原sde_euler_1的参数顺序,写一个包装函数调整参数传递顺序:
def sde_wrapper(dW, t0, t_end, dt, x0, mu, sigma): # 把dW作为最后一个参数传给原sde_euler_1 return sde_euler_1(t0, t_end, dt, x0, mu, sigma, dW)
在蒙特卡洛模拟函数中调用apply_along_axis时,将额外参数以位置参数形式传递:
def mc_sim(N, t0, t_end, dt, x0, mu, sigma): n_steps = int((t_end - t0) / dt) dW_samples = np.random.normal(0, np.sqrt(dt), (N, n_steps)) # 传递额外参数用位置参数,不要用关键字参数 results = np.apply_along_axis(sde_wrapper, 1, dW_samples, t0, t_end, dt, x0, mu, sigma) return results
方案2:调整原函数的参数顺序
如果可以修改sde_euler_1,直接把dW(即apply_along_axis要处理的数组元素)放在参数列表第一位:
def sde_euler_1(dW, t0, t_end, dt, x0, mu, sigma): # 原实现逻辑,调整dW的使用位置即可 n_steps = len(dW) t = np.linspace(t0, t_end, n_steps + 1) x = np.zeros(n_steps + 1) x[0] = x0 for i in range(n_steps): x[i+1] = x[i] + mu(x[i], t[i]) * dt + sigma(x[i], t[i]) * dW[i] return x
之后直接调用apply_along_axis并传递额外位置参数:
def mc_sim(N, t0, t_end, dt, x0, mu, sigma): n_steps = int((t_end - t0) / dt) dW_samples = np.random.normal(0, np.sqrt(dt), (N, n_steps)) results = np.apply_along_axis(sde_euler_1, 1, dW_samples, t0, t_end, dt, x0, mu, sigma) return results
方案3:向量化实现(更高效率)
如果追求极致性能,完全可以抛弃apply_along_axis,用numpy向量化操作实现,避免Python级别的循环开销:
def sde_euler_vectorized(t0, t_end, dt, x0, mu, sigma, dW_samples): n_steps = dW_samples.shape[1] t = np.linspace(t0, t_end, n_steps + 1) # 初始化所有模拟路径的数组 x = np.full((dW_samples.shape[0], n_steps + 1), x0) # 向量化迭代计算每一步 for i in range(n_steps): current_x = x[:, i] current_t = t[i] x[:, i+1] = current_x + mu(current_x, current_t) * dt + sigma(current_x, current_t) * dW_samples[:, i] return x def mc_sim(N, t0, t_end, dt, x0, mu, sigma): n_steps = int((t_end - t0) / dt) dW_samples = np.random.normal(0, np.sqrt(dt), (N, n_steps)) return sde_euler_vectorized(t0, t_end, dt, x0, mu, sigma, dW_samples)
这种方式效率远高于apply_along_axis,因为numpy的循环是底层C实现的,比Python循环或apply类方法快得多。
内容的提问来源于stack exchange,提问作者Landscape
相关产品推荐
相关产品推荐

