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

Jax版Heston模型路径生成慢于Numpy的原因及优化问询

问题原因与优化方案

核心原因

你的Jax版本代码效率低下主要有三个关键问题:

  1. 未利用JIT编译:Jax默认以解释模式执行Python代码,远不如Numpy的C实现高效;而Numpy的循环操作是底层C优化过的。
  2. 不可变数组的低效赋值:Jax数组是不可变的,S.at[i].set()每次都会创建全新数组,Python循环中反复执行会产生大量内存拷贝和解释层开销,而Numpy是原地修改数组,无此问题。
  3. Python显式循环的JIT缺陷:即使给函数加@jax.jit,Python循环会被编译器展开成庞大的计算图,编译时间长且运行效率远不如Jax原生的循环结构。

优化方案

通过JIT编译+**Jax原生循环lax.scan**重构代码,完全发挥Jax的性能优势:

优化后的代码

import jax.numpy as jnp
from jax import random, jit, lax

# 参数保持不变
S0 = 100.0          # 初始资产价格
K = 64              # 执行价格
T = 1.0             # 时间(年)
r = 0.05            # 无风险利率
N = 252             # 模拟时间步数
M = 100000          # 模拟次数

# Heston模型参数
kappa = 2           # 方差均值回归速率
theta = 0.05        # 方差长期均值
v0 = 0.05           # 初始方差
rho = -0.5          # 收益与方差的相关性
sigma = 0.3         # 波动率的波动率

@jit
def heston_model_sim_jax_opt(S0, v0, rho, kappa, theta, sigma, r, T, N, M):
    dt = T / N
    mu = jnp.array([0, 0])
    cov = jnp.array([[1, rho], [rho, 1]])
    
    # 生成相关布朗运动
    key = random.PRNGKey(0)
    Z = random.multivariate_normal(key, mean=mu, cov=cov, shape=(N, M))
    
    # 定义扫描迭代函数:输入上一步状态与当前布朗运动,输出新状态和当前步结果
    def step(carry, z):
        S_prev, v_prev = carry
        # 计算当前资产价格
        S_curr = S_prev * jnp.exp((r - 0.5 * v_prev) * dt + jnp.sqrt(v_prev * dt) * z[:, 0])
        # 计算当前方差(带截断避免负方差)
        v_curr = jnp.maximum(v_prev + kappa * (theta - v_prev) * dt + sigma * jnp.sqrt(v_prev * dt) * z[:, 1], 0)
        return (S_curr, v_curr), (S_curr, v_curr)
    
    # 初始状态:所有模拟路径的初始价格和方差
    initial_carry = (jnp.full(M, S0), jnp.full(M, v0))
    # 执行扫描,完成所有时间步的迭代
    final_carry, (S_steps, v_steps) = lax.scan(step, initial_carry, Z)
    
    # 拼接初始值与迭代结果,得到完整的(N+1, M)维度数组
    S = jnp.vstack([jnp.full(M, S0), S_steps])
    v = jnp.vstack([jnp.full(M, v0), v_steps])
    
    return S, v

优化点说明

  • @jit装饰器:将整个函数编译为机器码,彻底消除Python解释执行的开销,第一次运行会有编译时间,后续运行速度极快。
  • lax.scan替代Python循环:这是Jax专为序列迭代设计的原生操作,能被编译器优化为高效的向量化计算,既避免了不可变数组的频繁拷贝,也不会生成冗余的计算图。
  • 状态传递式迭代:通过carry传递上一步的价格和方差,不需要维护完整的历史数组再逐点赋值,内存效率和计算效率都显著提升。

性能测试提示

第一次调用优化后的函数会触发JIT编译(耗时约1-2秒),之后的重复调用速度会远快于Numpy版本(通常能达到数倍甚至十倍以上的加速),建议多次运行取平均时间以排除编译开销。

内容的提问来源于stack exchange,提问作者gd1234

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 21:50:20