【问题标题】:Python - Different regular/analytic functionsPython - 不同的常规/分析函数
【发布时间】: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 方法

标签: python numpy scipy


【解决方案1】:

正如 cmets 中已经提到的,您不能使用像 scipy.stats.norm.cdf 这样的 jax 库之外的方法。请改用jax.scipy.stats。同样,将 expsqrt 替换为它们的 jax 等效项 jnp.expjnp.sqrt

from jax import jit, grad, vmap
import jax.numpy as jnp
from jax.scipy.stats.norm import cdf

def analytical_call(s0):
    T, q, r, k, sigma = 1.0, 0.0, 0.0, 1.0, 0.4
    Kt = k*jnp.exp((q-r)*T)
    d = (jnp.log(Kt/s0)+(sigma**2)/2*T)/sigma
    result = (Kt * cdf((d / jnp.sqrt(T)), 0.0, 1.0)  - s0 * cdf(((d - sigma * T) / jnp.sqrt(T)), 0.0, 1.0)  ) * jnp.exp(-q * T) + jnp.exp(-q * T) * (s0 - Kt)
    return result

g = vmap(grad(analytical_call))
h = vmap(grad(grad(analytical_call)))
xi = jnp.linspace(1,1.5)

然后,您可以评估g(xi)h(xi)

【讨论】:

  • 我已经编辑了这个问题,只是为了附上我尝试过的代码
  • @John_maddon 当然。但是,为了清楚起见,我建议您发布一个单独的问题,而不是编辑您之前的帖子。
猜你喜欢
  • 2011-07-19
  • 2011-09-26
  • 2021-04-23
  • 1970-01-01
  • 1970-01-01
  • 2011-04-29
  • 2017-02-03
  • 1970-01-01
  • 1970-01-01
相关资源
最近更新 更多