【发布时间】:2021-07-30 16:12:11
【问题描述】:
为了执行衍生,我开发了以下代码:
import matplotlib.pyplot as plt
import numpy as np
from math import *
xi = jnp.linspace(-3,3)
def f(x):
a = x**3+5
return a
g1i = jax.vmap(jax.grad(f))(xi)
g2i = jax.vmap(jax.grad(jax.grad(f)))(xi)
g3i = jax.vmap(jax.grad(jax.grad(jax.grad(f))))(xi)
plt.plot(xi,yi, label = "f")
plt.plot(xi,g1i, label = "f'")
plt.plot(xi,g2i, label = "f''")
plt.plot(xi,g3i, label = "f'''")
plt.legend()
此代码有效,但现在我有兴趣应用以下代码来计算关于标的资产(即增量)的看涨价格的一阶导数,尝试使用以下代码,但它没有作品:
import scipy.stats as si
import sympy as sy
import sys
xi = jnp.linspace(1,1.5)
def analytical_call(s0):
T=1.
q=0.
r=0.
k=1.
sigma=0.4
Kt = k*exp((q-r)*T)
d = (log(Kt/s0)+(sigma**2)/2*T)/sigma
result = (Kt * si.norm.cdf((d / sqrt(T)), 0.0, 1.0) - s0 * si.norm.cdf(((d - sigma * T) / sqrt(T)), 0.0, 1.0) ) * exp(-q * T) + exp(-q * T) * (s0 - Kt)
return result
print(analytical_call(1))
g1i = jax.vmap(jax.grad(analytical_call))(xi)
g2i = jax.vmap(jax.grad(jax.grad(analytical_call)))(xi)
plt.plot(xi,yi, label = "f")
plt.plot(xi,g1i, label = "f'")
plt.legend()
你有什么提示吗?提前致谢!
【问题讨论】:
-
您的函数
analytical_call中没有delta,因此不清楚要区分哪个变量。你的意思是s0吗?另请注意,您不能将 scipy.stats 和 sympy 方法与 jax 混合使用。 -
是的,我的意思是区分一个关于 s0 流的调用,在代码中定义为“xi”@joni
-
Scipy 仅用于计算 d1.. 我该如何解决这个问题?因为我需要调用对底层证券流的敏感性,所以使用 AAD 方法