from functools import partial import numpy as np import jax.numpy as jnp import jax @partial(jax.jit, static_argnums=(2,)) def f(m, b, d=128): i = jnp.arange(d / 2) return jnp.cos(m[:, None] * b ** (-2 * i[None] / d)).sum(axis=1) @np.vectorize def fmin(L, b): return f(np.arange(L), b).min() def bmin(L): B = 1000 * L for k in range(1, 6): bs = np.linspace(0, 1, 10**k + 1)[1:] * B ys = fmin(L, bs) for b, y in zip(bs, ys): if y >= 0: B = b break return B bmin(1024 * 128)